unify production app configuration and ML routes

This commit is contained in:
2026-06-11 21:07:58 +02:00
parent 4b3dc3b7af
commit 3bed5e790a
8 changed files with 127 additions and 74 deletions

View File

@@ -2,38 +2,20 @@ from __future__ import annotations
from typing import List from typing import List
from fastapi import APIRouter, HTTPException, Request from fastapi import APIRouter, Depends
from app.dependencies import get_ha_reader
from app.ha.models import HaEntitySummary from app.ha.models import HaEntitySummary
from app.ha.reader import HaReader from app.ha.reader import HaReader
from app.rules.recommender import Recommender
router = APIRouter(prefix="/v1", tags=["entities"]) router = APIRouter(prefix="/v1", tags=["entities"])
def _state_ha_reader(request: Request) -> HaReader:
try:
return request.app.state.ha_reader
except AttributeError as exc:
raise HTTPException(status_code=503, detail="HA-Reader nicht initialisiert.") from exc
def _state_recommender(request: Request) -> Recommender:
try:
return request.app.state.recommender
except AttributeError as exc:
raise HTTPException(status_code=503, detail="Recommender nicht initialisiert.") from exc
@router.get( @router.get(
"/entities", "/entities",
summary="Home-Assistant-Entities auflisten", summary="Home-Assistant-Entities auflisten",
description="Gibt eine kompakte Zusammenfassung aller erreichbaren HA-Entitäten zurück.", description="Gibt eine kompakte Zusammenfassung aller erreichbaren HA-Entitäten zurück.",
response_model=List[HaEntitySummary], response_model=List[HaEntitySummary],
) )
def list_entities(request: Request) -> List[HaEntitySummary]: def list_entities(ha_reader: HaReader = Depends(get_ha_reader)) -> List[HaEntitySummary]:
ha_reader = _state_ha_reader(request) return list(ha_reader.read_entities())
recommender = _state_recommender(request)
entities = ha_reader.read_entities()
recommender.run(entities)
return entities

View File

@@ -8,6 +8,7 @@ from dataclasses import dataclass
class Settings: class Settings:
ha_url: str | None = None ha_url: str | None = None
ha_token: str | None = None ha_token: str | None = None
model_store: str = ".model_store"
@property @property
def ha_configured(self) -> bool: def ha_configured(self) -> bool:
@@ -18,4 +19,5 @@ def load_settings() -> Settings:
return Settings( return Settings(
ha_url=os.getenv("SILLYHOME_HA_URL") or os.getenv("HA_URL"), ha_url=os.getenv("SILLYHOME_HA_URL") or os.getenv("HA_URL"),
ha_token=os.getenv("SILLYHOME_HA_TOKEN") or os.getenv("HA_TOKEN"), ha_token=os.getenv("SILLYHOME_HA_TOKEN") or os.getenv("HA_TOKEN"),
model_store=os.getenv("SILLYHOME_MODEL_STORE", ".model_store"),
) )

View File

@@ -3,43 +3,39 @@ from contextlib import asynccontextmanager
from fastapi import FastAPI from fastapi import FastAPI
from app.api.v1.entities import router as entities_router from app.api.v1.entities import router as entities_router
from app.config import load_settings
from app.core.exception_handlers import register_exception_handlers from app.core.exception_handlers import register_exception_handlers
from app.ha.client import HaClient, HaClientSettings from app.ha.client import HaClient, HaClientSettings
from app.ha.reader import HaReader from app.ha.reader import HaReader
from app.rules.recommender import Recommender from backend.routes.ml import init_ml_routes
from app.rules.heating import HeatingRule
@asynccontextmanager @asynccontextmanager
async def lifespan(app: FastAPI): async def lifespan(app: FastAPI): # type: ignore[no-untyped-def]
settings = app.state.settings settings = app.state.settings
ha_url = getattr(settings, "ha_url", None) if hasattr(app.state, "ha_reader"):
ha_token = getattr(settings, "ha_token", None) del app.state.ha_reader
client = HaClient( if settings.ha_configured:
settings=HaClientSettings( client = HaClient(
url=ha_url or "", settings=HaClientSettings(
token=ha_token or "", url=settings.ha_url,
token=settings.ha_token,
)
) )
) app.state.ha_reader = HaReader(client=client)
app.state.ha_reader = HaReader(client=client)
app.state.recommender = Recommender(rules=[HeatingRule()])
yield yield
class Settings:
ha_url: str = "http://localhost:8123"
ha_token: str = ""
app = FastAPI( app = FastAPI(
title="SillyHome Next API", title="SillyHome Next API",
description="Lokales Smart-Home-Intelligenzsystem für Home Assistant.", description="Lokales Smart-Home-Intelligenzsystem für Home Assistant.",
version="0.1.0", version="0.1.0",
lifespan=lifespan, lifespan=lifespan,
) )
app.state.settings = Settings() app.state.settings = load_settings()
register_exception_handlers(app) register_exception_handlers(app)
app.include_router(entities_router) app.include_router(entities_router)
init_ml_routes(app, model_store=app.state.settings.model_store)
@app.get("/health") @app.get("/health")

View File

@@ -9,14 +9,10 @@ from app.ml.feature_store import FeatureStore, FeatureVector
def create_app() -> FastAPI: def create_app() -> FastAPI:
application = FastAPI(title="SillyHome Next ML") application = FastAPI(title="SillyHome Next ML")
init_ml_routes(application) init_ml_routes(application)
_seed_default_model(application.state if hasattr(application, "state") else application) _seed_default_model(application.state)
return application return application
def _app_state(): # noqa: ANN001
return app.state
def _seed_default_model(state) -> None: # noqa: ANN001 def _seed_default_model(state) -> None: # noqa: ANN001
registry = getattr(state, "registry", None) registry = getattr(state, "registry", None)
if registry is None: if registry is None:

View File

@@ -4,13 +4,12 @@ import logging
from datetime import datetime, timezone from datetime import datetime, timezone
from typing import List, Sequence from typing import List, Sequence
from fastapi import APIRouter from fastapi import APIRouter, FastAPI, HTTPException, Request, status
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from app.ml.feature_store import FeatureVector from app.ml.feature_store import FeatureVector
from app.ml.predictor import Predictor from app.ml.predictor import Predictor
from app.ml.registry.model_registry import ModelRegistry from app.ml.registry.model_registry import ModelRegistry
from app.ml.training import TrainedArtifact
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -52,53 +51,67 @@ def health() -> HealthResponse:
@router.get("/models", response_model=ModelsResponse, status_code=200) @router.get("/models", response_model=ModelsResponse, status_code=200)
def list_models() -> ModelsResponse: def list_models(request: Request) -> ModelsResponse:
registry = _require_registry() registry = _require_registry(request)
models = [artifact.artifact_id for artifact in registry.list_models()] models = [artifact.artifact_id for artifact in registry.list_models()]
return ModelsResponse(models=models) return ModelsResponse(models=models)
@router.post("/predict", response_model=PredictResponse, status_code=200) @router.post("/predict", response_model=PredictResponse, status_code=200)
def predict(request: PredictRequest) -> PredictResponse: def predict(payload: PredictRequest, request: Request) -> PredictResponse:
registry = _require_registry() registry = _require_registry(request)
predictor = Predictor(registry=registry) 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: try:
prediction = predictor.predict(request.model_id, vector) prediction = predictor.predict(payload.model_id, vector)
except KeyError as exc: # unknown artifact 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
return PredictResponse(model_id=request.model_id, sensor_id=request.sensor_id, prediction=prediction) 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) @router.post("/batch", response_model=BatchResponse, status_code=200)
def predict_batch(request: BatchRequest) -> BatchResponse: def predict_batch(payload: BatchRequest, request: Request) -> BatchResponse:
registry = _require_registry() registry = _require_registry(request)
predictor = Predictor(registry=registry) predictor = Predictor(registry=registry)
responses: List[PredictResponse] = [] responses: List[PredictResponse] = []
for item in request.requests: for item in payload.requests:
vector = FeatureVector(sensor_id=item.sensor_id, values=item.values) vector = FeatureVector(sensor_id=item.sensor_id, values=item.values)
try: try:
prediction = predictor.predict(item.model_id, vector) prediction = predictor.predict(item.model_id, vector)
except KeyError as exc: 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( responses.append(
PredictResponse(model_id=item.model_id, sensor_id=item.sensor_id, prediction=prediction) PredictResponse(model_id=item.model_id, sensor_id=item.sensor_id, prediction=prediction)
) )
return BatchResponse(predictions=responses) return BatchResponse(predictions=responses)
def _require_registry() -> ModelRegistry: def _require_registry(request: Request) -> ModelRegistry:
from backend.app import _app_state registry = getattr(request.app.state, "registry", None)
if not isinstance(registry, ModelRegistry):
state_obj = _app_state() raise HTTPException(
registry = getattr(state_obj, "registry", None) status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
if registry is None: detail="ML registry nicht initialisiert.",
raise RuntimeError("ML registry nicht initialisiert.") )
return registry return registry
def init_ml_routes(app) -> None: # noqa: ANN001 def init_ml_routes(app: FastAPI, model_store: str = ".model_store") -> None:
registry = ModelRegistry(".model_store") registry = ModelRegistry(model_store)
app.state.registry = registry app.state.registry = registry
app.include_router(router) app.include_router(router)
logger.info("ML routes registered") logger.info("ML routes registered")

View File

@@ -49,8 +49,6 @@ def test_entities_returns_reader_data() -> None:
def test_entities_returns_503_without_home_assistant_config() -> None: def test_entities_returns_503_without_home_assistant_config() -> None:
with TestClient(app) as client: with TestClient(app) as client:
if hasattr(app.state, "ha_reader"):
delattr(app.state, "ha_reader")
response = client.get("/v1/entities") response = client.get("/v1/entities")
assert response.status_code == 503 assert response.status_code == 503

View File

@@ -0,0 +1,50 @@
from __future__ import annotations
from fastapi.testclient import TestClient
from app.main import app
def test_ml_routes_are_exposed_by_production_app() -> None:
with TestClient(app) as client:
health = client.get("/ml/health")
models = client.get("/ml/models")
assert health.status_code == 200
assert models.status_code == 200
assert models.json() == {"models": []}
def test_unknown_model_returns_404() -> None:
with TestClient(app) as client:
response = client.post(
"/ml/predict",
json={
"modelId": "missing",
"sensor_id": "sensor.kitchen",
"values": {"temperature": 21.0},
},
)
assert response.status_code == 404
def test_unsupported_sensor_returns_422(tmp_path) -> None:
from app.ml.registry.model_registry import ModelRegistry
from app.ml.training import TrainedArtifact
registry = ModelRegistry(tmp_path)
registry.register(TrainedArtifact("default", ("sensor.kitchen",)))
with TestClient(app) as client:
app.state.registry = registry
response = client.post(
"/ml/predict",
json={
"modelId": "default",
"sensor_id": "sensor.unknown",
"values": {"temperature": 21.0},
},
)
assert response.status_code == 422

16
tests/test_config.py Normal file
View File

@@ -0,0 +1,16 @@
from __future__ import annotations
from app.config import load_settings
def test_load_settings_reads_documented_environment(monkeypatch) -> None:
monkeypatch.setenv("SILLYHOME_HA_URL", "http://ha.local:8123")
monkeypatch.setenv("SILLYHOME_HA_TOKEN", "secret")
monkeypatch.setenv("SILLYHOME_MODEL_STORE", "/tmp/models")
settings = load_settings()
assert settings.ha_url == "http://ha.local:8123"
assert settings.ha_token == "secret"
assert settings.model_store == "/tmp/models"
assert settings.ha_configured