hackathon-ENGIN/tests/test_model_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

115 lines
4.2 KiB
Python

from __future__ import annotations
import json
import tempfile
import unittest
from pathlib import Path
from unittest.mock import patch
import pandas as pd
from pandas.testing import assert_frame_equal
import app
from engin.artifact import (
MANIFEST_FILENAME,
MODEL_FILENAME,
ModelArtifactError,
load_model_artifact,
save_model_artifact,
)
from engin.inference import DiagnosticModels
from engin.model import SklearnPredictionModel
ROOT = Path(__file__).resolve().parents[1]
MODEL_VERSION = "engin-2026.08.25-1"
ARTIFACT_DIR = ROOT / "artifacts" / MODEL_VERSION
def _save_fake_artifact(directory: Path) -> None:
save_model_artifact(
DiagnosticModels("label", "transformer", "estimator"),
directory,
model_version="test-v1",
source_revision="unit-test",
training_data_sha256="a" * 64,
random_state=42,
n_jobs=1,
label_model="fake-label",
severity_model="fake-severity",
)
class ModelArtifactContractTests(unittest.TestCase):
def test_round_trip_preserves_metadata_and_models(self) -> None:
with tempfile.TemporaryDirectory() as temporary_directory:
directory = Path(temporary_directory)
_save_fake_artifact(directory)
loaded = load_model_artifact(
directory, expected_model_version="test-v1"
)
self.assertEqual(loaded.metadata.model_version, "test-v1")
self.assertEqual(loaded.models.label_pipeline, "label")
def test_corrupt_model_is_rejected_before_deserialization(self) -> None:
with tempfile.TemporaryDirectory() as temporary_directory:
directory = Path(temporary_directory)
_save_fake_artifact(directory)
with (directory / MODEL_FILENAME).open("ab") as handle:
handle.write(b"corruption")
with self.assertRaisesRegex(ModelArtifactError, "checksum"):
load_model_artifact(directory)
def test_runtime_version_mismatch_is_rejected(self) -> None:
with tempfile.TemporaryDirectory() as temporary_directory:
directory = Path(temporary_directory)
_save_fake_artifact(directory)
manifest_path = directory / MANIFEST_FILENAME
manifest = json.loads(manifest_path.read_text(encoding="utf-8"))
manifest["runtime_versions"]["scikit-learn"] = "0.0-invalid"
manifest_path.write_text(
json.dumps(manifest), encoding="utf-8"
)
with self.assertRaisesRegex(ModelArtifactError, "runtime versions"):
load_model_artifact(directory)
class ShippedArtifactTests(unittest.TestCase):
def test_artifact_reproduces_frozen_submission(self) -> None:
model = SklearnPredictionModel.from_artifact(
ARTIFACT_DIR,
expected_model_version=MODEL_VERSION,
)
test = pd.read_csv(ROOT / "test.csv").reset_index(drop=True)
expected = pd.read_csv(ROOT / "predictions.csv").reset_index(drop=True)
actual = model.predict(test).submission
assert_frame_equal(actual, expected, check_dtype=False, check_exact=True)
def test_diagnostics_identify_model_release(self) -> None:
model = SklearnPredictionModel.from_artifact(
ARTIFACT_DIR,
expected_model_version=MODEL_VERSION,
)
test = pd.read_csv(ROOT / "test.csv").reset_index(drop=True)
diagnostics = model.predict(test).diagnostics
self.assertEqual(set(diagnostics["model_version"]), {MODEL_VERSION})
self.assertEqual(
set(diagnostics["model_artifact_sha256"]),
{model.metadata.model_sha256},
)
def test_application_startup_does_not_read_training_data(self) -> None:
app.build_dependencies.clear()
with patch.object(
app.PandasCsvReader,
"read_path",
side_effect=AssertionError("runtime attempted to read training data"),
):
dependencies = app.build_dependencies()
app.build_dependencies.clear()
self.assertTrue(dependencies.live_inference)
self.assertEqual(dependencies.model_version, MODEL_VERSION)
if __name__ == "__main__":
unittest.main(verbosity=2)