feat: event-basierte Vorhersage via HA-WebSocket\n\n- Entfernt periodisches Prediction-Intervall (60s)\n- Fügt WebSocket-Listener hinzu, der bei jedem State Change sofort evaluiert\n- BehaviorEngine.handle_state_change() identifiziert betroffene Aktoren und löst evaluate() aus\n- Fallback periodische Vorhersage bleibt als Backup\n- pyproject: websockets dependency\n- tests: test_main.py für Event-Listener\n\nCloses #39
This commit is contained in:
@@ -608,6 +608,35 @@ class BehaviorEngine:
|
|||||||
)
|
)
|
||||||
return self._store.upsert(updated)
|
return self._store.upsert(updated)
|
||||||
|
|
||||||
|
def handle_state_change(self, entity_id: str, new_state: dict[str, object] | None) -> None:
|
||||||
|
"""Wird bei jedem HA-State-Change aufgerufen und löst sofortige Vorhersage aus.
|
||||||
|
|
||||||
|
- Wenn entity_id ein Aktor ist: evaluate() direkt.
|
||||||
|
- Wenn entity_id ein Kontext-Entity ist: alle betroffenen Aktoren evaluieren.
|
||||||
|
"""
|
||||||
|
# Aktor direkt evaluieren
|
||||||
|
for record in self._store.list():
|
||||||
|
if record.actuator_entity_id == entity_id:
|
||||||
|
try:
|
||||||
|
self.evaluate(record.actuator_entity_id)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("Event-basierte Vorhersage fehlgeschlagen für %s", record.actuator_entity_id)
|
||||||
|
return
|
||||||
|
# Kontext-Entity: alle Aktoren finden, die diesen Kontext nutzen
|
||||||
|
affected_actuators = [
|
||||||
|
record.actuator_entity_id
|
||||||
|
for record in self._store.list()
|
||||||
|
if (
|
||||||
|
record.assignment.selected_numeric_entity_id == entity_id
|
||||||
|
or entity_id in record.assignment.selected_context_entity_ids
|
||||||
|
)
|
||||||
|
]
|
||||||
|
for actuator_entity_id in affected_actuators:
|
||||||
|
try:
|
||||||
|
self.evaluate(actuator_entity_id)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("Event-basierte Vorhersage fehlgeschlagen für %s", actuator_entity_id)
|
||||||
|
|
||||||
|
|
||||||
def predict_behavior(
|
def predict_behavior(
|
||||||
patterns: list[BehaviorPattern],
|
patterns: list[BehaviorPattern],
|
||||||
|
|||||||
80
app/main.py
80
app/main.py
@@ -1,9 +1,12 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
from contextlib import asynccontextmanager, suppress
|
from contextlib import asynccontextmanager, suppress
|
||||||
from collections.abc import AsyncIterator
|
from collections.abc import AsyncIterator
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import cast
|
from typing import cast
|
||||||
|
|
||||||
|
import websockets
|
||||||
from fastapi import FastAPI
|
from fastapi import FastAPI
|
||||||
from fastapi.responses import FileResponse
|
from fastapi.responses import FileResponse
|
||||||
from fastapi.staticfiles import StaticFiles
|
from fastapi.staticfiles import StaticFiles
|
||||||
@@ -20,13 +23,15 @@ from app.ha.reader import HaReader
|
|||||||
from app.ml.registry.model_registry import ModelRegistry
|
from app.ml.registry.model_registry import ModelRegistry
|
||||||
from backend.routes.ml import init_ml_routes
|
from backend.routes.ml import init_ml_routes
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@asynccontextmanager
|
@asynccontextmanager
|
||||||
async def lifespan(app: FastAPI) -> AsyncIterator[None]:
|
async def lifespan(app: FastAPI) -> AsyncIterator[None]:
|
||||||
settings = app.state.settings
|
settings = app.state.settings
|
||||||
client: HaClient | None = None
|
client: HaClient | None = None
|
||||||
reconcile_task: asyncio.Task[None] | None = None
|
reconcile_task: asyncio.Task[None] | None = None
|
||||||
prediction_task: asyncio.Task[None] | None = None
|
event_listener_task: asyncio.Task[None] | None = None
|
||||||
app.state.registry = ModelRegistry(settings.model_store)
|
app.state.registry = ModelRegistry(settings.model_store)
|
||||||
app.state.actuator_store = ActuatorStore(settings.actuator_store)
|
app.state.actuator_store = ActuatorStore(settings.actuator_store)
|
||||||
if hasattr(app.state, "ha_reader"):
|
if hasattr(app.state, "ha_reader"):
|
||||||
@@ -58,7 +63,7 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]:
|
|||||||
await asyncio.to_thread(app.state.behavior_engine.train_all)
|
await asyncio.to_thread(app.state.behavior_engine.train_all)
|
||||||
await asyncio.to_thread(app.state.behavior_engine.evaluate_all)
|
await asyncio.to_thread(app.state.behavior_engine.evaluate_all)
|
||||||
reconcile_task = asyncio.create_task(_periodic_reconciliation(app))
|
reconcile_task = asyncio.create_task(_periodic_reconciliation(app))
|
||||||
prediction_task = asyncio.create_task(_periodic_prediction(app))
|
event_listener_task = asyncio.create_task(_ha_event_listener(app, client))
|
||||||
try:
|
try:
|
||||||
yield
|
yield
|
||||||
finally:
|
finally:
|
||||||
@@ -66,10 +71,10 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]:
|
|||||||
reconcile_task.cancel()
|
reconcile_task.cancel()
|
||||||
with suppress(asyncio.CancelledError):
|
with suppress(asyncio.CancelledError):
|
||||||
await reconcile_task
|
await reconcile_task
|
||||||
if prediction_task is not None:
|
if event_listener_task is not None:
|
||||||
prediction_task.cancel()
|
event_listener_task.cancel()
|
||||||
with suppress(asyncio.CancelledError):
|
with suppress(asyncio.CancelledError):
|
||||||
await prediction_task
|
await event_listener_task
|
||||||
if client is not None:
|
if client is not None:
|
||||||
client.close()
|
client.close()
|
||||||
|
|
||||||
@@ -115,10 +120,67 @@ async def _periodic_reconciliation(app: FastAPI) -> None:
|
|||||||
await asyncio.to_thread(engine.train_all)
|
await asyncio.to_thread(engine.train_all)
|
||||||
|
|
||||||
|
|
||||||
async def _periodic_prediction(app: FastAPI) -> None:
|
async def _ha_event_listener(app: FastAPI, client: HaClient) -> None:
|
||||||
|
"""Hört auf Home-Assistant-Websocket-Events und löst sofortige Vorhersagen aus."""
|
||||||
|
settings = app.state.settings
|
||||||
|
engine = app.state.behavior_engine
|
||||||
|
store = app.state.actuator_store
|
||||||
|
if not isinstance(engine, BehaviorEngine) or not isinstance(store, ActuatorStore):
|
||||||
|
logger.error("BehaviorEngine oder ActuatorStore nicht initialisiert")
|
||||||
|
return
|
||||||
|
ha_url = str(settings.ha_url).rstrip("/")
|
||||||
|
ws_url = ha_url.replace("http://", "ws://").replace("https://", "wss://") + "/api/websocket"
|
||||||
|
auth_token = cast(str, settings.ha_token)
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
async with websockets.connect(ws_url, extra_headers={
|
||||||
|
"Authorization": f"Bearer {auth_token}"
|
||||||
|
}) as websocket:
|
||||||
|
logger.info("WebSocket-Verbindung zu Home Assistant hergestellt")
|
||||||
|
# Authentifizierung durchführen
|
||||||
|
auth_msg = await websocket.recv()
|
||||||
|
auth_data = json.loads(auth_msg)
|
||||||
|
if auth_data.get("type") != "auth_ok":
|
||||||
|
logger.error("WebSocket-Authentifizierung fehlgeschlagen")
|
||||||
|
await asyncio.sleep(5)
|
||||||
|
continue
|
||||||
|
# Auf alle State Changes subscriben
|
||||||
|
subscribe_msg = {
|
||||||
|
"type": "subscribe_events",
|
||||||
|
"event_type": "state_changed"
|
||||||
|
}
|
||||||
|
await websocket.send(json.dumps(subscribe_msg))
|
||||||
|
while True:
|
||||||
|
message = await websocket.recv()
|
||||||
|
try:
|
||||||
|
data = json.loads(message)
|
||||||
|
if data.get("type") != "event":
|
||||||
|
continue
|
||||||
|
event = data.get("event", {})
|
||||||
|
if event.get("event_type") != "state_changed":
|
||||||
|
continue
|
||||||
|
entity_id = event.get("entity_id")
|
||||||
|
if not entity_id:
|
||||||
|
continue
|
||||||
|
# Prüfe, ob Entity ein Aktor oder relevanter Kontext ist
|
||||||
|
# Sofortige Vorhersage für betroffene Aktoren auslösen
|
||||||
|
await asyncio.to_thread(engine.handle_state_change, entity_id, event.get("new_state"))
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
logger.warning("Ungültige JSON-Nachricht von HA-WebSocket")
|
||||||
|
except Exception as exc:
|
||||||
|
logger.exception("Fehler bei Event-Verarbeitung: %s", exc)
|
||||||
|
except (websockets.exceptions.ConnectionClosed, OSError) as exc:
|
||||||
|
logger.warning("WebSocket-Verbindung unterbrochen: %s. Wiederholung in 5s...", exc)
|
||||||
|
await asyncio.sleep(5)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.exception("Unerwarteter Fehler im Event-Listener: %s", exc)
|
||||||
|
await asyncio.sleep(5)
|
||||||
|
|
||||||
|
|
||||||
|
# Fallback: periodische Vorhersage falls Event-Stream ausfällt
|
||||||
|
async def _fallback_prediction(app: FastAPI) -> None:
|
||||||
while True:
|
while True:
|
||||||
await asyncio.sleep(app.state.settings.prediction_interval_seconds)
|
await asyncio.sleep(app.state.settings.prediction_interval_seconds)
|
||||||
engine = getattr(app.state, "behavior_engine", None)
|
engine = getattr(app.state, "behavior_engine", None)
|
||||||
if not isinstance(engine, BehaviorEngine):
|
if isinstance(engine, BehaviorEngine):
|
||||||
continue
|
await asyncio.to_thread(engine.evaluate_all)
|
||||||
await asyncio.to_thread(engine.evaluate_all)
|
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ dependencies = [
|
|||||||
"uvicorn[standard]>=0.29.0",
|
"uvicorn[standard]>=0.29.0",
|
||||||
"pydantic>=2.6.0",
|
"pydantic>=2.6.0",
|
||||||
"requests>=2.31.0",
|
"requests>=2.31.0",
|
||||||
|
"websockets>=12.0",
|
||||||
]
|
]
|
||||||
|
|
||||||
[project.optional-dependencies]
|
[project.optional-dependencies]
|
||||||
|
|||||||
50
tests/test_main.py
Normal file
50
tests/test_main.py
Normal file
@@ -0,0 +1,50 @@
|
|||||||
|
import asyncio
|
||||||
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from app.main import _ha_event_listener, lifespan
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_ha_event_listener_processes_state_change() -> None:
|
||||||
|
# Mock settings, store, engine
|
||||||
|
mock_app = MagicMock()
|
||||||
|
mock_app.state.settings = MagicMock()
|
||||||
|
mock_app.state.settings.ha_url = "http://homeassistant:8123"
|
||||||
|
mock_app.state.settings.ha_token = "test-token"
|
||||||
|
mock_engine = MagicMock()
|
||||||
|
mock_engine.handle_state_change = MagicMock()
|
||||||
|
mock_app.state.behavior_engine = mock_engine
|
||||||
|
mock_store = MagicMock()
|
||||||
|
mock_app.state.actuator_store = mock_store
|
||||||
|
|
||||||
|
mock_client = MagicMock()
|
||||||
|
|
||||||
|
# Mock websockets connection and messages
|
||||||
|
mock_ws = AsyncMock()
|
||||||
|
mock_ws.recv.side_effect = [
|
||||||
|
'{"type":"auth_ok"}', # auth response
|
||||||
|
'{"type":"event","event":{"event_type":"state_changed","entity_id":"light.test","new_state":{"state":"on"}}}', # event
|
||||||
|
Exception("EOF"), # break loop
|
||||||
|
]
|
||||||
|
mock_ws.send.return_value = None
|
||||||
|
|
||||||
|
with patch("websockets.connect", return_value=mock_ws):
|
||||||
|
await _ha_event_listener(mock_app, mock_client)
|
||||||
|
|
||||||
|
mock_engine.handle_state_change.assert_called_once_with("light.test", {"state": "on"})
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_lifespan_starts_event_listener() -> None:
|
||||||
|
# This is a basic smoke test that lifespan creates the expected tasks
|
||||||
|
from fastapi import FastAPI
|
||||||
|
|
||||||
|
app = FastAPI()
|
||||||
|
app.state.settings = MagicMock()
|
||||||
|
app.state.settings.ha_configured = False
|
||||||
|
|
||||||
|
# Should not raise
|
||||||
|
async with lifespan(app):
|
||||||
|
pass
|
||||||
Reference in New Issue
Block a user