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 @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, engines=engine_list, class_ids=request.class_ids, engine_classes=request.engine_classes, target_class_names=request.target_class_names) 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"Could not inspect 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.8), 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, 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/sam3/playground-test") async def sam3_playground_test( file: UploadFile = File(...), prompts: str = Form(...), threshold: float = Form(0.35), iou_threshold: float = Form(0.8), ) -> 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") 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}") 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} @router.post("/api/batches/{batch_id}/approve") def approve_batch(batch_id: int) -> dict: try: return dataset.approve(batch_id) 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/projects/{project_id}/dataset/download") def dataset_download(project_id: int): project = project_or_404(project_id) try: path = dataset.zip_path(project) except dataset.DatasetError as exc: raise HTTPException(400, str(exc)) return FileResponse(path, media_type="application/zip", filename=f"{project['slug']}-dataset.zip") @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}