Files
SystemSimulationApp/tests/test_test_mql_output_validation.py
T

182 lines
6.8 KiB
Python

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()