hackathon-ENGIN/engin/artifact.py
Jakub Famulski 2 8e6dcb750d
All checks were successful
ENGIN CI / Build, test and smoke (push) Successful in 2m58s
prod: load versioned model artifact at runtime
2026-08-25 13:06:14 +02:00

258 lines
9.1 KiB
Python

"""Versioned, checksummed model-artifact persistence for runtime loading."""
from __future__ import annotations
import hashlib
import json
import os
import pickle
import platform
import tempfile
from dataclasses import asdict, dataclass
from datetime import UTC, datetime
from importlib.metadata import version
from pathlib import Path
from .inference import (
INFERENCE_COLUMNS,
OOD_OK_TO_UNKNOWN_RATIO_THRESHOLD,
OOD_OK_TO_UNKNOWN_THRESHOLD_MV,
DiagnosticModels,
)
ARTIFACT_FORMAT_VERSION = 1
MODEL_FILENAME = "model.pkl"
MANIFEST_FILENAME = "manifest.json"
RUNTIME_PACKAGES = (
"joblib",
"numpy",
"pandas",
"scikit-learn",
"scipy",
"threadpoolctl",
)
class ModelArtifactError(RuntimeError):
"""Raised when a model artifact is missing, corrupt, or incompatible."""
@dataclass(frozen=True)
class ArtifactMetadata:
artifact_format_version: int
model_version: str
created_at_utc: str
source_revision: str
training_data_sha256: str
model_sha256: str
input_columns: tuple[str, ...]
python_version: str
random_state: int
n_jobs: int
label_model: str
severity_model: str
ood_absolute_threshold_mv: float
ood_ratio_threshold: float
runtime_versions: dict[str, str]
@classmethod
def from_dict(cls, payload: dict[str, object]) -> "ArtifactMetadata":
try:
return cls(
artifact_format_version=int(payload["artifact_format_version"]),
model_version=str(payload["model_version"]),
created_at_utc=str(payload["created_at_utc"]),
source_revision=str(payload["source_revision"]),
training_data_sha256=str(payload["training_data_sha256"]),
model_sha256=str(payload["model_sha256"]),
input_columns=tuple(str(value) for value in payload["input_columns"]),
python_version=str(payload["python_version"]),
random_state=int(payload["random_state"]),
n_jobs=int(payload["n_jobs"]),
label_model=str(payload["label_model"]),
severity_model=str(payload["severity_model"]),
ood_absolute_threshold_mv=float(
payload["ood_absolute_threshold_mv"]
),
ood_ratio_threshold=float(payload["ood_ratio_threshold"]),
runtime_versions={
str(name): str(package_version)
for name, package_version in dict(
payload["runtime_versions"]
).items()
},
)
except (KeyError, TypeError, ValueError) as exc:
raise ModelArtifactError("Model manifest has an invalid schema.") from exc
def to_dict(self) -> dict[str, object]:
payload = asdict(self)
payload["input_columns"] = list(self.input_columns)
return payload
@dataclass(frozen=True)
class LoadedModelArtifact:
models: DiagnosticModels
metadata: ArtifactMetadata
def sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def _runtime_versions() -> dict[str, str]:
return {package: version(package) for package in RUNTIME_PACKAGES}
def _atomic_write(path: Path, payload: bytes) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
descriptor, temporary_name = tempfile.mkstemp(
dir=path.parent, prefix=f".{path.name}.", suffix=".tmp"
)
temporary_path = Path(temporary_name)
try:
with os.fdopen(descriptor, "wb") as handle:
handle.write(payload)
handle.flush()
os.fsync(handle.fileno())
os.replace(temporary_path, path)
except Exception:
temporary_path.unlink(missing_ok=True)
raise
def save_model_artifact(
models: DiagnosticModels,
directory: Path,
*,
model_version: str,
source_revision: str,
training_data_sha256: str,
random_state: int,
n_jobs: int,
label_model: str,
severity_model: str,
) -> ArtifactMetadata:
if not model_version.strip():
raise ValueError("model_version cannot be empty.")
if len(training_data_sha256) != 64 or any(
character not in "0123456789abcdef" for character in training_data_sha256
):
raise ValueError("training_data_sha256 must be a SHA-256 hex digest.")
envelope = {
"artifact_format_version": ARTIFACT_FORMAT_VERSION,
"model_version": model_version,
"models": models,
}
model_payload = pickle.dumps(envelope, protocol=pickle.HIGHEST_PROTOCOL)
model_sha256 = hashlib.sha256(model_payload).hexdigest()
metadata = ArtifactMetadata(
artifact_format_version=ARTIFACT_FORMAT_VERSION,
model_version=model_version,
created_at_utc=datetime.now(UTC).isoformat(),
source_revision=source_revision,
training_data_sha256=training_data_sha256,
model_sha256=model_sha256,
input_columns=tuple(INFERENCE_COLUMNS),
python_version=platform.python_version(),
random_state=random_state,
n_jobs=n_jobs,
label_model=label_model,
severity_model=severity_model,
ood_absolute_threshold_mv=OOD_OK_TO_UNKNOWN_THRESHOLD_MV,
ood_ratio_threshold=OOD_OK_TO_UNKNOWN_RATIO_THRESHOLD,
runtime_versions=_runtime_versions(),
)
directory.mkdir(parents=True, exist_ok=True)
_atomic_write(directory / MODEL_FILENAME, model_payload)
manifest_payload = json.dumps(
metadata.to_dict(), ensure_ascii=False, indent=2, sort_keys=True
).encode("utf-8") + b"\n"
_atomic_write(directory / MANIFEST_FILENAME, manifest_payload)
return metadata
def load_model_artifact(
directory: Path, *, expected_model_version: str | None = None
) -> LoadedModelArtifact:
manifest_path = directory / MANIFEST_FILENAME
model_path = directory / MODEL_FILENAME
try:
manifest_payload = json.loads(manifest_path.read_text(encoding="utf-8"))
except FileNotFoundError as exc:
raise ModelArtifactError(f"Model manifest is missing: {manifest_path}") from exc
except (OSError, json.JSONDecodeError) as exc:
raise ModelArtifactError("Model manifest cannot be read.") from exc
if not isinstance(manifest_payload, dict):
raise ModelArtifactError("Model manifest must be a JSON object.")
metadata = ArtifactMetadata.from_dict(manifest_payload)
if metadata.artifact_format_version != ARTIFACT_FORMAT_VERSION:
raise ModelArtifactError(
"Unsupported model artifact format: "
f"{metadata.artifact_format_version}."
)
if expected_model_version and metadata.model_version != expected_model_version:
raise ModelArtifactError(
f"Expected model {expected_model_version}, received {metadata.model_version}."
)
if metadata.input_columns != tuple(INFERENCE_COLUMNS):
raise ModelArtifactError("Model input schema does not match this application.")
expected_python = metadata.python_version.split(".")[:2]
installed_python = platform.python_version().split(".")[:2]
if expected_python != installed_python:
raise ModelArtifactError(
"Model Python version does not match the runtime. "
f"expected={metadata.python_version}, installed={platform.python_version()}"
)
current_versions = _runtime_versions()
if metadata.runtime_versions != current_versions:
raise ModelArtifactError(
"Model runtime versions do not match installed dependencies. "
f"expected={metadata.runtime_versions}, installed={current_versions}"
)
try:
model_payload = model_path.read_bytes()
except OSError as exc:
raise ModelArtifactError(f"Model file cannot be read: {model_path}") from exc
actual_sha256 = hashlib.sha256(model_payload).hexdigest()
if actual_sha256 != metadata.model_sha256:
raise ModelArtifactError("Model checksum verification failed.")
try:
envelope = pickle.loads(model_payload) # noqa: S301 - trusted, checksummed release artifact
except Exception as exc:
raise ModelArtifactError("Model payload cannot be deserialized.") from exc
if not isinstance(envelope, dict):
raise ModelArtifactError("Model payload has an invalid envelope.")
if envelope.get("artifact_format_version") != ARTIFACT_FORMAT_VERSION:
raise ModelArtifactError("Model payload format does not match the manifest.")
if envelope.get("model_version") != metadata.model_version:
raise ModelArtifactError("Model payload version does not match the manifest.")
models = envelope.get("models")
if not isinstance(models, DiagnosticModels):
raise ModelArtifactError("Model payload has an unexpected object type.")
return LoadedModelArtifact(models=models, metadata=metadata)
__all__ = [
"ARTIFACT_FORMAT_VERSION",
"MANIFEST_FILENAME",
"MODEL_FILENAME",
"ArtifactMetadata",
"LoadedModelArtifact",
"ModelArtifactError",
"load_model_artifact",
"save_model_artifact",
"sha256_file",
]