import json import os import shutil import tempfile from typing import Optional from fastapi import APIRouter, File, Form, HTTPException, Response, UploadFile from fastapi.responses import FileResponse from pydantic import BaseModel from backend import autolabel, dataset, library, projects from backend import batches as batch_store from backend import review as review_store from backend.api.common import project_or_404, thumbnail router = APIRouter(tags=["batches"]) class BatchRequest(BaseModel): rel: str start_sec: float = 0.0 end_sec: float fps: float = 1.0 class AutolabelRequest(BaseModel): engine: str = "sam3" engines: Optional[list[str]] = None class_ids: Optional[list[int]] = None engine_classes: Optional[dict[str, list[str]]] = None target_class_names: Optional[list[str]] = None threshold: float = autolabel.DEFAULT_THRESHOLD iou_threshold: float = autolabel.DEFAULT_IOU min_box_frac: float = 0.0 resume: bool = False append: bool = False custom_model_path: Optional[str] = None class Exemplar(BaseModel): box: list[float] # [cx, cy, w, h], normalized 0..1 positive: bool = True class PreviewRequest(BaseModel): frame_id: int engine: str threshold: float = autolabel.DEFAULT_THRESHOLD iou_threshold: float = autolabel.DEFAULT_IOU min_box_frac: float = 0.0 target_class_names: Optional[list[str]] = None custom_model_path: Optional[str] = None # Drawn box exemplars (REQ-172): normalized cxcywh, positive or negative. # Preview-only — `autolabel.start` deliberately has no equivalent. exemplars: Optional[list[Exemplar]] = None exemplar_class_name: Optional[str] = None @router.post("/api/projects/{project_id}/batches") def create_batch(project_id: int, request: BatchRequest) -> dict: try: return batch_store.create(project_id, request.rel, request.start_sec, request.end_sec, request.fps) except (batch_store.BatchError, library.LibraryError) as exc: raise HTTPException(400, str(exc)) @router.get("/api/projects/{project_id}/batches") def list_batches(project_id: int) -> dict: project_or_404(project_id) return {"batches": batch_store.listing(project_id)} class BatchPatch(BaseModel): batch_label: Optional[str] = None date_label: Optional[str] = None status: Optional[str] = None @router.get("/api/batches/{batch_id}") def read_batch(batch_id: int) -> dict: batch = batch_store.get(batch_id) if batch is None: raise HTTPException(404, "No such batch") return batch @router.patch("/api/batches/{batch_id}") def update_batch(batch_id: int, request: BatchPatch) -> dict: try: return batch_store.update(batch_id, request.model_dump(exclude_unset=True)) except batch_store.BatchError as exc: raise HTTPException(400, str(exc)) @router.delete("/api/batches/{batch_id}") def delete_batch(batch_id: int) -> dict: if not batch_store.delete(batch_id): raise HTTPException(404, "No such batch") return {"deleted": True} @router.get("/api/batches/{batch_id}/frames") def list_frames(batch_id: int) -> dict: if batch_store.get(batch_id) is None: raise HTTPException(404, "No such batch") return {"frames": batch_store.frames(batch_id)} @router.post("/api/batches/{batch_id}/autolabel") def start_autolabel(batch_id: int, request: AutolabelRequest) -> dict: try: engine_list = request.engines if (request.engines and len(request.engines) > 0) else [request.engine] return autolabel.start(batch_id, request.threshold, request.iou_threshold, request.min_box_frac, resume=request.resume, append=request.append, engine=request.engine, engines=engine_list, class_ids=request.class_ids, engine_classes=request.engine_classes, target_class_names=request.target_class_names, custom_model_path=request.custom_model_path) except batch_store.BatchError as exc: raise HTTPException(400, str(exc)) @router.post("/api/batches/inspect-model") async def inspect_model(file: UploadFile = File(...)) -> dict: if not (file.filename or "").endswith(".pt"): raise HTTPException(400, "Model must be a .pt file") with tempfile.NamedTemporaryFile(suffix=".pt", delete=False) as staged: shutil.copyfileobj(file.file, staged) staged_path = staged.name await file.close() try: classes = projects.read_model_classes(staged_path) return {"filename": file.filename, "classes": classes, "staged_path": staged_path} except Exception as exc: if os.path.exists(staged_path): os.unlink(staged_path) raise HTTPException(400, f"Invalid model: {exc}") @router.post("/api/batches/{batch_id}/autolabel-with-model") async def autolabel_with_model( batch_id: int, file: UploadFile = File(...), threshold: float = Form(0.35), iou_threshold: float = Form(0.0), selected_classes: str = Form("[]"), append: bool = Form(True), ) -> dict: if not (file.filename or "").endswith(".pt"): raise HTTPException(400, "Model must be a .pt file") with tempfile.NamedTemporaryFile(suffix=".pt", delete=False) as staged: shutil.copyfileobj(file.file, staged) staged_path = staged.name await file.close() try: target_classes = json.loads(selected_classes) if selected_classes else None return autolabel.start( batch_id, threshold=threshold, iou_threshold=iou_threshold, append=append, engine="custom", custom_model_path=staged_path, target_class_names=target_classes, ) except Exception as exc: if os.path.exists(staged_path): os.unlink(staged_path) raise HTTPException(400, f"Auto-annotation failed to start: {exc}") @router.post("/api/batches/{batch_id}/preview") def preview_autolabel(batch_id: int, request: PreviewRequest) -> dict: from backend import jobs, preview if not jobs.gpu_lock.acquire(timeout=20): busy = jobs.running_types() kind = busy[0] if busy else "background" raise HTTPException(409, f"The GPU is busy with a {kind} job — wait for it to finish") try: shapes = preview.preview_frame( batch_id=batch_id, frame_id=request.frame_id, engine=request.engine, threshold=request.threshold, iou_threshold=request.iou_threshold, min_box_frac=request.min_box_frac, target_class_names=request.target_class_names, custom_model_path=request.custom_model_path, exemplars=[e.model_dump() for e in request.exemplars or []], exemplar_class_name=request.exemplar_class_name, ) return {"shapes": shapes} except Exception as exc: raise HTTPException(400, str(exc)) finally: jobs.gpu_lock.release() @router.post("/api/sam3/playground-test") async def sam3_playground_test( file: UploadFile = File(...), prompts: str = Form(...), threshold: float = Form(0.35), iou_threshold: float = Form(0.0), ) -> dict: from PIL import Image from backend import labeling from backend.sam3_engine import get_engine try: image = Image.open(file.file).convert("RGB") except Exception as exc: raise HTTPException(400, f"Could not read image: {exc}") width, height = image.size prompt_list = [p.strip() for p in prompts.split(",") if p.strip()] if not prompt_list: raise HTTPException(400, "At least one text prompt is required") from backend import jobs if not jobs.gpu_lock.acquire(timeout=20): busy = jobs.running_types() kind = busy[0] if busy else "background" raise HTTPException(409, f"The GPU is busy with a {kind} job — wait for it to finish") try: engine = get_engine() raw_dets = engine.detect(image, prompt_list, threshold) kept_dets = labeling.deduplicate(raw_dets, iou_threshold=iou_threshold) except Exception as exc: raise HTTPException(500, f"SAM3 inference failed: {exc}") finally: jobs.gpu_lock.release() results = [] for det in kept_dets: norm_box = [ det.box[0] / width, det.box[1] / height, det.box[2] / width, det.box[3] / height, ] polys = [] if det.mask is not None: raw_polys = review_store.mask_to_polygons(det.mask) polys = [[[float(pt[0]), float(pt[1])] for pt in poly] for poly in raw_polys] results.append({ "class_id": det.class_id, "prompt": det.class_name, "score": round(float(det.score), 4), "box": [round(v, 5) for v in norm_box], "polygons": polys, }) return { "width": width, "height": height, "detections": results, } @router.post("/api/batches/{batch_id}/approve-all") def approve_all_batch_frames(batch_id: int) -> dict: if batch_store.get(batch_id) is None: raise HTTPException(404, "No such batch") updated = batch_store.approve_all_frames(batch_id) return {"approved_count": updated} class ApproveRequest(BaseModel): dataset_id: Optional[int] = None dataset_name: str = "" @router.post("/api/batches/{batch_ids}/approve") def approve_batch(batch_ids: str, request: ApproveRequest = ApproveRequest()) -> dict: """`batch_ids` is one id or a comma-separated selection — one merge, one dataset, however many batches Data Prep was tuned against (REQ-131).""" try: return dataset.approve(batch_ids, dataset_id=request.dataset_id, dataset_name=request.dataset_name) except dataset.DatasetError as exc: raise HTTPException(400, str(exc)) @router.get("/api/projects/{project_id}/dataset") def dataset_summary(project_id: int) -> dict: project_or_404(project_id) return dataset.summary(project_id) @router.get("/api/frames/{frame_id}/image") def frame_image(frame_id: int, w: int = 0): path = batch_store.frame_path(frame_id) if path is None or not os.path.isfile(path): raise HTTPException(404, "No such frame") if w and 16 <= w <= 2048: return Response(content=thumbnail(path, w), media_type="image/jpeg", headers={"Cache-Control": "public, max-age=3600"}) return FileResponse(path, media_type="image/jpeg", headers={"Cache-Control": "public, max-age=3600"}) @router.delete("/api/batches/{batch_id}/classes/{class_id}/annotations") def clear_batch_class_annotations(batch_id: int, class_id: int) -> dict: if batch_store.get(batch_id) is None: raise HTTPException(404, "No such batch") deleted = review_store.clear_batch_class_annotations(batch_id, class_id) return {"deleted": deleted} @router.post("/api/batches/{batch_id}/reset-auto-annotations") def reset_batch_auto_annotations(batch_id: int) -> dict: if batch_store.get(batch_id) is None: raise HTTPException(404, "No such batch") deleted = review_store.clear_batch_auto_annotations(batch_id) return {"deleted": deleted}