182 lines
6.8 KiB
Python
182 lines
6.8 KiB
Python
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()
|