Files
asus 5c7c122105 feat: add counting bench, triage, and dataset modules
This commit includes major additions and updates to the frontend and backend architectures, introducing new dataset management, live counting features, batch processing, and triage logic. Includes new UI pages, components, and API routes.
2026-08-14 16:28:52 +07:00

223 lines
7.3 KiB
Python

"""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"})