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