Files
sillyhome-next/backend/routes/ml.py

118 lines
3.7 KiB
Python

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