Compare commits
5 Commits
v0.1.0-rc1
...
feature/ha
| Author | SHA1 | Date | |
|---|---|---|---|
| 816a516106 | |||
| 1fbed37126 | |||
| dd496f9cc3 | |||
| 74b75de0fa | |||
| 840c404c1c |
1
.gitignore
vendored
1
.gitignore
vendored
@@ -4,6 +4,7 @@
|
||||
/.vscode
|
||||
__pycache__/
|
||||
*.pyc
|
||||
*.egg-info/
|
||||
.mypy_cache/
|
||||
.pytest_cache/
|
||||
.ruff_cache/
|
||||
|
||||
@@ -1,6 +1,11 @@
|
||||
# Changelog
|
||||
|
||||
## Unreleased
|
||||
- Klassifizierte Home-Assistant-Entity-Discovery mit Lernrelevanz und Filtern
|
||||
- Validierter Zugriff auf die Home-Assistant-History-API
|
||||
- Normalisierte, chronologisch sortierte numerische Zeitreihen über `/v1/history`
|
||||
|
||||
## 0.1.0 - 2026-06-13
|
||||
- Projektinitiierung
|
||||
- Architektur, ADRs und Roadmap
|
||||
- Einheitliche produktive FastAPI-App für HA- und ML-Routen
|
||||
@@ -8,3 +13,4 @@
|
||||
- Persistente, validierte und gegen Path Traversal gehärtete Model Registry
|
||||
- Reproduzierbares Packaging, CI-Gates und gehärteter non-root Container
|
||||
- Definierte API-Fehler und korrigierte Evaluationsmetriken
|
||||
- Scheduler-tauglicher Retraining-Service mit API und atomischem Registry-Update
|
||||
|
||||
@@ -42,7 +42,10 @@ uvicorn app.main:app --reload
|
||||
- `http://127.0.0.1:8000/health` - Health-Check
|
||||
- `http://127.0.0.1:8000/docs/` - OpenAPI-Dokumentation
|
||||
- `http://127.0.0.1:8000/v1/entities` - Home-Assistant-Entities
|
||||
- `http://127.0.0.1:8000/v1/discovery` - klassifizierte, filterbare Entities
|
||||
- `http://127.0.0.1:8000/v1/history` - normalisierte numerische Zeitreihen
|
||||
- `http://127.0.0.1:8000/ml/health` - Registry-/Serving-Health
|
||||
- `POST http://127.0.0.1:8000/ml/retrain` - Modell-Metadaten aktualisieren
|
||||
|
||||
Ohne vollständige HA-Konfiguration liefert `/v1/entities` bewusst `503`.
|
||||
|
||||
|
||||
@@ -1,10 +1,13 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import List
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||
|
||||
from app.dependencies import get_ha_reader
|
||||
from app.ha.discovery import DiscoveredEntity
|
||||
from app.ha.history import EntityHistorySeries
|
||||
from app.ha.models import HaEntitySummary
|
||||
from app.ha.reader import HaReader
|
||||
|
||||
@@ -19,3 +22,43 @@ router = APIRouter(prefix="/v1", tags=["entities"])
|
||||
)
|
||||
def list_entities(ha_reader: HaReader = Depends(get_ha_reader)) -> List[HaEntitySummary]:
|
||||
return list(ha_reader.read_entities())
|
||||
|
||||
|
||||
@router.get(
|
||||
"/discovery",
|
||||
summary="Home-Assistant-Entities klassifizieren",
|
||||
description="Klassifiziert Entities nach Lernrelevanz, Kontextquelle und Aktor-Rolle.",
|
||||
response_model=List[DiscoveredEntity],
|
||||
)
|
||||
def discovery(
|
||||
domain: List[str] | None = Query(default=None),
|
||||
learnable: bool | None = None,
|
||||
ha_reader: HaReader = Depends(get_ha_reader),
|
||||
) -> List[DiscoveredEntity]:
|
||||
return list(
|
||||
ha_reader.discover(
|
||||
domains=set(domain) if domain else None,
|
||||
learnable=learnable,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/history",
|
||||
summary="Numerische Home-Assistant-Historie lesen",
|
||||
description="Lädt und normalisiert numerische Zustände ausgewählter Entities.",
|
||||
response_model=List[EntityHistorySeries],
|
||||
)
|
||||
def history(
|
||||
entity_id: List[str] = Query(),
|
||||
start_time: datetime = Query(),
|
||||
end_time: datetime = Query(),
|
||||
ha_reader: HaReader = Depends(get_ha_reader),
|
||||
) -> List[EntityHistorySeries]:
|
||||
try:
|
||||
return list(ha_reader.read_history(entity_id, start_time, end_time))
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_CONTENT,
|
||||
detail=str(exc),
|
||||
) from exc
|
||||
|
||||
@@ -2,6 +2,9 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
import re
|
||||
from urllib.parse import quote
|
||||
|
||||
import requests
|
||||
|
||||
@@ -14,6 +17,9 @@ from app.ha.exceptions import (
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_ENTITY_ID_PATTERN = re.compile(r"^[a-z0-9_]+\.[a-z0-9_]+$")
|
||||
_MAX_HISTORY_SECONDS = 31 * 24 * 60 * 60
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class HaClientSettings:
|
||||
@@ -35,9 +41,58 @@ class HaClient:
|
||||
self._session.close()
|
||||
|
||||
def list_entities(self) -> list[dict[str, object]]:
|
||||
payload = self._get_json("/api/states")
|
||||
if not isinstance(payload, list):
|
||||
raise HaUnexpectedPayloadError(
|
||||
"Antwort von Home Assistant hat unerwartetes Format."
|
||||
)
|
||||
return payload
|
||||
|
||||
def get_history(
|
||||
self,
|
||||
entity_ids: list[str],
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
) -> list[object]:
|
||||
if not entity_ids:
|
||||
raise ValueError("Mindestens eine entity_id ist erforderlich.")
|
||||
if len(entity_ids) > 100:
|
||||
raise ValueError("Es können höchstens 100 Entities abgefragt werden.")
|
||||
if any(not _ENTITY_ID_PATTERN.fullmatch(entity_id) for entity_id in entity_ids):
|
||||
raise ValueError("entity_id enthält ein ungültiges Format.")
|
||||
if start_time.tzinfo is None or end_time.tzinfo is None:
|
||||
raise ValueError("start_time und end_time müssen eine Zeitzone enthalten.")
|
||||
if end_time <= start_time:
|
||||
raise ValueError("end_time muss nach start_time liegen.")
|
||||
if (end_time - start_time).total_seconds() > _MAX_HISTORY_SECONDS:
|
||||
raise ValueError("History-Abfragen sind auf 31 Tage begrenzt.")
|
||||
|
||||
start = quote(start_time.isoformat(), safe=":+")
|
||||
payload = self._get_json(
|
||||
f"/api/history/period/{start}",
|
||||
params={
|
||||
"filter_entity_id": ",".join(entity_ids),
|
||||
"end_time": end_time.isoformat(),
|
||||
"minimal_response": "1",
|
||||
"no_attributes": "1",
|
||||
},
|
||||
)
|
||||
if not isinstance(payload, list):
|
||||
raise HaUnexpectedPayloadError(
|
||||
"History-Antwort von Home Assistant hat unerwartetes Format."
|
||||
)
|
||||
return payload
|
||||
|
||||
def _get_json(
|
||||
self,
|
||||
path: str,
|
||||
*,
|
||||
params: dict[str, str] | None = None,
|
||||
) -> object:
|
||||
try:
|
||||
response = self._session.get(
|
||||
f"{self._settings.url.rstrip('/')}/api/states",
|
||||
f"{self._settings.url.rstrip('/')}{path}",
|
||||
params=params,
|
||||
timeout=self._settings.timeout_seconds,
|
||||
)
|
||||
except requests.Timeout as exc:
|
||||
@@ -69,9 +124,4 @@ class HaClient:
|
||||
"Antwort von Home Assistant ist kein gültiges JSON."
|
||||
) from exc
|
||||
|
||||
if not isinstance(payload, list):
|
||||
raise HaUnexpectedPayloadError(
|
||||
"Antwort von Home Assistant hat unerwartetes Format."
|
||||
)
|
||||
|
||||
return payload
|
||||
|
||||
184
app/ha/discovery.py
Normal file
184
app/ha/discovery.py
Normal file
@@ -0,0 +1,184 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import StrEnum
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from app.ha.models import HaEntitySummary
|
||||
|
||||
|
||||
class EntityRole(StrEnum):
|
||||
MEASUREMENT = "measurement"
|
||||
BINARY_CONTEXT = "binary_context"
|
||||
CONTEXT = "context"
|
||||
ACTUATOR = "actuator"
|
||||
UNSUPPORTED = "unsupported"
|
||||
|
||||
|
||||
class DiscoveredEntity(BaseModel):
|
||||
entity_id: str
|
||||
domain: str
|
||||
device_class: str | None = None
|
||||
state_class: str | None = None
|
||||
unit_of_measurement: str | None = None
|
||||
role: EntityRole
|
||||
learnable: bool
|
||||
reason: str
|
||||
|
||||
|
||||
_MEASUREMENT_CLASSES = frozenset({
|
||||
"apparent_power",
|
||||
"atmospheric_pressure",
|
||||
"battery",
|
||||
"carbon_dioxide",
|
||||
"carbon_monoxide",
|
||||
"current",
|
||||
"distance",
|
||||
"duration",
|
||||
"energy",
|
||||
"frequency",
|
||||
"gas",
|
||||
"humidity",
|
||||
"illuminance",
|
||||
"moisture",
|
||||
"monetary",
|
||||
"nitrogen_dioxide",
|
||||
"nitrogen_monoxide",
|
||||
"nitrous_oxide",
|
||||
"ozone",
|
||||
"pm1",
|
||||
"pm10",
|
||||
"pm25",
|
||||
"power",
|
||||
"precipitation",
|
||||
"pressure",
|
||||
"reactive_power",
|
||||
"signal_strength",
|
||||
"sound_pressure",
|
||||
"speed",
|
||||
"sulphur_dioxide",
|
||||
"temperature",
|
||||
"volatile_organic_compounds",
|
||||
"voltage",
|
||||
"volume",
|
||||
"volume_flow_rate",
|
||||
"water",
|
||||
"weight",
|
||||
"wind_speed",
|
||||
})
|
||||
_BINARY_CONTEXT_CLASSES = frozenset({
|
||||
"door",
|
||||
"garage_door",
|
||||
"lock",
|
||||
"motion",
|
||||
"occupancy",
|
||||
"opening",
|
||||
"presence",
|
||||
"problem",
|
||||
"safety",
|
||||
"smoke",
|
||||
"sound",
|
||||
"vibration",
|
||||
"window",
|
||||
})
|
||||
_ACTUATOR_DOMAINS = frozenset({
|
||||
"button",
|
||||
"climate",
|
||||
"cover",
|
||||
"fan",
|
||||
"humidifier",
|
||||
"light",
|
||||
"lock",
|
||||
"scene",
|
||||
"select",
|
||||
"siren",
|
||||
"switch",
|
||||
"valve",
|
||||
})
|
||||
_CONTEXT_DOMAINS = frozenset({"device_tracker", "person", "sun", "weather", "zone"})
|
||||
_LEARNABLE_CONTEXT_DOMAINS = frozenset({"device_tracker", "person", "weather"})
|
||||
_NUMERIC_STATE_CLASSES = frozenset({"measurement", "total", "total_increasing"})
|
||||
|
||||
|
||||
def classify_entity(entity: HaEntitySummary) -> DiscoveredEntity:
|
||||
if entity.domain == "sensor" and (
|
||||
entity.state_class in _NUMERIC_STATE_CLASSES
|
||||
or entity.device_class in _MEASUREMENT_CLASSES
|
||||
or entity.unit_of_measurement is not None
|
||||
):
|
||||
return _result(
|
||||
entity,
|
||||
EntityRole.MEASUREMENT,
|
||||
learnable=True,
|
||||
reason="Numerischer Messsensor für Zeitreihen und Training.",
|
||||
)
|
||||
|
||||
if entity.domain == "binary_sensor" and entity.device_class in _BINARY_CONTEXT_CLASSES:
|
||||
return _result(
|
||||
entity,
|
||||
EntityRole.BINARY_CONTEXT,
|
||||
learnable=True,
|
||||
reason="Binärer Kontextsensor für Zustands- und Anwesenheitsmuster.",
|
||||
)
|
||||
|
||||
if entity.domain in _CONTEXT_DOMAINS:
|
||||
learnable = entity.domain in _LEARNABLE_CONTEXT_DOMAINS
|
||||
return _result(
|
||||
entity,
|
||||
EntityRole.CONTEXT,
|
||||
learnable=learnable,
|
||||
reason=(
|
||||
"Kontextquelle für Training und Erklärungen."
|
||||
if learnable
|
||||
else "Kontextquelle ohne direkte Trainingsfreigabe."
|
||||
),
|
||||
)
|
||||
|
||||
if entity.domain in _ACTUATOR_DOMAINS:
|
||||
return _result(
|
||||
entity,
|
||||
EntityRole.ACTUATOR,
|
||||
learnable=False,
|
||||
reason="Aktor ist ein mögliches Automationsziel, aber kein Trainingssensor.",
|
||||
)
|
||||
|
||||
return _result(
|
||||
entity,
|
||||
EntityRole.UNSUPPORTED,
|
||||
learnable=False,
|
||||
reason="Entity-Typ ist noch nicht für Lernen oder Automationen klassifiziert.",
|
||||
)
|
||||
|
||||
|
||||
def discover_entities(
|
||||
entities: list[HaEntitySummary],
|
||||
domains: set[str] | None = None,
|
||||
learnable: bool | None = None,
|
||||
) -> list[DiscoveredEntity]:
|
||||
normalized_domains = {domain.strip().lower() for domain in domains or set() if domain.strip()}
|
||||
discovered = [classify_entity(entity) for entity in entities]
|
||||
return [
|
||||
entity
|
||||
for entity in discovered
|
||||
if (not normalized_domains or entity.domain in normalized_domains)
|
||||
and (learnable is None or entity.learnable is learnable)
|
||||
]
|
||||
|
||||
|
||||
def _result(
|
||||
entity: HaEntitySummary,
|
||||
role: EntityRole,
|
||||
*,
|
||||
learnable: bool,
|
||||
reason: str,
|
||||
) -> DiscoveredEntity:
|
||||
return DiscoveredEntity(
|
||||
entity_id=entity.entity_id,
|
||||
domain=entity.domain,
|
||||
device_class=entity.device_class,
|
||||
state_class=entity.state_class,
|
||||
unit_of_measurement=entity.unit_of_measurement,
|
||||
role=role,
|
||||
learnable=learnable,
|
||||
reason=reason,
|
||||
)
|
||||
91
app/ha/history.py
Normal file
91
app/ha/history.py
Normal file
@@ -0,0 +1,91 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from app.ha.exceptions import HaUnexpectedPayloadError
|
||||
|
||||
|
||||
class NumericHistoryPoint(BaseModel):
|
||||
timestamp: datetime
|
||||
value: float
|
||||
|
||||
|
||||
class EntityHistorySeries(BaseModel):
|
||||
entity_id: str
|
||||
points: list[NumericHistoryPoint]
|
||||
|
||||
|
||||
def normalize_history_payload(payload: object) -> list[EntityHistorySeries]:
|
||||
if not isinstance(payload, list):
|
||||
raise HaUnexpectedPayloadError("History-Payload muss eine Liste sein.")
|
||||
|
||||
normalized: list[EntityHistorySeries] = []
|
||||
for raw_series in payload:
|
||||
if not isinstance(raw_series, list):
|
||||
raise HaUnexpectedPayloadError("History-Serie muss eine Liste sein.")
|
||||
series = _normalize_series(raw_series)
|
||||
if series is not None:
|
||||
normalized.append(series)
|
||||
|
||||
return sorted(normalized, key=lambda item: item.entity_id)
|
||||
|
||||
|
||||
def _normalize_series(raw_series: list[object]) -> EntityHistorySeries | None:
|
||||
entity_id: str | None = None
|
||||
points: list[NumericHistoryPoint] = []
|
||||
|
||||
for raw_entry in raw_series:
|
||||
if not isinstance(raw_entry, dict):
|
||||
raise HaUnexpectedPayloadError("History-Eintrag muss ein Objekt sein.")
|
||||
|
||||
raw_entity_id = raw_entry.get("entity_id")
|
||||
if raw_entity_id is not None:
|
||||
if not isinstance(raw_entity_id, str) or "." not in raw_entity_id:
|
||||
raise HaUnexpectedPayloadError("History-Eintrag enthält ungültige entity_id.")
|
||||
if entity_id is not None and entity_id != raw_entity_id:
|
||||
raise HaUnexpectedPayloadError("History-Serie enthält mehrere Entities.")
|
||||
entity_id = raw_entity_id
|
||||
|
||||
raw_state = raw_entry.get("state")
|
||||
value = _finite_float(raw_state)
|
||||
if value is None:
|
||||
continue
|
||||
if entity_id is None:
|
||||
raise HaUnexpectedPayloadError("History-Serie enthält keine entity_id.")
|
||||
|
||||
raw_timestamp = raw_entry.get("last_changed") or raw_entry.get("last_updated")
|
||||
timestamp = _parse_timestamp(raw_timestamp)
|
||||
points.append(NumericHistoryPoint(timestamp=timestamp, value=value))
|
||||
|
||||
if entity_id is None or not points:
|
||||
return None
|
||||
|
||||
points.sort(key=lambda point: point.timestamp)
|
||||
return EntityHistorySeries(entity_id=entity_id, points=points)
|
||||
|
||||
|
||||
def _finite_float(value: object) -> float | None:
|
||||
if isinstance(value, bool) or value is None:
|
||||
return None
|
||||
if not isinstance(value, (str, int, float)):
|
||||
return None
|
||||
try:
|
||||
converted = float(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
return converted if math.isfinite(converted) else None
|
||||
|
||||
|
||||
def _parse_timestamp(value: object) -> datetime:
|
||||
if not isinstance(value, str):
|
||||
raise HaUnexpectedPayloadError("Numerischer History-Eintrag enthält keinen Zeitstempel.")
|
||||
try:
|
||||
parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||
except ValueError as exc:
|
||||
raise HaUnexpectedPayloadError("History-Eintrag enthält ungültigen Zeitstempel.") from exc
|
||||
if parsed.tzinfo is None:
|
||||
raise HaUnexpectedPayloadError("History-Zeitstempel muss eine Zeitzone enthalten.")
|
||||
return parsed
|
||||
@@ -1,9 +1,12 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from app.ha.client import HaClient
|
||||
from app.ha.discovery import DiscoveredEntity, discover_entities
|
||||
from app.ha.history import EntityHistorySeries, normalize_history_payload
|
||||
from app.ha.models import HaEntitySummary
|
||||
|
||||
|
||||
@@ -33,6 +36,22 @@ class HaReader:
|
||||
)
|
||||
return summaries
|
||||
|
||||
def discover(
|
||||
self,
|
||||
domains: set[str] | None = None,
|
||||
learnable: bool | None = None,
|
||||
) -> Sequence[DiscoveredEntity]:
|
||||
return discover_entities(list(self.read_entities()), domains=domains, learnable=learnable)
|
||||
|
||||
def read_history(
|
||||
self,
|
||||
entity_ids: list[str],
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
) -> Sequence[EntityHistorySeries]:
|
||||
payload = self._client.get_history(entity_ids, start_time, end_time)
|
||||
return normalize_history_payload(payload)
|
||||
|
||||
|
||||
def _optional_str(value: object) -> str | None:
|
||||
if value is None or value == "":
|
||||
|
||||
@@ -1,5 +1,14 @@
|
||||
|
||||
"""Machine-Learning-Grundbausteine für SillyHome Next."""
|
||||
__all__ = ["FeatureStore", "FeatureVector", "TrainedArtifact", "TrainingPipeline"]
|
||||
__all__ = [
|
||||
"FeatureStore",
|
||||
"FeatureVector",
|
||||
"RetrainingResult",
|
||||
"RetrainingService",
|
||||
"TrainedArtifact",
|
||||
"TrainingPipeline",
|
||||
"retrain_model",
|
||||
]
|
||||
from app.ml.feature_store import FeatureStore, FeatureVector
|
||||
from app.ml.retraining import RetrainingResult, RetrainingService, retrain_model
|
||||
from app.ml.training import TrainedArtifact, TrainingPipeline
|
||||
|
||||
@@ -5,6 +5,7 @@ import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
import re
|
||||
from threading import RLock
|
||||
from collections.abc import Iterable
|
||||
|
||||
from app.ml.training import TrainedArtifact
|
||||
@@ -19,22 +20,31 @@ class ModelRegistry:
|
||||
self._root = Path(root).resolve()
|
||||
self._root.mkdir(parents=True, exist_ok=True)
|
||||
self._artifacts: dict[str, TrainedArtifact] = {}
|
||||
self._lock = RLock()
|
||||
self._load_existing()
|
||||
|
||||
def register(self, artifact: TrainedArtifact) -> TrainedArtifact:
|
||||
registered, _ = self.register_with_status(artifact)
|
||||
return registered
|
||||
|
||||
def register_with_status(self, artifact: TrainedArtifact) -> tuple[TrainedArtifact, bool]:
|
||||
self._validate_artifact_id(artifact.artifact_id)
|
||||
self._persist(artifact)
|
||||
self._artifacts[artifact.artifact_id] = artifact
|
||||
return artifact
|
||||
with self._lock:
|
||||
replaced = artifact.artifact_id in self._artifacts
|
||||
self._persist(artifact)
|
||||
self._artifacts[artifact.artifact_id] = artifact
|
||||
return artifact, replaced
|
||||
|
||||
def load_artifact(self, artifact_id: str) -> TrainedArtifact:
|
||||
self._validate_artifact_id(artifact_id)
|
||||
if artifact_id not in self._artifacts:
|
||||
raise KeyError(f"Artifact '{artifact_id}' nicht registriert.")
|
||||
return self._artifacts[artifact_id]
|
||||
with self._lock:
|
||||
if artifact_id not in self._artifacts:
|
||||
raise KeyError(f"Artifact '{artifact_id}' nicht registriert.")
|
||||
return self._artifacts[artifact_id]
|
||||
|
||||
def list_models(self) -> Iterable[TrainedArtifact]:
|
||||
return list(self._artifacts.values())
|
||||
with self._lock:
|
||||
return [self._artifacts[key] for key in sorted(self._artifacts)]
|
||||
|
||||
def _load_existing(self) -> None:
|
||||
for source in sorted(self._root.glob("*.json")):
|
||||
|
||||
43
app/ml/retraining.py
Normal file
43
app/ml/retraining.py
Normal file
@@ -0,0 +1,43 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable
|
||||
from dataclasses import dataclass
|
||||
|
||||
from app.ml.feature_store import FeatureStore, FeatureVector
|
||||
from app.ml.registry.model_registry import ModelRegistry
|
||||
from app.ml.training import TrainedArtifact, TrainingPipeline
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RetrainingResult:
|
||||
artifact: TrainedArtifact
|
||||
replaced: bool
|
||||
|
||||
|
||||
class RetrainingService:
|
||||
"""Runs one retraining cycle without owning scheduling or background threads."""
|
||||
|
||||
def __init__(self, registry: ModelRegistry) -> None:
|
||||
self._registry = registry
|
||||
|
||||
def retrain(
|
||||
self,
|
||||
artifact_id: str,
|
||||
vectors: Iterable[FeatureVector],
|
||||
) -> RetrainingResult:
|
||||
store = FeatureStore()
|
||||
store.add_batch(vectors)
|
||||
pipeline = TrainingPipeline(store)
|
||||
artifact = pipeline.run(artifact_id)
|
||||
_, replaced = self._registry.register_with_status(artifact)
|
||||
return RetrainingResult(artifact=artifact, replaced=replaced)
|
||||
|
||||
|
||||
def retrain_model(
|
||||
registry: ModelRegistry,
|
||||
artifact_id: str,
|
||||
vectors: Iterable[FeatureVector],
|
||||
) -> RetrainingResult:
|
||||
"""Scheduler-compatible entry point for exactly one retraining run."""
|
||||
|
||||
return RetrainingService(registry).retrain(artifact_id, vectors)
|
||||
@@ -10,6 +10,7 @@ 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.retraining import retrain_model
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -45,6 +46,23 @@ class ModelsResponse(BaseModel):
|
||||
models: list[str]
|
||||
|
||||
|
||||
class TrainingSample(BaseModel):
|
||||
sensor_id: str = Field(min_length=1)
|
||||
values: dict[str, float]
|
||||
label: str | None = None
|
||||
|
||||
|
||||
class RetrainRequest(BaseModel):
|
||||
model_id: str = Field(..., alias="modelId", min_length=1, max_length=128)
|
||||
samples: list[TrainingSample] = Field(min_length=1)
|
||||
|
||||
|
||||
class RetrainResponse(BaseModel):
|
||||
model_id: str
|
||||
supported_sensors: list[str]
|
||||
replaced: bool
|
||||
|
||||
|
||||
@router.get("/health", response_model=HealthResponse, status_code=200)
|
||||
def health() -> HealthResponse:
|
||||
return HealthResponse(status="ok")
|
||||
@@ -57,6 +75,31 @@ def list_models(request: Request) -> ModelsResponse:
|
||||
return ModelsResponse(models=models)
|
||||
|
||||
|
||||
@router.post("/retrain", response_model=RetrainResponse, status_code=200)
|
||||
def retrain(payload: RetrainRequest, request: Request) -> RetrainResponse:
|
||||
registry = _require_registry(request)
|
||||
vectors = [
|
||||
FeatureVector(
|
||||
sensor_id=sample.sensor_id,
|
||||
values=sample.values,
|
||||
label=sample.label,
|
||||
)
|
||||
for sample in payload.samples
|
||||
]
|
||||
try:
|
||||
result = retrain_model(registry, payload.model_id, vectors)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_CONTENT,
|
||||
detail=str(exc),
|
||||
) from exc
|
||||
return RetrainResponse(
|
||||
model_id=result.artifact.artifact_id,
|
||||
supported_sensors=list(result.artifact.supported_sensors),
|
||||
replaced=result.replaced,
|
||||
)
|
||||
|
||||
|
||||
@router.post("/predict", response_model=PredictResponse, status_code=200)
|
||||
def predict(payload: PredictRequest, request: Request) -> PredictResponse:
|
||||
registry = _require_registry(request)
|
||||
|
||||
42
docs/ha_data.md
Normal file
42
docs/ha_data.md
Normal file
@@ -0,0 +1,42 @@
|
||||
# Home-Assistant-Datenpipeline
|
||||
|
||||
SillyHome Next trennt aktuelle Entity-Metadaten, Discovery und historische
|
||||
Messwerte. Dadurch gelangen nur klassifizierte, geeignete Daten in spätere
|
||||
Trainings- und Erklärungsprozesse.
|
||||
|
||||
## Entity Discovery
|
||||
|
||||
`GET /v1/discovery` klassifiziert Home-Assistant-Entities in:
|
||||
|
||||
- `measurement`: numerische Messsensoren, für Training geeignet
|
||||
- `binary_context`: binäre Kontextsensoren wie Bewegung oder Anwesenheit
|
||||
- `context`: Personen-, Wetter- und Standortkontext
|
||||
- `actuator`: mögliche Automationsziele, nicht als Trainingssensor verwendet
|
||||
- `unsupported`: noch nicht klassifizierte Entity-Typen
|
||||
|
||||
Optionale Query-Parameter:
|
||||
|
||||
- `domain=sensor` kann mehrfach angegeben werden
|
||||
- `learnable=true|false` filtert nach Trainingsrelevanz
|
||||
|
||||
## Historische Daten
|
||||
|
||||
Historische Zustände werden über Home Assistants
|
||||
`/api/history/period/<start>`-Schnittstelle geladen. Abfragen verlangen:
|
||||
|
||||
- mindestens eine Entity-ID, maximal 100
|
||||
- zeitzonenbehaftete Start- und Endzeit
|
||||
- ein Enddatum nach dem Startdatum
|
||||
- maximal 31 Tage pro Abfrage
|
||||
|
||||
Die Normalisierung übernimmt nur endliche numerische Zustände. `unknown`,
|
||||
`unavailable`, nichtnumerische Werte, `NaN` und unendliche Werte werden nicht
|
||||
als Trainingsdaten verwendet. Ergebnisse werden je Entity chronologisch
|
||||
sortiert.
|
||||
|
||||
## Datenschutz und Betrieb
|
||||
|
||||
Die Daten bleiben lokal. Home-Assistant-Tokens gehören ausschließlich in die
|
||||
Umgebungskonfiguration und dürfen nicht protokolliert oder versioniert werden.
|
||||
Die API sollte nur lokal oder hinter einem authentifizierenden Reverse Proxy
|
||||
erreichbar sein.
|
||||
@@ -12,6 +12,7 @@ Modell-Artefakt- und Vorhersage-Schnittstelle.
|
||||
- Standard: `http://127.0.0.1:8000/ml`
|
||||
- Health: `/health`
|
||||
- Modelle: `/models`
|
||||
- Retraining: `/retrain`
|
||||
- Einzelvorhersage: `/predict`
|
||||
- Batchvorhersage: `/batch`
|
||||
|
||||
@@ -65,6 +66,35 @@ Einzelne Vorhersage für einen Sensor.
|
||||
}
|
||||
```
|
||||
|
||||
### `POST /ml/retrain`
|
||||
|
||||
Trainiert die Artefakt-Metadaten aus neuen Sensordaten. Existiert `modelId`
|
||||
bereits, wird das Artefakt atomisch ersetzt und beim nächsten Prozessstart aus
|
||||
dem Modellverzeichnis geladen.
|
||||
|
||||
**Request**
|
||||
```json
|
||||
{
|
||||
"modelId": "home-model",
|
||||
"samples": [
|
||||
{
|
||||
"sensor_id": "sensor.kitchen",
|
||||
"values": {"temperature": 21.0},
|
||||
"label": "occupied"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
**Antwort**
|
||||
```json
|
||||
{
|
||||
"model_id": "home-model",
|
||||
"supported_sensors": ["sensor.kitchen"],
|
||||
"replaced": false
|
||||
}
|
||||
```
|
||||
|
||||
### `POST /ml/batch`
|
||||
|
||||
Batch-Vorhersage für mehrere Sensorwerte.
|
||||
@@ -114,11 +144,15 @@ Batch-Vorhersage für mehrere Sensorwerte.
|
||||
## Betrieb
|
||||
|
||||
Die produktive App lädt Artefakte aus `SILLYHOME_MODEL_STORE`. Neue Artefakte
|
||||
werden derzeit intern über `ModelRegistry.register(...)` registriert. Die
|
||||
Registry speichert validiertes JSON atomisch und lädt es beim Neustart.
|
||||
werden über `/ml/retrain`, `RetrainingService` oder direkt über
|
||||
`ModelRegistry.register(...)` registriert. Die Registry speichert validiertes
|
||||
JSON atomisch und lädt es beim Neustart. Die API sollte nur in einem
|
||||
vertrauenswürdigen Netz oder hinter einem authentifizierenden Reverse Proxy
|
||||
erreichbar sein.
|
||||
|
||||
## Verweise
|
||||
|
||||
- `app/ml/predictor.py`
|
||||
- `app/ml/retraining.py`
|
||||
- `app/ml/registry/model_registry.py`
|
||||
- `backend/routes/ml.py`
|
||||
|
||||
@@ -38,6 +38,21 @@ Der Report enthält:
|
||||
|
||||
Das trainierte Artefakt kann anschließend über `ModelRegistry.register(artifact)` bereitgestellt werden. Die ML-Serving-API stellt es unter `/ml/predict` und `/ml/batch` zur Verfügung.
|
||||
|
||||
## 5. Retraining ausführen
|
||||
|
||||
`RetrainingService.retrain(...)` führt genau einen Trainingslauf aus und ersetzt
|
||||
ein vorhandenes Artefakt mit derselben ID atomisch in der Registry:
|
||||
|
||||
```python
|
||||
service = RetrainingService(registry)
|
||||
result = service.retrain("home-model", vectors)
|
||||
```
|
||||
|
||||
Scheduler, Cronjobs oder Home-Assistant-Automationen können alternativ die
|
||||
zustandslose Funktion `retrain_model(registry, artifact_id, vectors)` aufrufen.
|
||||
Der Service startet bewusst keinen eigenen Hintergrundprozess. Über
|
||||
`POST /ml/retrain` kann derselbe Ablauf per API angestoßen werden.
|
||||
|
||||
## Hinweise
|
||||
- Für reproduzierbare Sensor-Reihenfolgen wird in `TrainingPipeline.run(...)` eine sortierte Sensor-Liste verwendet.
|
||||
- Fehlende Trainingsdaten lösen `ValueError` aus; nicht registrierte Artefakte lösen `KeyError` aus.
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
from collections.abc import Sequence
|
||||
from datetime import datetime
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from app.ha.exceptions import HaTimeoutError
|
||||
from app.ha.discovery import DiscoveredEntity, EntityRole
|
||||
from app.ha.history import EntityHistorySeries, NumericHistoryPoint
|
||||
from app.ha.models import HaEntitySummary
|
||||
from app.ha.reader import HaReader
|
||||
from app.main import app
|
||||
@@ -15,6 +18,38 @@ class FakeHaReader(HaReader):
|
||||
def read_entities(self) -> Sequence[HaEntitySummary]:
|
||||
return [HaEntitySummary(entity_id="sensor.temperature", domain="sensor")]
|
||||
|
||||
def discover(
|
||||
self,
|
||||
domains: set[str] | None = None,
|
||||
learnable: bool | None = None,
|
||||
) -> Sequence[DiscoveredEntity]:
|
||||
result = DiscoveredEntity(
|
||||
entity_id="sensor.temperature",
|
||||
domain="sensor",
|
||||
device_class="temperature",
|
||||
role=EntityRole.MEASUREMENT,
|
||||
learnable=True,
|
||||
reason="Numerischer Messsensor für Zeitreihen und Training.",
|
||||
)
|
||||
if domains and result.domain not in domains:
|
||||
return []
|
||||
if learnable is not None and result.learnable is not learnable:
|
||||
return []
|
||||
return [result]
|
||||
|
||||
def read_history(
|
||||
self,
|
||||
entity_ids: list[str],
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
) -> Sequence[EntityHistorySeries]:
|
||||
return [
|
||||
EntityHistorySeries(
|
||||
entity_id=entity_ids[0],
|
||||
points=[NumericHistoryPoint(timestamp=start_time, value=21.5)],
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
class TimeoutHaReader(HaReader):
|
||||
def __init__(self) -> None:
|
||||
@@ -59,3 +94,44 @@ def test_entities_maps_ha_errors_without_leaking_details() -> None:
|
||||
response = client.get("/v1/entities")
|
||||
assert response.status_code == 504
|
||||
assert response.json() == {"detail": "Home Assistant request timed out."}
|
||||
|
||||
|
||||
def test_discovery_filters_entities() -> None:
|
||||
with TestClient(app) as client:
|
||||
app.state.ha_reader = FakeHaReader()
|
||||
response = client.get("/v1/discovery?domain=sensor&learnable=true")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json() == [
|
||||
{
|
||||
"entity_id": "sensor.temperature",
|
||||
"domain": "sensor",
|
||||
"device_class": "temperature",
|
||||
"state_class": None,
|
||||
"unit_of_measurement": None,
|
||||
"role": "measurement",
|
||||
"learnable": True,
|
||||
"reason": "Numerischer Messsensor für Zeitreihen und Training.",
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_history_returns_normalized_series() -> None:
|
||||
with TestClient(app) as client:
|
||||
app.state.ha_reader = FakeHaReader()
|
||||
response = client.get(
|
||||
"/v1/history",
|
||||
params=[
|
||||
("entity_id", "sensor.temperature"),
|
||||
("start_time", "2026-06-01T00:00:00Z"),
|
||||
("end_time", "2026-06-02T00:00:00Z"),
|
||||
],
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json() == [
|
||||
{
|
||||
"entity_id": "sensor.temperature",
|
||||
"points": [{"timestamp": "2026-06-01T00:00:00Z", "value": 21.5}],
|
||||
}
|
||||
]
|
||||
|
||||
@@ -50,3 +50,60 @@ def test_unsupported_sensor_returns_422(tmp_path: Path) -> None:
|
||||
)
|
||||
|
||||
assert response.status_code == 422
|
||||
|
||||
|
||||
def test_retrain_creates_and_replaces_persisted_model(tmp_path: Path) -> None:
|
||||
from app.ml.registry.model_registry import ModelRegistry
|
||||
|
||||
registry = ModelRegistry(tmp_path)
|
||||
with TestClient(app) as client:
|
||||
app.state.registry = registry
|
||||
created = client.post(
|
||||
"/ml/retrain",
|
||||
json={
|
||||
"modelId": "home-model",
|
||||
"samples": [
|
||||
{
|
||||
"sensor_id": "sensor.kitchen",
|
||||
"values": {"temperature": 21.0},
|
||||
}
|
||||
],
|
||||
},
|
||||
)
|
||||
replaced = client.post(
|
||||
"/ml/retrain",
|
||||
json={
|
||||
"modelId": "home-model",
|
||||
"samples": [
|
||||
{
|
||||
"sensor_id": "sensor.bedroom",
|
||||
"values": {"temperature": 18.0},
|
||||
}
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
assert created.status_code == 200
|
||||
assert created.json() == {
|
||||
"model_id": "home-model",
|
||||
"supported_sensors": ["sensor.kitchen"],
|
||||
"replaced": False,
|
||||
}
|
||||
assert replaced.status_code == 200
|
||||
assert replaced.json() == {
|
||||
"model_id": "home-model",
|
||||
"supported_sensors": ["sensor.bedroom"],
|
||||
"replaced": True,
|
||||
}
|
||||
restarted = ModelRegistry(tmp_path)
|
||||
assert restarted.load_artifact("home-model").supported_sensors == ("sensor.bedroom",)
|
||||
|
||||
|
||||
def test_retrain_rejects_empty_samples() -> None:
|
||||
with TestClient(app) as client:
|
||||
response = client.post(
|
||||
"/ml/retrain",
|
||||
json={"modelId": "home-model", "samples": []},
|
||||
)
|
||||
|
||||
assert response.status_code == 422
|
||||
|
||||
76
tests/ha/test_discovery.py
Normal file
76
tests/ha/test_discovery.py
Normal file
@@ -0,0 +1,76 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from app.ha.discovery import EntityRole, classify_entity, discover_entities
|
||||
from app.ha.models import HaEntitySummary
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("entity", "role", "learnable"),
|
||||
[
|
||||
(
|
||||
HaEntitySummary(
|
||||
entity_id="sensor.temperature",
|
||||
domain="sensor",
|
||||
device_class="temperature",
|
||||
state_class="measurement",
|
||||
unit_of_measurement="°C",
|
||||
),
|
||||
EntityRole.MEASUREMENT,
|
||||
True,
|
||||
),
|
||||
(
|
||||
HaEntitySummary(
|
||||
entity_id="binary_sensor.motion",
|
||||
domain="binary_sensor",
|
||||
device_class="motion",
|
||||
),
|
||||
EntityRole.BINARY_CONTEXT,
|
||||
True,
|
||||
),
|
||||
(
|
||||
HaEntitySummary(entity_id="person.simon", domain="person"),
|
||||
EntityRole.CONTEXT,
|
||||
True,
|
||||
),
|
||||
(
|
||||
HaEntitySummary(entity_id="light.living_room", domain="light"),
|
||||
EntityRole.ACTUATOR,
|
||||
False,
|
||||
),
|
||||
(
|
||||
HaEntitySummary(entity_id="camera.driveway", domain="camera"),
|
||||
EntityRole.UNSUPPORTED,
|
||||
False,
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_classify_entity(
|
||||
entity: HaEntitySummary,
|
||||
role: EntityRole,
|
||||
learnable: bool,
|
||||
) -> None:
|
||||
result = classify_entity(entity)
|
||||
assert result.role is role
|
||||
assert result.learnable is learnable
|
||||
|
||||
|
||||
def test_discovery_filters_domain_and_learnable() -> None:
|
||||
entities = [
|
||||
HaEntitySummary(
|
||||
entity_id="sensor.temperature",
|
||||
domain="sensor",
|
||||
device_class="temperature",
|
||||
),
|
||||
HaEntitySummary(entity_id="sensor.status", domain="sensor"),
|
||||
HaEntitySummary(
|
||||
entity_id="binary_sensor.motion",
|
||||
domain="binary_sensor",
|
||||
device_class="motion",
|
||||
),
|
||||
]
|
||||
|
||||
result = discover_entities(entities, domains={" SENSOR "}, learnable=True)
|
||||
|
||||
assert [item.entity_id for item in result] == ["sensor.temperature"]
|
||||
@@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
@@ -68,4 +69,60 @@ def test_list_entities_rejects_invalid_json() -> None:
|
||||
def test_list_entities_rejects_non_list_payload() -> None:
|
||||
client = _client_with_response(_response(payload={"entity_id": "sensor.temperature"}))
|
||||
with pytest.raises(HaUnexpectedPayloadError):
|
||||
client.list_entities()
|
||||
client.list_entities()
|
||||
|
||||
|
||||
def test_get_history_calls_home_assistant_history_api() -> None:
|
||||
response = _response(payload=[[{"entity_id": "sensor.temperature", "state": "21.0"}]])
|
||||
client = _client_with_response(response)
|
||||
start = datetime(2026, 6, 1, tzinfo=timezone.utc)
|
||||
end = datetime(2026, 6, 2, tzinfo=timezone.utc)
|
||||
|
||||
payload = client.get_history(["sensor.temperature"], start, end)
|
||||
|
||||
assert payload == [[{"entity_id": "sensor.temperature", "state": "21.0"}]]
|
||||
client._session.get.assert_called_once() # type: ignore[attr-defined]
|
||||
call = client._session.get.call_args # type: ignore[attr-defined]
|
||||
assert "/api/history/period/2026-06-01T00:00:00+00:00" in call.args[0]
|
||||
assert call.kwargs["params"]["filter_entity_id"] == "sensor.temperature"
|
||||
assert call.kwargs["params"]["end_time"] == "2026-06-02T00:00:00+00:00"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("entity_ids", "start", "end"),
|
||||
[
|
||||
(
|
||||
[],
|
||||
datetime(2026, 6, 1, tzinfo=timezone.utc),
|
||||
datetime(2026, 6, 2, tzinfo=timezone.utc),
|
||||
),
|
||||
(
|
||||
["sensor.temperature"],
|
||||
datetime(2026, 6, 1),
|
||||
datetime(2026, 6, 2, tzinfo=timezone.utc),
|
||||
),
|
||||
(
|
||||
["sensor.temperature"],
|
||||
datetime(2026, 6, 2, tzinfo=timezone.utc),
|
||||
datetime(2026, 6, 1, tzinfo=timezone.utc),
|
||||
),
|
||||
(
|
||||
["invalid entity"],
|
||||
datetime(2026, 6, 1, tzinfo=timezone.utc),
|
||||
datetime(2026, 6, 2, tzinfo=timezone.utc),
|
||||
),
|
||||
(
|
||||
["sensor.temperature"],
|
||||
datetime(2026, 5, 1, tzinfo=timezone.utc),
|
||||
datetime(2026, 6, 2, tzinfo=timezone.utc),
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_get_history_validates_request(
|
||||
entity_ids: list[str],
|
||||
start: datetime,
|
||||
end: datetime,
|
||||
) -> None:
|
||||
client = HaClient(HaClientSettings(url="http://ha.local", token="test-token"))
|
||||
with pytest.raises(ValueError):
|
||||
client.get_history(entity_ids, start, end)
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from app.ha.client import HaClient, HaClientSettings
|
||||
from app.ha.reader import HaReader
|
||||
|
||||
@@ -26,6 +28,22 @@ class FakeHaClient(HaClient):
|
||||
},
|
||||
]
|
||||
|
||||
def get_history(
|
||||
self,
|
||||
entity_ids: list[str],
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
) -> list[object]:
|
||||
return [
|
||||
[
|
||||
{
|
||||
"entity_id": entity_ids[0],
|
||||
"state": "21.5",
|
||||
"last_changed": start_time.isoformat(),
|
||||
}
|
||||
]
|
||||
]
|
||||
|
||||
|
||||
def test_ha_reader_returns_summaries() -> None:
|
||||
reader = HaReader(FakeHaClient())
|
||||
@@ -35,3 +53,25 @@ def test_ha_reader_returns_summaries() -> None:
|
||||
assert domains == {"sensor", "light"}
|
||||
sensor = next(item for item in summaries if item.entity_id == "sensor.temperature")
|
||||
assert sensor.unit_of_measurement == "°C"
|
||||
|
||||
|
||||
def test_ha_reader_discovers_learnable_sensors() -> None:
|
||||
reader = HaReader(FakeHaClient())
|
||||
|
||||
discovered = reader.discover(learnable=True)
|
||||
|
||||
assert [entity.entity_id for entity in discovered] == ["sensor.temperature"]
|
||||
|
||||
|
||||
def test_ha_reader_normalizes_history() -> None:
|
||||
reader = HaReader(FakeHaClient())
|
||||
start = datetime(2026, 6, 1, tzinfo=timezone.utc)
|
||||
|
||||
history = reader.read_history(
|
||||
["sensor.temperature"],
|
||||
start,
|
||||
datetime(2026, 6, 2, tzinfo=timezone.utc),
|
||||
)
|
||||
|
||||
assert history[0].entity_id == "sensor.temperature"
|
||||
assert history[0].points[0].value == 21.5
|
||||
|
||||
92
tests/ha/test_history.py
Normal file
92
tests/ha/test_history.py
Normal file
@@ -0,0 +1,92 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
|
||||
from app.ha.exceptions import HaUnexpectedPayloadError
|
||||
from app.ha.history import normalize_history_payload
|
||||
|
||||
|
||||
def test_normalize_history_payload_groups_and_sorts_numeric_states() -> None:
|
||||
payload = [
|
||||
[
|
||||
{
|
||||
"entity_id": "sensor.temperature",
|
||||
"state": "22.5",
|
||||
"last_changed": "2026-06-01T12:15:00+00:00",
|
||||
},
|
||||
{
|
||||
"state": "21.0",
|
||||
"last_changed": "2026-06-01T12:00:00Z",
|
||||
},
|
||||
],
|
||||
[
|
||||
{
|
||||
"entity_id": "sensor.humidity",
|
||||
"state": 45,
|
||||
"last_updated": "2026-06-01T12:00:00+00:00",
|
||||
}
|
||||
],
|
||||
]
|
||||
|
||||
result = normalize_history_payload(payload)
|
||||
|
||||
assert [series.entity_id for series in result] == [
|
||||
"sensor.humidity",
|
||||
"sensor.temperature",
|
||||
]
|
||||
temperature = result[1]
|
||||
assert [point.value for point in temperature.points] == [21.0, 22.5]
|
||||
assert temperature.points[0].timestamp == datetime(
|
||||
2026, 6, 1, 12, 0, tzinfo=timezone.utc
|
||||
)
|
||||
|
||||
|
||||
def test_normalize_history_payload_skips_non_numeric_and_non_finite_states() -> None:
|
||||
payload = [
|
||||
[
|
||||
{
|
||||
"entity_id": "sensor.temperature",
|
||||
"state": state,
|
||||
"last_changed": "2026-06-01T12:00:00+00:00",
|
||||
}
|
||||
for state in ("unknown", "unavailable", "nan", "inf", "-inf", True, None)
|
||||
]
|
||||
]
|
||||
|
||||
assert normalize_history_payload(payload) == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"payload",
|
||||
[
|
||||
{},
|
||||
[{}],
|
||||
[["invalid"]],
|
||||
[[{"entity_id": "invalid", "state": "21", "last_changed": "2026-06-01"}]],
|
||||
[[{"entity_id": "sensor.a", "state": "21", "last_changed": "invalid"}]],
|
||||
[[{"state": "21", "last_changed": "2026-06-01T12:00:00+00:00"}]],
|
||||
[
|
||||
[
|
||||
{
|
||||
"entity_id": "sensor.a",
|
||||
"state": "21",
|
||||
"last_changed": "2026-06-01T12:00:00+00:00",
|
||||
},
|
||||
{
|
||||
"entity_id": "sensor.b",
|
||||
"state": "22",
|
||||
"last_changed": "2026-06-01T12:01:00+00:00",
|
||||
},
|
||||
]
|
||||
],
|
||||
],
|
||||
)
|
||||
def test_normalize_history_payload_rejects_malformed_structure(payload: object) -> None:
|
||||
with pytest.raises(HaUnexpectedPayloadError):
|
||||
normalize_history_payload(payload)
|
||||
|
||||
|
||||
def test_normalize_history_payload_accepts_empty_series() -> None:
|
||||
assert normalize_history_payload([[]]) == []
|
||||
@@ -19,6 +19,17 @@ def test_registry_loads_persisted_artifacts_after_restart(tmp_path: Path) -> Non
|
||||
assert restarted.load_artifact("model-v1") == artifact
|
||||
|
||||
|
||||
def test_registry_replaces_persisted_artifact_after_restart(tmp_path: Path) -> None:
|
||||
registry = ModelRegistry(tmp_path)
|
||||
registry.register(TrainedArtifact("model-v1", ("sensor.kitchen",)))
|
||||
replacement = TrainedArtifact("model-v1", ("sensor.bedroom",))
|
||||
|
||||
registry.register(replacement)
|
||||
|
||||
assert registry.load_artifact("model-v1") == replacement
|
||||
assert ModelRegistry(tmp_path).load_artifact("model-v1") == replacement
|
||||
|
||||
|
||||
@pytest.mark.parametrize("artifact_id", ["../escape", "nested/model", "..", ""])
|
||||
def test_registry_rejects_unsafe_artifact_ids(tmp_path: Path, artifact_id: str) -> None:
|
||||
registry = ModelRegistry(tmp_path)
|
||||
|
||||
39
tests/ml/test_retraining.py
Normal file
39
tests/ml/test_retraining.py
Normal file
@@ -0,0 +1,39 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from app.ml.feature_store import FeatureVector
|
||||
from app.ml.registry.model_registry import ModelRegistry
|
||||
from app.ml.retraining import RetrainingService, retrain_model
|
||||
|
||||
|
||||
def _vector(sensor_id: str) -> FeatureVector:
|
||||
return FeatureVector(sensor_id=sensor_id, values={"temperature": 21.0})
|
||||
|
||||
|
||||
def test_retraining_registers_new_artifact(tmp_path: Path) -> None:
|
||||
registry = ModelRegistry(tmp_path)
|
||||
|
||||
result = retrain_model(registry, "home-model", [_vector("sensor.kitchen")])
|
||||
|
||||
assert result.replaced is False
|
||||
assert registry.load_artifact("home-model") == result.artifact
|
||||
|
||||
|
||||
def test_retraining_replaces_existing_artifact(tmp_path: Path) -> None:
|
||||
registry = ModelRegistry(tmp_path)
|
||||
service = RetrainingService(registry)
|
||||
service.retrain("home-model", [_vector("sensor.kitchen")])
|
||||
|
||||
result = service.retrain("home-model", [_vector("sensor.bedroom")])
|
||||
|
||||
assert result.replaced is True
|
||||
assert result.artifact.supported_sensors == ("sensor.bedroom",)
|
||||
assert ModelRegistry(tmp_path).load_artifact("home-model") == result.artifact
|
||||
|
||||
|
||||
def test_retraining_rejects_empty_training_data(tmp_path: Path) -> None:
|
||||
with pytest.raises(ValueError, match="keine Trainingsdaten"):
|
||||
retrain_model(ModelRegistry(tmp_path), "home-model", [])
|
||||
Reference in New Issue
Block a user