from __future__ import annotations from uuid import UUID from fastapi import APIRouter, Depends from sqlalchemy.orm import Session from app.db.session import get_db from app.schemas import SegmentationQaRequest, SegmentationRunRequest from app.services.model_registry_service import ModelRegistryService from app.services.segmentation_service import SegmentationService from app.utils.response import envelope router = APIRouter(prefix="/segmentation", tags=["segmentation"]) @router.get("/models", response_model=dict) def list_segmentation_models() -> dict: return envelope({"models": [model.model_dump() for model in ModelRegistryService.list_model_capabilities(task_type="segmentation")]}) @router.post("/run", response_model=dict) def run_segmentation(payload: SegmentationRunRequest, db: Session = Depends(get_db)) -> dict: result = SegmentationService.run_segmentation( db=db, project_id=payload.project_id, dataset_id=payload.dataset_id, model_id=payload.model_id, confidence_threshold=payload.confidence_threshold, class_filter=payload.class_filter, tile_manifest_path=payload.tile_manifest_path, parameters_json=payload.parameters_json, ) return envelope(result.model_dump()) @router.get("/runs", response_model=dict) def list_segmentation_runs( project_id: UUID | None = None, dataset_id: UUID | None = None, db: Session = Depends(get_db), ) -> dict: return envelope(SegmentationService.list_runs(db, project_id=project_id, dataset_id=dataset_id).model_dump()) @router.get("/runs/{analysis_run_id}", response_model=dict) def get_segmentation_run(analysis_run_id: UUID, db: Session = Depends(get_db)) -> dict: return envelope(SegmentationService.get_run(db, analysis_run_id).model_dump()) @router.get("/runs/{analysis_run_id}/segmentations", response_model=dict) def list_segmentation_run_outputs( analysis_run_id: UUID, dataset_id: UUID | None = None, class_name: str | None = None, min_confidence: float | None = None, db: Session = Depends(get_db), ) -> dict: return envelope( SegmentationService.list_segmentations( db, analysis_run_id=analysis_run_id, dataset_id=dataset_id, class_name=class_name, min_confidence=min_confidence, ).model_dump() ) @router.get("/datasets/{dataset_id}/segmentations", response_model=dict) def list_dataset_segmentations( dataset_id: UUID, analysis_run_id: UUID | None = None, class_name: str | None = None, min_confidence: float | None = None, db: Session = Depends(get_db), ) -> dict: return envelope( SegmentationService.list_segmentations( db, analysis_run_id=analysis_run_id, dataset_id=dataset_id, class_name=class_name, min_confidence=min_confidence, ).model_dump() ) @router.get("/segmentations/{segmentation_id}", response_model=dict) def get_segmentation(segmentation_id: UUID, db: Session = Depends(get_db)) -> dict: return envelope(SegmentationService.get_segmentation(db, segmentation_id).model_dump()) @router.get("/runs/{analysis_run_id}/geojson", response_model=dict) def get_segmentation_run_geojson( analysis_run_id: UUID, class_name: str | None = None, min_confidence: float | None = None, db: Session = Depends(get_db), ) -> dict: return envelope( SegmentationService.segmentations_to_geojson( db, analysis_run_id=analysis_run_id, class_name=class_name, min_confidence=min_confidence, ) ) @router.get("/datasets/{dataset_id}/geojson", response_model=dict) def get_dataset_segmentation_geojson( dataset_id: UUID, analysis_run_id: UUID | None = None, class_name: str | None = None, min_confidence: float | None = None, db: Session = Depends(get_db), ) -> dict: return envelope( SegmentationService.segmentations_to_geojson( db, analysis_run_id=analysis_run_id, dataset_id=dataset_id, class_name=class_name, min_confidence=min_confidence, ) ) @router.post("/runs/{analysis_run_id}/qa/reference", response_model=dict) def compare_segmentation_run_with_reference( analysis_run_id: UUID, payload: SegmentationQaRequest, db: Session = Depends(get_db), ) -> dict: return envelope( SegmentationService.compare_segmentations_with_reference( db=db, analysis_run_id=analysis_run_id, reference_dataset_id=payload.reference_dataset_id, iou_threshold=payload.iou_threshold, class_name=payload.class_name, min_confidence=payload.min_confidence, ) )