from __future__ import annotations import unittest from pathlib import Path from app.simulation.reporting.amesim_results import load_test_mql_amesim_results from app.simulation.reporting.test_mql_observations import build_test_mql_observation_catalog from app.simulation.reporting.test_mql_output_schema import build_test_mql_output_schema from app.simulation.reporting.test_mql_output_validation import ( TestMqlOutputValidationError, compare_validated_test_mql_output, validate_test_mql_output, ) REPO_ROOT = Path(__file__).resolve().parents[1] TEST_MQL_AME = REPO_ROOT / "AmesimModels" / "test_mql.ame" class TestMqlOutputValidationTests(unittest.TestCase): @classmethod def setUpClass(cls) -> None: cls.amesim_results = load_test_mql_amesim_results(TEST_MQL_AME) cls.observation_catalog = build_test_mql_observation_catalog(cls.amesim_results) cls.schema = build_test_mql_output_schema( cls.amesim_results, observation_catalog=cls.observation_catalog, ) def test_validates_selected_python_output_series(self) -> None: times = [0.0, 0.1, 0.2] series = { "press@pn_c1_8": [1.0, 2.0, 3.0], "vol1@pn_brp2_8": [4, 5, 6], } validated = validate_test_mql_output( times=times, series_by_data_path=series, schema=self.schema, ) self.assertEqual(validated.times, (0.0, 0.1, 0.2)) self.assertEqual(validated.data_paths, ("press@pn_c1_8", "vol1@pn_brp2_8")) self.assertEqual(validated.series("press@pn_c1_8"), (1.0, 2.0, 3.0)) self.assertEqual(validated.series("vol1@pn_brp2_8"), (4.0, 5.0, 6.0)) def test_can_validate_explicit_data_path_order(self) -> None: validated = validate_test_mql_output( times=(0.0, 0.1), series_by_data_path={ "press@pn_c1_8": (10.0, 11.0), "vol1@pn_brp2_8": (20.0, 21.0), }, schema=self.schema, data_paths=("vol1@pn_brp2_8", "press@pn_c1_8"), ) self.assertEqual(validated.data_paths, ("vol1@pn_brp2_8", "press@pn_c1_8")) def test_rejects_output_paths_outside_schema_by_default(self) -> None: with self.assertRaisesRegex(TestMqlOutputValidationError, "outside test_mql schema"): validate_test_mql_output( times=(0.0,), series_by_data_path={"not_a_signal@not_an_owner": (1.0,)}, schema=self.schema, ) def test_can_ignore_extra_output_paths_when_requested(self) -> None: validated = validate_test_mql_output( times=(0.0,), series_by_data_path={ "press@pn_c1_8": (1.0,), "not_a_signal@not_an_owner": (2.0,), }, schema=self.schema, allow_extra_paths=True, ) self.assertEqual(validated.data_paths, ("press@pn_c1_8",)) def test_rejects_missing_selected_data_path(self) -> None: with self.assertRaisesRegex(TestMqlOutputValidationError, "missing selected"): validate_test_mql_output( times=(0.0,), series_by_data_path={"press@pn_c1_8": (1.0,)}, schema=self.schema, data_paths=("press@pn_c1_8", "vol1@pn_brp2_8"), ) def test_rejects_incomplete_full_schema_output(self) -> None: with self.assertRaisesRegex(TestMqlOutputValidationError, "missing required"): validate_test_mql_output( times=(0.0,), series_by_data_path={"press@pn_c1_8": (1.0,)}, schema=self.schema, require_all_schema_paths=True, ) def test_rejects_length_mismatch(self) -> None: with self.assertRaisesRegex(TestMqlOutputValidationError, "length mismatch"): validate_test_mql_output( times=(0.0, 0.1), series_by_data_path={"press@pn_c1_8": (1.0,)}, schema=self.schema, ) def test_rejects_non_finite_values(self) -> None: with self.assertRaisesRegex(TestMqlOutputValidationError, "non-finite"): validate_test_mql_output( times=(0.0,), series_by_data_path={"press@pn_c1_8": (float("nan"),)}, schema=self.schema, ) def test_rejects_non_monotonic_time_axis(self) -> None: with self.assertRaisesRegex(TestMqlOutputValidationError, "monotonically increasing"): validate_test_mql_output( times=(0.0, 0.2, 0.1), series_by_data_path={"press@pn_c1_8": (1.0, 2.0, 3.0)}, schema=self.schema, ) def test_compares_validated_amesim_baseline_with_zero_error(self) -> None: selected_paths = ( "press@pn_c1_8", "dm2@pn_morifice_11", "x1@mass_friction_endstops_10", ) baseline = self.observation_catalog.baseline_series_by_data_path( self.amesim_results, selected_paths, ) comparison = compare_validated_test_mql_output( times=self.amesim_results.times, series_by_data_path=baseline, schema=self.schema, amesim_results=self.amesim_results, data_paths=selected_paths, ) self.assertEqual(len(comparison.metrics), len(selected_paths)) self.assertEqual(comparison.max_abs_error, 0.0) self.assertEqual(comparison.max_rel_error, 0.0) def test_compares_validated_output_and_reports_perturbation_error(self) -> None: selected_paths = ("press@pn_c1_8",) baseline = self.observation_catalog.baseline_series_by_data_path( self.amesim_results, selected_paths, ) perturbed = dict(baseline) values = list(perturbed["press@pn_c1_8"]) values[-1] += 10.0 perturbed["press@pn_c1_8"] = tuple(values) comparison = compare_validated_test_mql_output( times=self.amesim_results.times, series_by_data_path=perturbed, schema=self.schema, amesim_results=self.amesim_results, data_paths=selected_paths, ) metric = comparison.metric("press@pn_c1_8") self.assertAlmostEqual(metric.final_abs_error, 10.0) self.assertAlmostEqual(metric.max_abs_error, 10.0) def test_compare_validated_output_reuses_schema_validation(self) -> None: with self.assertRaisesRegex(TestMqlOutputValidationError, "outside test_mql schema"): compare_validated_test_mql_output( times=(0.0,), series_by_data_path={"not_a_signal@not_an_owner": (1.0,)}, schema=self.schema, amesim_results=self.amesim_results, ) if __name__ == "__main__": unittest.main()