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)