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 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
|
|
||||||
|
|||||||
@@ -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"),
|
||||||
)
|
)
|
||||||
|
|||||||
32
app/main.py
32
app/main.py
@@ -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")
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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")
|
||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
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