from __future__ import annotations import logging from datetime import datetime, timezone from typing import List, Sequence from fastapi import APIRouter from pydantic import BaseModel, Field from app.ml.feature_store import FeatureVector from app.ml.predictor import Predictor from app.ml.registry.model_registry import ModelRegistry from app.ml.training import TrainedArtifact logger = logging.getLogger(__name__) router = APIRouter(prefix="/ml", tags=["ml"]) class HealthResponse(BaseModel): status: str updated_at: datetime = Field(default_factory=lambda: datetime.now(timezone.utc)) class PredictRequest(BaseModel): model_id: str = Field(..., alias="modelId") sensor_id: str values: dict class PredictResponse(BaseModel): model_id: str sensor_id: str prediction: str class BatchRequest(BaseModel): requests: Sequence[PredictRequest] class BatchResponse(BaseModel): predictions: Sequence[PredictResponse] class ModelsResponse(BaseModel): models: List[str] @router.get("/health", response_model=HealthResponse, status_code=200) def health() -> HealthResponse: return HealthResponse(status="ok") @router.get("/models", response_model=ModelsResponse, status_code=200) def list_models() -> ModelsResponse: registry = _require_registry() models = [artifact.artifact_id for artifact in registry.list_models()] return ModelsResponse(models=models) @router.post("/predict", response_model=PredictResponse, status_code=200) def predict(request: PredictRequest) -> PredictResponse: registry = _require_registry() predictor = Predictor(registry=registry) vector = FeatureVector(sensor_id=request.sensor_id, values=request.values) try: prediction = predictor.predict(request.model_id, vector) except KeyError as exc: # unknown artifact raise _not_found_error(str(exc)) from exc return PredictResponse(model_id=request.model_id, sensor_id=request.sensor_id, prediction=prediction) @router.post("/batch", response_model=BatchResponse, status_code=200) def predict_batch(request: BatchRequest) -> BatchResponse: registry = _require_registry() predictor = Predictor(registry=registry) responses: List[PredictResponse] = [] for item in request.requests: vector = FeatureVector(sensor_id=item.sensor_id, values=item.values) try: prediction = predictor.predict(item.model_id, vector) except KeyError as exc: raise _not_found_error(str(exc)) from exc responses.append( PredictResponse(model_id=item.model_id, sensor_id=item.sensor_id, prediction=prediction) ) return BatchResponse(predictions=responses) def _require_registry() -> ModelRegistry: from backend.app import _app_state state_obj = _app_state() registry = getattr(state_obj, "registry", None) if registry is None: raise RuntimeError("ML registry nicht initialisiert.") return registry def init_ml_routes(app) -> None: # noqa: ANN001 registry = ModelRegistry(".model_store") app.state.registry = registry app.include_router(router) logger.info("ML routes registered")