All checks were successful
ENGIN CI / Build, test and smoke (push) Successful in 2m58s
258 lines
9.1 KiB
Python
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",
|
|
]
|