104 lines
3.1 KiB
Python
104 lines
3.1 KiB
Python
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") |