ML-007: add retraining pipeline and API
Some checks failed
quality / test (3.11) (push) Has been cancelled
quality / test (3.13) (push) Has been cancelled

Closes #13
This commit is contained in:
2026-06-13 19:10:17 +02:00
parent ecd32d4813
commit 840c404c1c
12 changed files with 274 additions and 10 deletions

View File

@@ -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")):