"""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", ]