182 lines
7.5 KiB
Python
182 lines
7.5 KiB
Python
from __future__ import annotations
|
|
|
|
import json
|
|
import logging
|
|
import math
|
|
import os
|
|
from pathlib import Path
|
|
import re
|
|
from threading import RLock
|
|
from collections.abc import Iterable
|
|
|
|
from app.ml.training import FeatureModel, TrainedArtifact
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
_ARTIFACT_ID_PATTERN = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]{0,127}$")
|
|
|
|
|
|
class ModelRegistry:
|
|
def __init__(self, root: str | Path) -> None:
|
|
self._root = Path(root).resolve()
|
|
self._root.mkdir(parents=True, exist_ok=True)
|
|
self._archive_root = self._root / "archive"
|
|
self._archive_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)
|
|
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)
|
|
with self._lock:
|
|
if artifact_id not in self._artifacts:
|
|
raise KeyError(f"Artifact '{artifact_id}' nicht registriert.")
|
|
return self._artifacts[artifact_id]
|
|
|
|
def get_optional(self, artifact_id: str) -> TrainedArtifact | None:
|
|
self._validate_artifact_id(artifact_id)
|
|
with self._lock:
|
|
return self._artifacts.get(artifact_id)
|
|
|
|
def list_models(self) -> Iterable[TrainedArtifact]:
|
|
with self._lock:
|
|
return [self._artifacts[key] for key in sorted(self._artifacts)]
|
|
|
|
def archive(self, artifact_id: str) -> bool:
|
|
self._validate_artifact_id(artifact_id)
|
|
with self._lock:
|
|
artifact = self._artifacts.pop(artifact_id, None)
|
|
source = self._root / f"{artifact_id}.json"
|
|
if not source.exists():
|
|
return artifact is not None
|
|
target = self._archive_root / f"{artifact_id}.json"
|
|
os.replace(source, target)
|
|
logger.info("Modell archiviert: %s", target)
|
|
return True
|
|
|
|
def _load_existing(self) -> None:
|
|
for source in sorted(self._root.glob("*.json")):
|
|
try:
|
|
raw = json.loads(source.read_text(encoding="utf-8"))
|
|
artifact_id = raw["artifact_id"]
|
|
supported_sensors = raw["supported_sensors"]
|
|
model_type = raw.get("model_type", "metadata")
|
|
raw_feature_models = raw.get("feature_models", {})
|
|
if not isinstance(artifact_id, str) or not isinstance(supported_sensors, list):
|
|
raise ValueError("invalid artifact structure")
|
|
if not isinstance(model_type, str):
|
|
raise ValueError("model_type must be a string")
|
|
self._validate_artifact_id(artifact_id)
|
|
if source.name != f"{artifact_id}.json":
|
|
raise ValueError("artifact id does not match filename")
|
|
if not all(isinstance(sensor, str) for sensor in supported_sensors):
|
|
raise ValueError("supported_sensors must contain strings")
|
|
feature_models = _deserialize_feature_models(raw_feature_models)
|
|
except (KeyError, TypeError, ValueError, json.JSONDecodeError) as exc:
|
|
raise ValueError(f"Ungültiges Modell-Artefakt: {source.name}") from exc
|
|
|
|
self._artifacts[artifact_id] = TrainedArtifact(
|
|
artifact_id=artifact_id,
|
|
supported_sensors=tuple(supported_sensors),
|
|
feature_models=feature_models,
|
|
model_type=model_type,
|
|
)
|
|
|
|
def _persist(self, artifact: TrainedArtifact) -> None:
|
|
target = self._root / f"{artifact.artifact_id}.json"
|
|
temporary = target.with_suffix(".json.tmp")
|
|
payload = {
|
|
"artifact_id": artifact.artifact_id,
|
|
"supported_sensors": list(artifact.supported_sensors),
|
|
"model_type": artifact.model_type,
|
|
"feature_models": {
|
|
sensor_id: {
|
|
feature_name: {
|
|
"sample_count": model.sample_count,
|
|
"mean": model.mean,
|
|
"standard_deviation": model.standard_deviation,
|
|
"minimum": model.minimum,
|
|
"maximum": model.maximum,
|
|
"slope": model.slope,
|
|
"intercept": model.intercept,
|
|
}
|
|
for feature_name, model in sorted(models.items())
|
|
}
|
|
for sensor_id, models in sorted(artifact.feature_models.items())
|
|
},
|
|
}
|
|
temporary.write_text(
|
|
json.dumps(payload, ensure_ascii=True, sort_keys=True) + "\n",
|
|
encoding="utf-8",
|
|
)
|
|
os.replace(temporary, target)
|
|
logger.info("Modell gespeichert: %s", target)
|
|
|
|
@staticmethod
|
|
def _validate_artifact_id(artifact_id: str) -> None:
|
|
if not _ARTIFACT_ID_PATTERN.fullmatch(artifact_id) or ".." in artifact_id:
|
|
raise ValueError(
|
|
"artifact_id darf nur Buchstaben, Ziffern, Punkt, Unterstrich "
|
|
"und Bindestrich enthalten."
|
|
)
|
|
|
|
|
|
def _deserialize_feature_models(raw: object) -> dict[str, dict[str, FeatureModel]]:
|
|
if not isinstance(raw, dict):
|
|
raise ValueError("feature_models must be an object")
|
|
|
|
result: dict[str, dict[str, FeatureModel]] = {}
|
|
for sensor_id, raw_features in raw.items():
|
|
if not isinstance(sensor_id, str) or not isinstance(raw_features, dict):
|
|
raise ValueError("invalid sensor feature models")
|
|
features: dict[str, FeatureModel] = {}
|
|
for feature_name, raw_model in raw_features.items():
|
|
if not isinstance(feature_name, str) or not isinstance(raw_model, dict):
|
|
raise ValueError("invalid feature model")
|
|
sample_count = raw_model.get("sample_count")
|
|
if not isinstance(sample_count, int) or isinstance(sample_count, bool) or sample_count < 1:
|
|
raise ValueError("sample_count must be a positive integer")
|
|
values = {
|
|
key: _finite_number(raw_model.get(key))
|
|
for key in (
|
|
"mean",
|
|
"standard_deviation",
|
|
"minimum",
|
|
"maximum",
|
|
"slope",
|
|
"intercept",
|
|
)
|
|
}
|
|
features[feature_name] = FeatureModel(
|
|
sample_count=sample_count,
|
|
mean=values["mean"],
|
|
standard_deviation=values["standard_deviation"],
|
|
minimum=values["minimum"],
|
|
maximum=values["maximum"],
|
|
slope=values["slope"],
|
|
intercept=values["intercept"],
|
|
)
|
|
result[sensor_id] = features
|
|
return result
|
|
|
|
|
|
def _finite_number(value: object) -> float:
|
|
if not isinstance(value, (int, float)) or isinstance(value, bool):
|
|
raise ValueError("feature model values must be finite numbers")
|
|
converted = float(value)
|
|
if not math.isfinite(converted):
|
|
raise ValueError("feature model values must be finite numbers")
|
|
return converted
|