diff --git a/.gitignore b/.gitignore index c101618..a06e3ca 100644 --- a/.gitignore +++ b/.gitignore @@ -21,6 +21,8 @@ checkpoints/ *.pt *.pth *.onnx +*.db +*.db-journal # OS .DS_Store diff --git a/labeling/backend/README.md b/labeling/backend/README.md new file mode 100644 index 0000000..74a36fe --- /dev/null +++ b/labeling/backend/README.md @@ -0,0 +1,58 @@ +# Labeling Tool — Backend + +API + data model for the bounding-box labeling tool. See CLAUDE.md §3.2 for the +full requirements this is built against. + +**Status:** core data model + CRUD API implemented (sets, hierarchical classes, +frames, labels, per-frame/per-set status). **Not yet implemented:** dataset +versioning/promotion (`enemy-v1`, `enemy-v2`, ...) — deferred since it needs its +own design pass (snapshot semantics: labels are editable at any time, but a +promoted dataset version must stay reproducible). Also not yet implemented: any +frontend, auth/login (see below), or the active-learning auto-label workflow. + +## Stack + +- **FastAPI** + **SQLAlchemy** (2.0), **SQLite** by default (`./labeling.db`), + swappable via the `DATABASE_URL` env var (e.g. to Postgres later without code + changes — one Postgres-compatible ORM). +- **Multi-user, no auth yet:** labels/status changes take a plain `created_by` / + `updated_by` username string, resolved via get-or-create (`users.py`). There's no + login flow — attribution only, since there's no UI yet that would need real auth. + +## Data model + +- `LabelSet` — a label set (`Enemy`, `Items`, `Traps`, ...), one per detector model. +- `MainClass` / `SubClass` — the per-set Main → Sub class hierarchy (e.g. `Enemy` → + `Bat`, `Snake`), created ad hoc via the API, no migration needed to add classes. +- `Frame` — one labelable image, identified by `(session_name, frame_index)` — + matches the recording tool's frame-extraction output 1:1. +- `Label` — one bounding box (`x, y, width, height` in pixel space), scoped to a + frame + set + sub-class. +- `FrameSetStatus` — per-frame, per-set label state (`unlabeled` / `auto_labeled` / + `reviewed`). + +## Setup + +``` +python -m venv .venv +source .venv/bin/activate +pip install -e ".[dev]" +``` + +## Run + +``` +spelunkai-labeling-backend +# or: uvicorn spelunkai_labeling_backend.main:create_app --factory --reload +``` + +Interactive API docs at `http://127.0.0.1:8000/docs` once running. + +## Testing + +``` +pytest +``` + +Each test gets a fully isolated app + SQLite file via `create_app(database_url=...)` +(see `tests/conftest.py`) — no shared state between tests, no real server needed. diff --git a/labeling/backend/pyproject.toml b/labeling/backend/pyproject.toml index 6f16138..2e23f88 100644 --- a/labeling/backend/pyproject.toml +++ b/labeling/backend/pyproject.toml @@ -3,7 +3,17 @@ name = "spelunkai-labeling-backend" version = "0.0.0" description = "SpelunkAI Labeling Tool backend: API + data model for bounding-box labels" requires-python = ">=3.10" -dependencies = [] +dependencies = [ + "fastapi>=0.115", + "uvicorn[standard]>=0.32", + "sqlalchemy>=2.0", +] + +[project.optional-dependencies] +dev = ["pytest>=8", "httpx>=0.27"] + +[project.scripts] +spelunkai-labeling-backend = "spelunkai_labeling_backend.cli:main" [build-system] requires = ["setuptools>=68"] diff --git a/labeling/backend/src/spelunkai_labeling_backend/__main__.py b/labeling/backend/src/spelunkai_labeling_backend/__main__.py new file mode 100644 index 0000000..9ae637f --- /dev/null +++ b/labeling/backend/src/spelunkai_labeling_backend/__main__.py @@ -0,0 +1,4 @@ +from .cli import main + +if __name__ == "__main__": + main() diff --git a/labeling/backend/src/spelunkai_labeling_backend/cli.py b/labeling/backend/src/spelunkai_labeling_backend/cli.py new file mode 100644 index 0000000..3bc4713 --- /dev/null +++ b/labeling/backend/src/spelunkai_labeling_backend/cli.py @@ -0,0 +1,14 @@ +"""Convenience entry point to run the labeling backend with uvicorn.""" +from __future__ import annotations + +import uvicorn + +from .main import create_app + + +def main() -> None: + uvicorn.run(create_app(), host="127.0.0.1", port=8000) + + +if __name__ == "__main__": + main() diff --git a/labeling/backend/src/spelunkai_labeling_backend/db.py b/labeling/backend/src/spelunkai_labeling_backend/db.py new file mode 100644 index 0000000..62d2ac7 --- /dev/null +++ b/labeling/backend/src/spelunkai_labeling_backend/db.py @@ -0,0 +1,40 @@ +"""Database engine/session setup. + +Each `Database` instance owns its own engine and session factory, so a +production run and a test can each get a fully isolated database instead of +sharing process-global state (see `main.create_app`). +""" +from __future__ import annotations + +import os +from typing import Iterator, Optional + +from sqlalchemy import create_engine +from sqlalchemy.orm import Session, sessionmaker + +from .models import Base + +DEFAULT_DATABASE_URL = "sqlite:///./labeling.db" + + +class Database: + def __init__(self, database_url: Optional[str] = None): + self.url = database_url or os.environ.get("DATABASE_URL", DEFAULT_DATABASE_URL) + connect_args = {"check_same_thread": False} if self.url.startswith("sqlite") else {} + self.engine = create_engine(self.url, connect_args=connect_args) + self.session_factory = sessionmaker(autocommit=False, autoflush=False, bind=self.engine) + + def init_models(self) -> None: + Base.metadata.create_all(bind=self.engine) + + def get_session(self) -> Iterator[Session]: + session = self.session_factory() + try: + yield session + finally: + session.close() + + +def get_db() -> Iterator[Session]: + """Marker dependency; `create_app` overrides this with a real `Database.get_session`.""" + raise NotImplementedError("get_db must be overridden via app.dependency_overrides") diff --git a/labeling/backend/src/spelunkai_labeling_backend/main.py b/labeling/backend/src/spelunkai_labeling_backend/main.py new file mode 100644 index 0000000..e6574f3 --- /dev/null +++ b/labeling/backend/src/spelunkai_labeling_backend/main.py @@ -0,0 +1,38 @@ +"""FastAPI app factory for the labeling backend. + +Run with `uvicorn spelunkai_labeling_backend.main:create_app --factory` (or +`spelunkai-labeling-backend`) so the app — and its database — is constructed +lazily instead of at import time; this also means importing this module has +no side effects, and tests can call `create_app(database_url=...)` to get a +fully isolated instance. +""" +from __future__ import annotations + +from typing import Optional + +from fastapi import FastAPI + +from .db import Database, get_db +from .routers.frames import router as frames_router +from .routers.labels import router as labels_router +from .routers.sets import main_classes_router, sets_router + + +def create_app(database_url: Optional[str] = None) -> FastAPI: + database = Database(database_url) + database.init_models() + + app = FastAPI(title="SpelunkAI Labeling Backend") + app.state.database = database + app.dependency_overrides[get_db] = database.get_session + + app.include_router(sets_router) + app.include_router(main_classes_router) + app.include_router(frames_router) + app.include_router(labels_router) + + @app.get("/health") + def health() -> dict: + return {"status": "ok"} + + return app diff --git a/labeling/backend/src/spelunkai_labeling_backend/models.py b/labeling/backend/src/spelunkai_labeling_backend/models.py new file mode 100644 index 0000000..692c86c --- /dev/null +++ b/labeling/backend/src/spelunkai_labeling_backend/models.py @@ -0,0 +1,121 @@ +"""SQLAlchemy ORM models for the labeling data model. + +Main/sub classes are data rows (not enum columns), so users can add new +classes on the fly per CLAUDE.md §3.2 without a schema migration. +""" +from __future__ import annotations + +import datetime +import enum + +from sqlalchemy import Column, DateTime, Enum as SqlEnum, Float, ForeignKey, Integer, String, UniqueConstraint +from sqlalchemy.orm import declarative_base, relationship + +Base = declarative_base() + + +def _utcnow() -> datetime.datetime: + return datetime.datetime.now(datetime.timezone.utc) + + +class LabelStatus(str, enum.Enum): + UNLABELED = "unlabeled" + AUTO_LABELED = "auto_labeled" + REVIEWED = "reviewed" + + +class LabelSource(str, enum.Enum): + MANUAL = "manual" + AUTO = "auto" + + +class User(Base): + __tablename__ = "users" + + id = Column(Integer, primary_key=True) + username = Column(String, unique=True, nullable=False) + created_at = Column(DateTime, default=_utcnow) + + +class LabelSet(Base): + __tablename__ = "label_sets" + + id = Column(Integer, primary_key=True) + name = Column(String, unique=True, nullable=False) + description = Column(String, nullable=True) + created_at = Column(DateTime, default=_utcnow) + + main_classes = relationship("MainClass", back_populates="label_set", cascade="all, delete-orphan") + + +class MainClass(Base): + __tablename__ = "main_classes" + __table_args__ = (UniqueConstraint("set_id", "name", name="uq_main_class_set_name"),) + + id = Column(Integer, primary_key=True) + set_id = Column(Integer, ForeignKey("label_sets.id"), nullable=False) + name = Column(String, nullable=False) + + label_set = relationship("LabelSet", back_populates="main_classes") + sub_classes = relationship("SubClass", back_populates="main_class", cascade="all, delete-orphan") + + +class SubClass(Base): + __tablename__ = "sub_classes" + __table_args__ = (UniqueConstraint("main_class_id", "name", name="uq_sub_class_main_name"),) + + id = Column(Integer, primary_key=True) + main_class_id = Column(Integer, ForeignKey("main_classes.id"), nullable=False) + name = Column(String, nullable=False) + + main_class = relationship("MainClass", back_populates="sub_classes") + + +class Frame(Base): + """A single labelable image, produced by the recording tool's frame extraction.""" + + __tablename__ = "frames" + __table_args__ = (UniqueConstraint("session_name", "frame_index", name="uq_frame_session_index"),) + + id = Column(Integer, primary_key=True) + session_name = Column(String, nullable=False) + frame_index = Column(Integer, nullable=False) + image_path = Column(String, nullable=False) + width = Column(Integer, nullable=False) + height = Column(Integer, nullable=False) + created_at = Column(DateTime, default=_utcnow) + + +class FrameSetStatus(Base): + """Per-image, per-set label state (unlabeled / auto_labeled / reviewed).""" + + __tablename__ = "frame_set_status" + __table_args__ = (UniqueConstraint("frame_id", "set_id", name="uq_frame_set_status"),) + + id = Column(Integer, primary_key=True) + frame_id = Column(Integer, ForeignKey("frames.id"), nullable=False) + set_id = Column(Integer, ForeignKey("label_sets.id"), nullable=False) + status = Column(SqlEnum(LabelStatus), nullable=False, default=LabelStatus.UNLABELED) + updated_at = Column(DateTime, default=_utcnow, onupdate=_utcnow) + updated_by_id = Column(Integer, ForeignKey("users.id"), nullable=True) + + +class Label(Base): + """A single bounding-box annotation, scoped to one frame within one label set.""" + + __tablename__ = "labels" + + id = Column(Integer, primary_key=True) + frame_id = Column(Integer, ForeignKey("frames.id"), nullable=False) + set_id = Column(Integer, ForeignKey("label_sets.id"), nullable=False) + sub_class_id = Column(Integer, ForeignKey("sub_classes.id"), nullable=False) + x = Column(Float, nullable=False) + y = Column(Float, nullable=False) + width = Column(Float, nullable=False) + height = Column(Float, nullable=False) + source = Column(SqlEnum(LabelSource), nullable=False, default=LabelSource.MANUAL) + created_by_id = Column(Integer, ForeignKey("users.id"), nullable=True) + created_at = Column(DateTime, default=_utcnow) + updated_at = Column(DateTime, default=_utcnow, onupdate=_utcnow) + + sub_class = relationship("SubClass") diff --git a/labeling/backend/src/spelunkai_labeling_backend/routers/__init__.py b/labeling/backend/src/spelunkai_labeling_backend/routers/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/labeling/backend/src/spelunkai_labeling_backend/routers/frames.py b/labeling/backend/src/spelunkai_labeling_backend/routers/frames.py new file mode 100644 index 0000000..83f290d --- /dev/null +++ b/labeling/backend/src/spelunkai_labeling_backend/routers/frames.py @@ -0,0 +1,50 @@ +"""Frame ingestion and listing.""" +from __future__ import annotations + +from typing import Optional + +from fastapi import APIRouter, Depends, HTTPException +from sqlalchemy.exc import IntegrityError +from sqlalchemy.orm import Session + +from .. import models, schemas +from ..db import get_db + +router = APIRouter(prefix="/frames", tags=["frames"]) + + +@router.post("", response_model=schemas.FrameRead, status_code=201) +def create_frame(payload: schemas.FrameCreate, db: Session = Depends(get_db)): + existing = ( + db.query(models.Frame) + .filter_by(session_name=payload.session_name, frame_index=payload.frame_index) + .first() + ) + if existing is not None: + return existing + + frame = models.Frame(**payload.model_dump()) + db.add(frame) + try: + db.commit() + except IntegrityError: + db.rollback() + raise HTTPException(409, "frame already exists") + db.refresh(frame) + return frame + + +@router.get("", response_model=list[schemas.FrameRead]) +def list_frames(session_name: Optional[str] = None, db: Session = Depends(get_db)): + query = db.query(models.Frame) + if session_name is not None: + query = query.filter_by(session_name=session_name) + return query.order_by(models.Frame.frame_index).all() + + +@router.get("/{frame_id}", response_model=schemas.FrameRead) +def get_frame(frame_id: int, db: Session = Depends(get_db)): + frame = db.get(models.Frame, frame_id) + if frame is None: + raise HTTPException(404, "frame not found") + return frame diff --git a/labeling/backend/src/spelunkai_labeling_backend/routers/labels.py b/labeling/backend/src/spelunkai_labeling_backend/routers/labels.py new file mode 100644 index 0000000..464a5be --- /dev/null +++ b/labeling/backend/src/spelunkai_labeling_backend/routers/labels.py @@ -0,0 +1,97 @@ +"""Bounding-box labels and per-frame/per-set label status.""" +from __future__ import annotations + +from fastapi import APIRouter, Depends, HTTPException +from sqlalchemy.orm import Session + +from .. import models, schemas +from ..db import get_db +from ..users import get_or_create_user + +router = APIRouter(tags=["labels"]) + + +def _require_frame_and_set(db: Session, frame_id: int, set_id: int) -> None: + if db.get(models.Frame, frame_id) is None: + raise HTTPException(404, "frame not found") + if db.get(models.LabelSet, set_id) is None: + raise HTTPException(404, "set not found") + + +@router.post("/frames/{frame_id}/sets/{set_id}/labels", response_model=schemas.LabelRead, status_code=201) +def create_label(frame_id: int, set_id: int, payload: schemas.LabelCreate, db: Session = Depends(get_db)): + _require_frame_and_set(db, frame_id, set_id) + + sub_class = db.get(models.SubClass, payload.sub_class_id) + if sub_class is None or sub_class.main_class.set_id != set_id: + raise HTTPException(422, "sub_class_id does not belong to this set") + + label = models.Label( + frame_id=frame_id, + set_id=set_id, + sub_class_id=payload.sub_class_id, + x=payload.x, + y=payload.y, + width=payload.width, + height=payload.height, + source=payload.source, + created_by_id=get_or_create_user(db, payload.created_by).id if payload.created_by else None, + ) + db.add(label) + db.commit() + db.refresh(label) + return label + + +@router.get("/frames/{frame_id}/sets/{set_id}/labels", response_model=list[schemas.LabelRead]) +def list_labels(frame_id: int, set_id: int, db: Session = Depends(get_db)): + _require_frame_and_set(db, frame_id, set_id) + return db.query(models.Label).filter_by(frame_id=frame_id, set_id=set_id).all() + + +@router.patch("/labels/{label_id}", response_model=schemas.LabelRead) +def update_label(label_id: int, payload: schemas.LabelUpdate, db: Session = Depends(get_db)): + label = db.get(models.Label, label_id) + if label is None: + raise HTTPException(404, "label not found") + for field, value in payload.model_dump(exclude_unset=True).items(): + setattr(label, field, value) + db.commit() + db.refresh(label) + return label + + +@router.delete("/labels/{label_id}", status_code=204) +def delete_label(label_id: int, db: Session = Depends(get_db)): + label = db.get(models.Label, label_id) + if label is None: + raise HTTPException(404, "label not found") + db.delete(label) + db.commit() + + +@router.get("/frames/{frame_id}/sets/{set_id}/status", response_model=schemas.FrameSetStatusRead) +def get_status(frame_id: int, set_id: int, db: Session = Depends(get_db)): + _require_frame_and_set(db, frame_id, set_id) + status = db.query(models.FrameSetStatus).filter_by(frame_id=frame_id, set_id=set_id).first() + if status is None: + status = models.FrameSetStatus(frame_id=frame_id, set_id=set_id, status=models.LabelStatus.UNLABELED) + db.add(status) + db.commit() + db.refresh(status) + return status + + +@router.put("/frames/{frame_id}/sets/{set_id}/status", response_model=schemas.FrameSetStatusRead) +def set_status(frame_id: int, set_id: int, payload: schemas.FrameSetStatusUpdate, db: Session = Depends(get_db)): + _require_frame_and_set(db, frame_id, set_id) + status = db.query(models.FrameSetStatus).filter_by(frame_id=frame_id, set_id=set_id).first() + if status is None: + status = models.FrameSetStatus(frame_id=frame_id, set_id=set_id) + db.add(status) + status.status = payload.status + if payload.updated_by: + status.updated_by_id = get_or_create_user(db, payload.updated_by).id + db.commit() + db.refresh(status) + return status diff --git a/labeling/backend/src/spelunkai_labeling_backend/routers/sets.py b/labeling/backend/src/spelunkai_labeling_backend/routers/sets.py new file mode 100644 index 0000000..1978805 --- /dev/null +++ b/labeling/backend/src/spelunkai_labeling_backend/routers/sets.py @@ -0,0 +1,63 @@ +"""Label sets and their Main -> Sub class hierarchy.""" +from __future__ import annotations + +from fastapi import APIRouter, Depends, HTTPException +from sqlalchemy.orm import Session + +from .. import models, schemas +from ..db import get_db + +sets_router = APIRouter(prefix="/sets", tags=["sets"]) +main_classes_router = APIRouter(prefix="/main-classes", tags=["sets"]) + + +@sets_router.post("", response_model=schemas.LabelSetRead, status_code=201) +def create_set(payload: schemas.LabelSetCreate, db: Session = Depends(get_db)): + if db.query(models.LabelSet).filter_by(name=payload.name).first(): + raise HTTPException(409, f"set '{payload.name}' already exists") + label_set = models.LabelSet(name=payload.name, description=payload.description) + db.add(label_set) + db.commit() + db.refresh(label_set) + return label_set + + +@sets_router.get("", response_model=list[schemas.LabelSetRead]) +def list_sets(db: Session = Depends(get_db)): + return db.query(models.LabelSet).all() + + +@sets_router.get("/{set_id}", response_model=schemas.LabelSetRead) +def get_set(set_id: int, db: Session = Depends(get_db)): + label_set = db.get(models.LabelSet, set_id) + if label_set is None: + raise HTTPException(404, "set not found") + return label_set + + +@sets_router.post("/{set_id}/main-classes", response_model=schemas.MainClassRead, status_code=201) +def create_main_class(set_id: int, payload: schemas.MainClassCreate, db: Session = Depends(get_db)): + label_set = db.get(models.LabelSet, set_id) + if label_set is None: + raise HTTPException(404, "set not found") + if db.query(models.MainClass).filter_by(set_id=set_id, name=payload.name).first(): + raise HTTPException(409, f"main class '{payload.name}' already exists in this set") + main_class = models.MainClass(set_id=set_id, name=payload.name) + db.add(main_class) + db.commit() + db.refresh(main_class) + return main_class + + +@main_classes_router.post("/{main_class_id}/sub-classes", response_model=schemas.SubClassRead, status_code=201) +def create_sub_class(main_class_id: int, payload: schemas.SubClassCreate, db: Session = Depends(get_db)): + main_class = db.get(models.MainClass, main_class_id) + if main_class is None: + raise HTTPException(404, "main class not found") + if db.query(models.SubClass).filter_by(main_class_id=main_class_id, name=payload.name).first(): + raise HTTPException(409, f"sub class '{payload.name}' already exists under this main class") + sub_class = models.SubClass(main_class_id=main_class_id, name=payload.name) + db.add(sub_class) + db.commit() + db.refresh(sub_class) + return sub_class diff --git a/labeling/backend/src/spelunkai_labeling_backend/schemas.py b/labeling/backend/src/spelunkai_labeling_backend/schemas.py new file mode 100644 index 0000000..5177137 --- /dev/null +++ b/labeling/backend/src/spelunkai_labeling_backend/schemas.py @@ -0,0 +1,109 @@ +"""Pydantic request/response schemas.""" +from __future__ import annotations + +from typing import Optional + +from pydantic import BaseModel, ConfigDict + +from .models import LabelSource, LabelStatus + + +class SubClassCreate(BaseModel): + name: str + + +class SubClassRead(BaseModel): + model_config = ConfigDict(from_attributes=True) + + id: int + name: str + + +class MainClassCreate(BaseModel): + name: str + + +class MainClassRead(BaseModel): + model_config = ConfigDict(from_attributes=True) + + id: int + name: str + sub_classes: list[SubClassRead] = [] + + +class LabelSetCreate(BaseModel): + name: str + description: Optional[str] = None + + +class LabelSetRead(BaseModel): + model_config = ConfigDict(from_attributes=True) + + id: int + name: str + description: Optional[str] = None + main_classes: list[MainClassRead] = [] + + +class FrameCreate(BaseModel): + session_name: str + frame_index: int + image_path: str + width: int + height: int + + +class FrameRead(BaseModel): + model_config = ConfigDict(from_attributes=True) + + id: int + session_name: str + frame_index: int + image_path: str + width: int + height: int + + +class LabelCreate(BaseModel): + sub_class_id: int + x: float + y: float + width: float + height: float + source: LabelSource = LabelSource.MANUAL + created_by: Optional[str] = None + + +class LabelUpdate(BaseModel): + sub_class_id: Optional[int] = None + x: Optional[float] = None + y: Optional[float] = None + width: Optional[float] = None + height: Optional[float] = None + + +class LabelRead(BaseModel): + model_config = ConfigDict(from_attributes=True) + + id: int + frame_id: int + set_id: int + sub_class_id: int + x: float + y: float + width: float + height: float + source: LabelSource + + +class FrameSetStatusRead(BaseModel): + model_config = ConfigDict(from_attributes=True) + + frame_id: int + set_id: int + status: LabelStatus + + +class FrameSetStatusUpdate(BaseModel): + status: LabelStatus + updated_by: Optional[str] = None diff --git a/labeling/backend/src/spelunkai_labeling_backend/users.py b/labeling/backend/src/spelunkai_labeling_backend/users.py new file mode 100644 index 0000000..bbc266c --- /dev/null +++ b/labeling/backend/src/spelunkai_labeling_backend/users.py @@ -0,0 +1,22 @@ +"""Lightweight user resolution. + +No login/auth flow yet — just get-or-create by username, so labels and +status changes can be attributed to a person (CLAUDE.md §3.2 multi-user +support) without building a full auth system before there's a UI to need it. +""" +from __future__ import annotations + +from sqlalchemy.orm import Session + +from . import models + + +def get_or_create_user(db: Session, username: str) -> models.User: + user = db.query(models.User).filter_by(username=username).first() + if user is not None: + return user + user = models.User(username=username) + db.add(user) + db.commit() + db.refresh(user) + return user diff --git a/labeling/backend/tests/conftest.py b/labeling/backend/tests/conftest.py new file mode 100644 index 0000000..b78434f --- /dev/null +++ b/labeling/backend/tests/conftest.py @@ -0,0 +1,12 @@ +import pytest +from fastapi.testclient import TestClient + +from spelunkai_labeling_backend.main import create_app + + +@pytest.fixture() +def client(tmp_path): + db_path = tmp_path / "test.db" + app = create_app(database_url=f"sqlite:///{db_path}") + with TestClient(app) as test_client: + yield test_client diff --git a/labeling/backend/tests/test_frames.py b/labeling/backend/tests/test_frames.py new file mode 100644 index 0000000..b78578b --- /dev/null +++ b/labeling/backend/tests/test_frames.py @@ -0,0 +1,32 @@ +def test_create_and_list_frames(client): + payload = { + "session_name": "run01", + "frame_index": 0, + "image_path": "run01_frames/frame_000000.png", + "width": 1280, + "height": 720, + } + resp = client.post("/frames", json=payload) + assert resp.status_code == 201 + frame_id = resp.json()["id"] + + resp = client.post("/frames", json=payload) + assert resp.status_code == 201 + assert resp.json()["id"] == frame_id + + resp = client.get("/frames", params={"session_name": "run01"}) + assert resp.status_code == 200 + assert len(resp.json()) == 1 + + +def test_list_frames_filters_by_session(client): + client.post("/frames", json={"session_name": "run01", "frame_index": 0, "image_path": "a.png", "width": 1, "height": 1}) + client.post("/frames", json={"session_name": "run02", "frame_index": 0, "image_path": "b.png", "width": 1, "height": 1}) + + resp = client.get("/frames", params={"session_name": "run02"}) + assert [f["image_path"] for f in resp.json()] == ["b.png"] + + +def test_get_missing_frame(client): + resp = client.get("/frames/999") + assert resp.status_code == 404 diff --git a/labeling/backend/tests/test_labels.py b/labeling/backend/tests/test_labels.py new file mode 100644 index 0000000..610ab2a --- /dev/null +++ b/labeling/backend/tests/test_labels.py @@ -0,0 +1,75 @@ +def _make_set_with_subclass(client, name="Enemy", sub_name="Bat"): + set_id = client.post("/sets", json={"name": name}).json()["id"] + main_class_id = client.post(f"/sets/{set_id}/main-classes", json={"name": name}).json()["id"] + sub_class_id = client.post(f"/main-classes/{main_class_id}/sub-classes", json={"name": sub_name}).json()["id"] + return set_id, sub_class_id + + +def _make_frame(client, session_name="run01", frame_index=0): + payload = {"session_name": session_name, "frame_index": frame_index, "image_path": "x.png", "width": 1280, "height": 720} + return client.post("/frames", json=payload).json()["id"] + + +def test_create_list_update_delete_label(client): + set_id, sub_class_id = _make_set_with_subclass(client) + frame_id = _make_frame(client) + + resp = client.post( + f"/frames/{frame_id}/sets/{set_id}/labels", + json={"sub_class_id": sub_class_id, "x": 10, "y": 20, "width": 30, "height": 40, "created_by": "jonas"}, + ) + assert resp.status_code == 201 + label_id = resp.json()["id"] + + resp = client.get(f"/frames/{frame_id}/sets/{set_id}/labels") + assert len(resp.json()) == 1 + + resp = client.patch(f"/labels/{label_id}", json={"x": 15}) + assert resp.status_code == 200 + assert resp.json()["x"] == 15 + + resp = client.delete(f"/labels/{label_id}") + assert resp.status_code == 204 + + resp = client.get(f"/frames/{frame_id}/sets/{set_id}/labels") + assert resp.json() == [] + + +def test_label_rejects_sub_class_from_another_set(client): + _, sub_class_id = _make_set_with_subclass(client, name="Enemy", sub_name="Bat") + other_set_id, _ = _make_set_with_subclass(client, name="Items", sub_name="Gold") + frame_id = _make_frame(client) + + resp = client.post( + f"/frames/{frame_id}/sets/{other_set_id}/labels", + json={"sub_class_id": sub_class_id, "x": 0, "y": 0, "width": 1, "height": 1}, + ) + assert resp.status_code == 422 + + +def test_labels_require_existing_frame_and_set(client): + _, sub_class_id = _make_set_with_subclass(client) + resp = client.post( + "/frames/999/sets/999/labels", + json={"sub_class_id": sub_class_id, "x": 0, "y": 0, "width": 1, "height": 1}, + ) + assert resp.status_code == 404 + + +def test_frame_set_status_defaults_and_updates(client): + set_id, _ = _make_set_with_subclass(client) + frame_id = _make_frame(client) + + resp = client.get(f"/frames/{frame_id}/sets/{set_id}/status") + assert resp.status_code == 200 + assert resp.json()["status"] == "unlabeled" + + resp = client.put( + f"/frames/{frame_id}/sets/{set_id}/status", + json={"status": "reviewed", "updated_by": "jonas"}, + ) + assert resp.status_code == 200 + assert resp.json()["status"] == "reviewed" + + resp = client.get(f"/frames/{frame_id}/sets/{set_id}/status") + assert resp.json()["status"] == "reviewed" diff --git a/labeling/backend/tests/test_sets.py b/labeling/backend/tests/test_sets.py new file mode 100644 index 0000000..2094f2a --- /dev/null +++ b/labeling/backend/tests/test_sets.py @@ -0,0 +1,47 @@ +def test_create_and_list_sets(client): + resp = client.post("/sets", json={"name": "Enemy", "description": "Enemy sprites"}) + assert resp.status_code == 201 + + resp = client.get("/sets") + assert resp.status_code == 200 + assert [s["name"] for s in resp.json()] == ["Enemy"] + + +def test_create_set_rejects_duplicate_name(client): + client.post("/sets", json={"name": "Enemy"}) + resp = client.post("/sets", json={"name": "Enemy"}) + assert resp.status_code == 409 + + +def test_get_missing_set(client): + resp = client.get("/sets/999") + assert resp.status_code == 404 + + +def test_create_main_and_sub_class(client): + set_id = client.post("/sets", json={"name": "Enemy"}).json()["id"] + + resp = client.post(f"/sets/{set_id}/main-classes", json={"name": "Enemy"}) + assert resp.status_code == 201 + main_class_id = resp.json()["id"] + + resp = client.post(f"/main-classes/{main_class_id}/sub-classes", json={"name": "Bat"}) + assert resp.status_code == 201 + assert resp.json()["name"] == "Bat" + + resp = client.get(f"/sets/{set_id}") + assert resp.json()["main_classes"][0]["sub_classes"][0]["name"] == "Bat" + + +def test_main_class_requires_existing_set(client): + resp = client.post("/sets/999/main-classes", json={"name": "Enemy"}) + assert resp.status_code == 404 + + +def test_sub_class_rejects_duplicate_name_under_same_main_class(client): + set_id = client.post("/sets", json={"name": "Enemy"}).json()["id"] + main_class_id = client.post(f"/sets/{set_id}/main-classes", json={"name": "Enemy"}).json()["id"] + client.post(f"/main-classes/{main_class_id}/sub-classes", json={"name": "Bat"}) + + resp = client.post(f"/main-classes/{main_class_id}/sub-classes", json={"name": "Bat"}) + assert resp.status_code == 409