unify production app configuration and ML routes
This commit is contained in:
@@ -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
|
||||
def list_entities(ha_reader: HaReader = Depends(get_ha_reader)) -> List[HaEntitySummary]:
|
||||
return list(ha_reader.read_entities())
|
||||
|
||||
@@ -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"),
|
||||
)
|
||||
|
||||
34
app/main.py
34
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"}
|
||||
return {"service": "sillyhome-next", "docs": "/docs"}
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
50
tests/api/test_ml_routes.py
Normal file
50
tests/api/test_ml_routes.py
Normal 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
16
tests/test_config.py
Normal 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
|
||||
Reference in New Issue
Block a user