"""Triage routes: rules, hand overrides, and the per-batch shape view (REQ-100…108).""" import io import os import shutil from fastapi import APIRouter, File, Form, HTTPException, UploadFile from fastapi.responses import StreamingResponse from pydantic import BaseModel from typing import List, Optional from backend import augment from backend import batches as batch_store from backend import triage router = APIRouter() class AugmentRequest(BaseModel): settings: dict @router.get("/api/projects/{project_id}/base-datasets") def list_base_datasets(project_id: int) -> dict: from backend import base_dataset return {"base_datasets": base_dataset.listing(project_id)} @router.delete("/api/base-datasets/{base_id}") def delete_base_dataset(base_id: int) -> dict: from backend import base_dataset, projects record = base_dataset.get(base_id) if record is None: raise HTTPException(404, "No such base dataset") project = projects.get(record["project_id"]) if not base_dataset.delete(base_id, project["slug"]): raise HTTPException(404, "No such base dataset") return {"deleted": True} @router.get("/api/projects/{project_id}/augment") def get_augment(project_id: int) -> dict: try: return augment.get(project_id) except augment.AugmentError as exc: raise HTTPException(404, str(exc)) @router.put("/api/projects/{project_id}/augment") def put_augment(project_id: int, body: AugmentRequest) -> dict: try: return augment.save(project_id, body.settings) except augment.AugmentError as exc: raise HTTPException(400, str(exc)) class Rule(BaseModel): name: str = "" predicate: dict = {} action: str target_class: Optional[int] = None class RuleList(BaseModel): rules: List[Rule] class OverrideRequest(BaseModel): annotation_ids: List[int] verdict: str target_class: Optional[int] = None class ClearRequest(BaseModel): annotation_ids: List[int] @router.get("/api/projects/{project_id}/triage/rules") def get_rules(project_id: int) -> dict: return {"rules": triage.rules(project_id)} @router.put("/api/projects/{project_id}/triage/rules") def put_rules(project_id: int, body: RuleList) -> dict: try: stored = triage.replace_rules(project_id, [item.model_dump() for item in body.rules]) except triage.TriageError as exc: raise HTTPException(400, str(exc)) return {"rules": stored} @router.get("/api/batches/{batch_ids}/triage/summary") def batch_summary(batch_ids: str) -> dict: """`batch_ids` is one id or a comma-separated selection (REQ-130).""" try: return triage.batch_summary(batch_ids) except triage.TriageError as exc: raise HTTPException(404, str(exc)) @router.get("/api/batches/{batch_ids}/triage/shapes") def batch_page(batch_ids: str, sort: str = "score", offset: int = 0, limit: int = 120) -> dict: try: return triage.batch_page(batch_ids, sort=sort, offset=offset, limit=min(limit, 500)) except triage.TriageError as exc: raise HTTPException(400, str(exc)) @router.post("/api/triage/overrides") def set_overrides(body: OverrideRequest) -> dict: try: triage.set_overrides(body.annotation_ids, body.verdict, body.target_class) except triage.TriageError as exc: raise HTTPException(400, str(exc)) return {"updated": len(body.annotation_ids)} @router.delete("/api/triage/overrides") def clear_overrides(body: ClearRequest) -> dict: return {"cleared": triage.clear_overrides(body.annotation_ids)} @router.get("/api/batches/{batch_ids}/triage/suggest") def suggest(batch_ids: str) -> dict: try: return triage.suggest(batch_ids) except triage.TriageError as exc: raise HTTPException(404, str(exc)) @router.post("/api/batches/{batch_ids}/triage/simulate") def simulate(batch_ids: str, body: RuleList) -> dict: try: return triage.simulate(batch_ids, [item.model_dump() for item in body.rules]) except triage.TriageError as exc: raise HTTPException(400, str(exc)) @router.get("/api/projects/{project_id}/triage/preview") def preview(project_id: int) -> dict: return triage.preview(project_id) @router.get("/api/projects/{project_id}/export") def export_annotated(project_id: int, batch_ids: str = "", approved_only: bool = False, include_empty: bool = False): """Download annotated frames as a YOLO zip, merged or not — the user's own backup.""" from fastapi.responses import FileResponse from backend import export ids = [int(part) for part in batch_ids.split(",") if part.strip().isdigit()] try: path = export.build_zip(project_id, ids or None, approved_only=approved_only, include_empty=include_empty) except export.ExportError as exc: raise HTTPException(400, str(exc)) return FileResponse(path, media_type="application/zip", filename=os.path.basename(path)) @router.post("/api/projects/{project_id}/import") async def import_annotated(project_id: int, file: UploadFile = File(...), batch_label: str = Form("")) -> dict: """Load a previously exported zip back in, as a new batch to keep working on.""" import tempfile from backend import export staged = tempfile.NamedTemporaryFile(suffix=".zip", delete=False) try: shutil.copyfileobj(file.file, staged) staged.close() return export.restore_zip(project_id, staged.name, batch_label=batch_label) except export.ExportError as exc: raise HTTPException(400, str(exc)) except Exception as exc: raise HTTPException(400, f"Could not read that zip: {exc}") finally: if os.path.exists(staged.name): os.unlink(staged.name) @router.get("/api/annotations/{annotation_id}/crop") def crop(annotation_id: int, pad: float = 0.08): """The shape itself, cropped out of its frame — the crop grid judges objects, not whole frames (REQ-106).""" from PIL import Image from backend import db, review with db.cursor() as cur: cur.execute("SELECT frame_id, geometry FROM annotations WHERE id = ?", (annotation_id,)) row = cur.fetchone() if row is None: raise HTTPException(404, "No such annotation") import json box = review.to_box(json.loads(row["geometry"])) path = batch_store.frame_path(row["frame_id"]) if not path or not os.path.isfile(path): raise HTTPException(404, "The frame image is missing") with Image.open(path) as handle: image = handle.convert("RGB") width, height = image.size x0, y0, x1, y1 = box px, py = (x1 - x0) * pad, (y1 - y0) * pad crop_box = ( max(0, int((x0 - px) * width)), max(0, int((y0 - py) * height)), min(width, int((x1 + px) * width)), min(height, int((y1 + py) * height)), ) if crop_box[2] <= crop_box[0] or crop_box[3] <= crop_box[1]: raise HTTPException(400, "This shape has no area to crop") cropped = image.crop(crop_box) cropped.thumbnail((192, 192)) buffer = io.BytesIO() cropped.save(buffer, format="JPEG", quality=80) buffer.seek(0) return StreamingResponse(buffer, media_type="image/jpeg", headers={"Cache-Control": "public, max-age=86400"})