unify production app configuration and ML routes
This commit is contained in:
@@ -9,14 +9,10 @@ from app.ml.feature_store import FeatureStore, FeatureVector
|
||||
def create_app() -> FastAPI:
|
||||
application = FastAPI(title="SillyHome Next ML")
|
||||
init_ml_routes(application)
|
||||
_seed_default_model(application.state if hasattr(application, "state") else application)
|
||||
_seed_default_model(application.state)
|
||||
return application
|
||||
|
||||
|
||||
def _app_state(): # noqa: ANN001
|
||||
return app.state
|
||||
|
||||
|
||||
def _seed_default_model(state) -> None: # noqa: ANN001
|
||||
registry = getattr(state, "registry", None)
|
||||
if registry is None:
|
||||
@@ -34,4 +30,4 @@ def _seed_default_model(state) -> None: # noqa: ANN001
|
||||
registry.register(artifact)
|
||||
|
||||
|
||||
app = create_app()
|
||||
app = create_app()
|
||||
|
||||
@@ -4,13 +4,12 @@ import logging
|
||||
from datetime import datetime, timezone
|
||||
from typing import List, Sequence
|
||||
|
||||
from fastapi import APIRouter
|
||||
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
|
||||
from app.ml.training import TrainedArtifact
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -52,53 +51,67 @@ def health() -> HealthResponse:
|
||||
|
||||
|
||||
@router.get("/models", response_model=ModelsResponse, status_code=200)
|
||||
def list_models() -> ModelsResponse:
|
||||
registry = _require_registry()
|
||||
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(request: PredictRequest) -> PredictResponse:
|
||||
registry = _require_registry()
|
||||
def predict(payload: PredictRequest, request: Request) -> PredictResponse:
|
||||
registry = _require_registry(request)
|
||||
predictor = Predictor(registry=registry)
|
||||
vector = FeatureVector(sensor_id=request.sensor_id, values=request.values)
|
||||
vector = FeatureVector(sensor_id=payload.sensor_id, values=payload.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)
|
||||
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(request: BatchRequest) -> BatchResponse:
|
||||
registry = _require_registry()
|
||||
def predict_batch(payload: BatchRequest, request: Request) -> BatchResponse:
|
||||
registry = _require_registry(request)
|
||||
predictor = Predictor(registry=registry)
|
||||
responses: List[PredictResponse] = []
|
||||
for item in request.requests:
|
||||
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 _not_found_error(str(exc)) from 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() -> 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.")
|
||||
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) -> None: # noqa: ANN001
|
||||
registry = ModelRegistry(".model_store")
|
||||
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")
|
||||
logger.info("ML routes registered")
|
||||
|
||||
Reference in New Issue
Block a user