from __future__ import annotations import logging from datetime import datetime, timezone from typing import List, Sequence from fastapi import APIRouter, FastAPI, HTTPException, Request, status 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 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(request: Request) -> ModelsResponse: registry = _require_registry(request) 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(payload: PredictRequest, request: Request) -> PredictResponse: registry = _require_registry(request) predictor = Predictor(registry=registry) vector = FeatureVector(sensor_id=payload.sensor_id, values=payload.values) try: prediction = predictor.predict(payload.model_id, vector) except KeyError as exc: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc except ValueError as exc: raise HTTPException( status_code=status.HTTP_422_UNPROCESSABLE_CONTENT, detail=str(exc), ) from exc return PredictResponse( model_id=payload.model_id, sensor_id=payload.sensor_id, prediction=prediction, ) @router.post("/batch", response_model=BatchResponse, status_code=200) def predict_batch(payload: BatchRequest, request: Request) -> BatchResponse: registry = _require_registry(request) predictor = Predictor(registry=registry) responses: List[PredictResponse] = [] for item in payload.requests: vector = FeatureVector(sensor_id=item.sensor_id, values=item.values) try: prediction = predictor.predict(item.model_id, vector) except KeyError as exc: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc except ValueError as exc: raise HTTPException( status_code=status.HTTP_422_UNPROCESSABLE_CONTENT, detail=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(request: Request) -> ModelRegistry: registry = getattr(request.app.state, "registry", None) if not isinstance(registry, ModelRegistry): raise HTTPException( status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="ML registry nicht initialisiert.", ) return registry def init_ml_routes(app: FastAPI, model_store: str = ".model_store") -> None: registry = ModelRegistry(model_store) app.state.registry = registry app.include_router(router) logger.info("ML routes registered")