diff --git a/app/api/v1/entities.py b/app/api/v1/entities.py index e05ad73..2c08dac 100644 --- a/app/api/v1/entities.py +++ b/app/api/v1/entities.py @@ -2,38 +2,20 @@ from __future__ import annotations 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.reader import HaReader -from app.rules.recommender import Recommender 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( "/entities", summary="Home-Assistant-Entities auflisten", description="Gibt eine kompakte Zusammenfassung aller erreichbaren HA-Entitäten zurück.", response_model=List[HaEntitySummary], ) -def list_entities(request: Request) -> List[HaEntitySummary]: - ha_reader = _state_ha_reader(request) - recommender = _state_recommender(request) - entities = ha_reader.read_entities() - recommender.run(entities) - return entities \ No newline at end of file +def list_entities(ha_reader: HaReader = Depends(get_ha_reader)) -> List[HaEntitySummary]: + return list(ha_reader.read_entities()) diff --git a/app/config.py b/app/config.py index fd9d619..24de378 100644 --- a/app/config.py +++ b/app/config.py @@ -8,6 +8,7 @@ from dataclasses import dataclass class Settings: ha_url: str | None = None ha_token: str | None = None + model_store: str = ".model_store" @property def ha_configured(self) -> bool: @@ -18,4 +19,5 @@ def load_settings() -> Settings: return Settings( ha_url=os.getenv("SILLYHOME_HA_URL") or os.getenv("HA_URL"), ha_token=os.getenv("SILLYHOME_HA_TOKEN") or os.getenv("HA_TOKEN"), + model_store=os.getenv("SILLYHOME_MODEL_STORE", ".model_store"), ) diff --git a/app/main.py b/app/main.py index 4d8c33c..592c78a 100644 --- a/app/main.py +++ b/app/main.py @@ -3,43 +3,39 @@ from contextlib import asynccontextmanager from fastapi import FastAPI 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.ha.client import HaClient, HaClientSettings from app.ha.reader import HaReader -from app.rules.recommender import Recommender -from app.rules.heating import HeatingRule +from backend.routes.ml import init_ml_routes @asynccontextmanager -async def lifespan(app: FastAPI): +async def lifespan(app: FastAPI): # type: ignore[no-untyped-def] settings = app.state.settings - ha_url = getattr(settings, "ha_url", None) - ha_token = getattr(settings, "ha_token", None) - client = HaClient( - settings=HaClientSettings( - url=ha_url or "", - token=ha_token or "", + if hasattr(app.state, "ha_reader"): + del app.state.ha_reader + if settings.ha_configured: + client = HaClient( + settings=HaClientSettings( + url=settings.ha_url, + token=settings.ha_token, + ) ) - ) - app.state.ha_reader = HaReader(client=client) - app.state.recommender = Recommender(rules=[HeatingRule()]) + app.state.ha_reader = HaReader(client=client) yield -class Settings: - ha_url: str = "http://localhost:8123" - ha_token: str = "" - - app = FastAPI( title="SillyHome Next API", description="Lokales Smart-Home-Intelligenzsystem für Home Assistant.", version="0.1.0", lifespan=lifespan, ) -app.state.settings = Settings() +app.state.settings = load_settings() register_exception_handlers(app) app.include_router(entities_router) +init_ml_routes(app, model_store=app.state.settings.model_store) @app.get("/health") @@ -49,4 +45,4 @@ def health() -> dict[str, str]: @app.get("/") def root() -> dict[str, str]: - return {"service": "sillyhome-next", "docs": "/docs"} \ No newline at end of file + return {"service": "sillyhome-next", "docs": "/docs"} diff --git a/backend/app.py b/backend/app.py index af0d9a3..e61d761 100644 --- a/backend/app.py +++ b/backend/app.py @@ -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() \ No newline at end of file +app = create_app() diff --git a/backend/routes/ml.py b/backend/routes/ml.py index 7044724..700c452 100644 --- a/backend/routes/ml.py +++ b/backend/routes/ml.py @@ -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") \ No newline at end of file + logger.info("ML routes registered") diff --git a/tests/api/test_entities.py b/tests/api/test_entities.py index 6df2771..1c57112 100644 --- a/tests/api/test_entities.py +++ b/tests/api/test_entities.py @@ -49,8 +49,6 @@ def test_entities_returns_reader_data() -> None: def test_entities_returns_503_without_home_assistant_config() -> None: with TestClient(app) as client: - if hasattr(app.state, "ha_reader"): - delattr(app.state, "ha_reader") response = client.get("/v1/entities") assert response.status_code == 503 diff --git a/tests/api/test_ml_routes.py b/tests/api/test_ml_routes.py new file mode 100644 index 0000000..3358edb --- /dev/null +++ b/tests/api/test_ml_routes.py @@ -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 diff --git a/tests/test_config.py b/tests/test_config.py new file mode 100644 index 0000000..d81ee63 --- /dev/null +++ b/tests/test_config.py @@ -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