"""Training and model-version routes (REQ-060…065).""" import os from typing import Optional, Union from fastapi import APIRouter, HTTPException from fastapi.responses import FileResponse from pydantic import BaseModel from backend import hardware, training from backend.api.common import project_or_404 router = APIRouter(tags=["models"]) class TrainRequest(BaseModel): epochs: int = 50 batch: Optional[int] = None imgsz: Optional[int] = None device: Optional[Union[int, str]] = None batch_ids: Optional[list] = None class_ids: Optional[list] = None dataset_ids: Optional[list] = None base_dataset_ids: Optional[list] = None @router.get("/api/hardware") def read_hardware() -> dict: return hardware.defaults() @router.post("/api/projects/{project_id}/train") def start_training(project_id: int, request: TrainRequest) -> dict: project_or_404(project_id) try: return training.start( project_id, request.epochs, {"batch": request.batch, "imgsz": request.imgsz, "device": request.device}, batch_ids=request.batch_ids, class_ids=request.class_ids, dataset_ids=request.dataset_ids, base_dataset_ids=request.base_dataset_ids, ) except training.TrainingError as exc: raise HTTPException(400, str(exc)) @router.get("/api/projects/{project_id}/models") def list_models(project_id: int) -> dict: project = project_or_404(project_id) return {"models": training.listing(project_id), "base_model_kind": project["base_model_kind"]} @router.get("/api/models/{model_id}/weights") def download_weights(model_id: int): version = training.get_version(model_id) if version is None or not os.path.isfile(version["weights_path"]): raise HTTPException(404, "No weights for that version") name = version.get('name') or f"v{version['version']}" return FileResponse(version["weights_path"], media_type="application/octet-stream", filename=f"{name}-best.pt") @router.post("/api/models/{model_id}/promote") def promote_model(model_id: int) -> dict: try: return training.promote(model_id) except training.TrainingError as exc: raise HTTPException(400, str(exc)) class RenameRequest(BaseModel): name: str @router.patch("/api/models/{model_id}/rename") def rename_model(model_id: int, request: RenameRequest) -> dict: try: version = training.rename(model_id, request.name) if version is None: raise HTTPException(404, "No such model version") return version except training.TrainingError as exc: raise HTTPException(400, str(exc))