From 6ccbf7eb5e43a76b7283d0c61b7277e55c07e052 Mon Sep 17 00:00:00 2001 From: huojiarong Date: Thu, 16 Jul 2026 02:34:32 +0000 Subject: [PATCH] =?UTF-8?q?=E8=A1=A5=E5=85=85test=5Fmql=E8=BE=93=E5=87=BA?= =?UTF-8?q?=E6=A0=A1=E9=AA=8C=E5=AF=B9=E6=AF=94=E5=85=A5=E5=8F=A3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../reporting/test_mql_output_validation.py | 161 ++++++++++++++++ tests/test_test_mql_output_validation.py | 181 ++++++++++++++++++ 2 files changed, 342 insertions(+) create mode 100644 PythonModels/reporting/test_mql_output_validation.py create mode 100644 tests/test_test_mql_output_validation.py diff --git a/PythonModels/reporting/test_mql_output_validation.py b/PythonModels/reporting/test_mql_output_validation.py new file mode 100644 index 0000000..39333e2 --- /dev/null +++ b/PythonModels/reporting/test_mql_output_validation.py @@ -0,0 +1,161 @@ +from __future__ import annotations + +from dataclasses import dataclass +from math import isfinite + +from PythonModels.reporting.amesim_results import AmesimResults +from PythonModels.reporting.test_mql_comparison import ( + TestMqlComparisonResult, + compare_test_mql_series, +) +from PythonModels.reporting.test_mql_output_schema import TestMqlOutputSchema + + +class TestMqlOutputValidationError(ValueError): + """Raised when a Python test_mql output does not satisfy the AMESim output contract.""" + + +@dataclass(frozen=True) +class TestMqlValidatedOutput: + times: tuple[float, ...] + series_by_data_path: dict[str, tuple[float, ...]] + data_paths: tuple[str, ...] + + def series(self, data_path: str) -> tuple[float, ...]: + if data_path not in self.series_by_data_path: + raise KeyError(data_path) + return self.series_by_data_path[data_path] + + +def validate_test_mql_output( + *, + times: tuple[float, ...] | list[float], + series_by_data_path: dict[str, tuple[float, ...] | list[float]], + schema: TestMqlOutputSchema, + data_paths: tuple[str, ...] | list[str] | None = None, + allow_extra_paths: bool = False, + require_all_schema_paths: bool = False, +) -> TestMqlValidatedOutput: + validated_times = _validate_time_axis(times) + selected_paths = _select_paths( + series_by_data_path=series_by_data_path, + schema=schema, + data_paths=data_paths, + allow_extra_paths=allow_extra_paths, + require_all_schema_paths=require_all_schema_paths, + ) + validated_series = { + data_path: _validate_series( + data_path=data_path, + values=series_by_data_path[data_path], + expected_count=len(validated_times), + ) + for data_path in selected_paths + } + return TestMqlValidatedOutput( + times=validated_times, + series_by_data_path=validated_series, + data_paths=selected_paths, + ) + + +def compare_validated_test_mql_output( + *, + times: tuple[float, ...] | list[float], + series_by_data_path: dict[str, tuple[float, ...] | list[float]], + schema: TestMqlOutputSchema, + amesim_results: AmesimResults, + data_paths: tuple[str, ...] | list[str] | None = None, + allow_extra_paths: bool = False, + require_all_schema_paths: bool = False, + relative_floor: float = 1.0e-12, +) -> TestMqlComparisonResult: + validated = validate_test_mql_output( + times=times, + series_by_data_path=series_by_data_path, + schema=schema, + data_paths=data_paths, + allow_extra_paths=allow_extra_paths, + require_all_schema_paths=require_all_schema_paths, + ) + return compare_test_mql_series( + python_times=validated.times, + python_series_by_data_path=validated.series_by_data_path, + amesim_results=amesim_results, + data_paths=validated.data_paths, + relative_floor=relative_floor, + ) + + +def _validate_time_axis(times: tuple[float, ...] | list[float]) -> tuple[float, ...]: + if not times: + raise TestMqlOutputValidationError("Python time axis is empty.") + validated = tuple(_finite_float("time", value) for value in times) + previous = validated[0] + for value in validated[1:]: + if value < previous: + raise TestMqlOutputValidationError("Python time axis must be monotonically increasing.") + previous = value + return validated + + +def _select_paths( + *, + series_by_data_path: dict[str, tuple[float, ...] | list[float]], + schema: TestMqlOutputSchema, + data_paths: tuple[str, ...] | list[str] | None, + allow_extra_paths: bool, + require_all_schema_paths: bool, +) -> tuple[str, ...]: + schema_paths = set(schema.data_paths()) + provided_paths = set(series_by_data_path) + if not allow_extra_paths: + extra_paths = sorted(provided_paths - schema_paths) + if extra_paths: + raise TestMqlOutputValidationError( + f"Python output contains Data_Path values outside test_mql schema: {extra_paths}" + ) + if require_all_schema_paths: + missing_schema_paths = sorted(schema_paths - provided_paths) + if missing_schema_paths: + raise TestMqlOutputValidationError( + f"Python output is missing required test_mql schema Data_Path values: {missing_schema_paths}" + ) + selected_paths = tuple(data_paths) if data_paths is not None else tuple(sorted(provided_paths & schema_paths)) + if not selected_paths: + raise TestMqlOutputValidationError("no test_mql schema Data_Path values are available.") + unknown_selected = [data_path for data_path in selected_paths if data_path not in schema_paths] + if unknown_selected: + raise TestMqlOutputValidationError( + f"Requested Data_Path values are outside test_mql schema: {unknown_selected}" + ) + missing_selected = [data_path for data_path in selected_paths if data_path not in series_by_data_path] + if missing_selected: + raise TestMqlOutputValidationError( + f"Python output is missing selected Data_Path values: {missing_selected}" + ) + return selected_paths + + +def _validate_series( + *, + data_path: str, + values: tuple[float, ...] | list[float], + expected_count: int, +) -> tuple[float, ...]: + if len(values) != expected_count: + raise TestMqlOutputValidationError( + f"Python series length mismatch for {data_path!r}: " + f"{len(values)} values for {expected_count} time samples." + ) + return tuple(_finite_float(data_path, value) for value in values) + + +def _finite_float(label: str, value: float) -> float: + try: + numeric_value = float(value) + except (TypeError, ValueError) as exc: + raise TestMqlOutputValidationError(f"{label!r} contains a non-numeric value: {value!r}") from exc + if not isfinite(numeric_value): + raise TestMqlOutputValidationError(f"{label!r} contains a non-finite value: {value!r}") + return numeric_value diff --git a/tests/test_test_mql_output_validation.py b/tests/test_test_mql_output_validation.py new file mode 100644 index 0000000..e280371 --- /dev/null +++ b/tests/test_test_mql_output_validation.py @@ -0,0 +1,181 @@ +from __future__ import annotations + +import unittest +from pathlib import Path + +from PythonModels.reporting.amesim_results import load_test_mql_amesim_results +from PythonModels.reporting.test_mql_observations import build_test_mql_observation_catalog +from PythonModels.reporting.test_mql_output_schema import build_test_mql_output_schema +from PythonModels.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()