from __future__ import annotations import tempfile import unittest from pathlib import Path from PythonModels.core.solver import SolveIVPConfig from PythonModels.scripts.run_test_mql_full_state_comparison import ( TestMqlFullStateComparisonExecutionConfig, TestMqlFullStateComparisonPathConfig, TestMqlFullStateComparisonScriptConfig, format_test_mql_full_state_comparison_summary, run_test_mql_full_state_comparison, ) REPO_ROOT = Path(__file__).resolve().parents[1] TEST_MQL_AME = REPO_ROOT / "AmesimModels" / "test_mql.ame" class RunTestMqlFullStateComparisonScriptTests(unittest.TestCase): def test_full_state_script_writes_summary_and_comparison_csv(self) -> None: with tempfile.TemporaryDirectory() as tmpdir: output_dir = Path(tmpdir) run, returned_output_dir = run_test_mql_full_state_comparison( TestMqlFullStateComparisonScriptConfig( paths=TestMqlFullStateComparisonPathConfig( archive_path=TEST_MQL_AME, output_dir=output_dir, ), execution=TestMqlFullStateComparisonExecutionConfig( data_paths=("press@pn_c1_8", "vol1@pn_brp2_8"), solver=SolveIVPConfig(t_start=0.0, t_stop=0.0), t_eval=(0.0,), ), ) ) summary_path = output_dir / "test_mql_full_state_comparison_summary.txt" csv_path = output_dir / "test_mql_amesim_comparison.csv" csv_summary_path = output_dir / "test_mql_amesim_comparison_summary.txt" summary = summary_path.read_text(encoding="utf-8") csv_header = csv_path.read_text(encoding="utf-8").splitlines()[0] csv_summary_exists = csv_summary_path.exists() self.assertEqual(returned_output_dir, output_dir) self.assertEqual(run.sample_count, 1) self.assertEqual(run.signal_count, 2) self.assertAlmostEqual(run.comparison.max_abs_error, 0.0, delta=1.0e-8) self.assertEqual( run.metrics_by_max_abs_error()[0].data_path, "press@pn_c1_8", ) diagnostic = run.signal_diagnostic("press@pn_c1_8") self.assertEqual(diagnostic.data_path, "press@pn_c1_8") self.assertAlmostEqual(diagnostic.initial_python_value, -1300.0) self.assertAlmostEqual(diagnostic.final_python_value, -1300.0) self.assertAlmostEqual(diagnostic.final_abs_error, 0.0, delta=1.0e-8) self.assertEqual( run.diagnostics_by_final_abs_error()[0].data_path, "press@pn_c1_8", ) self.assertTrue(csv_summary_exists) self.assertIn("Mode: Python 132 full-state closure comparison", summary) self.assertIn("Compared signals: 2", summary) self.assertIn("Largest absolute error: press@pn_c1_8=", summary) self.assertIn("Largest final endpoint error: press@pn_c1_8=", summary) self.assertIn("Metrics by max absolute error:", summary) self.assertIn("Endpoint diagnostics by final absolute error:", summary) self.assertIn("final_python=", summary) self.assertIn("final_amesim=", summary) self.assertIn("python.press@pn_c1_8", csv_header) self.assertIn("amesim.vol1@pn_brp2_8", csv_header) def test_full_state_script_can_skip_artifact_files(self) -> None: with tempfile.TemporaryDirectory() as tmpdir: output_dir = Path(tmpdir) run, returned_output_dir = run_test_mql_full_state_comparison( TestMqlFullStateComparisonScriptConfig( paths=TestMqlFullStateComparisonPathConfig( archive_path=TEST_MQL_AME, output_dir=output_dir, ), execution=TestMqlFullStateComparisonExecutionConfig( write_summary=False, write_comparison_csv=False, data_paths=("press@pn_c1_8",), solver=SolveIVPConfig(t_start=0.0, t_stop=0.0), t_eval=(0.0,), ), ) ) self.assertEqual(returned_output_dir, output_dir) self.assertEqual(run.signal_count, 1) self.assertFalse((output_dir / "test_mql_full_state_comparison_summary.txt").exists()) self.assertFalse((output_dir / "test_mql_amesim_comparison.csv").exists()) def test_format_test_mql_full_state_comparison_summary(self) -> None: run, _ = run_test_mql_full_state_comparison( TestMqlFullStateComparisonScriptConfig( paths=TestMqlFullStateComparisonPathConfig(archive_path=TEST_MQL_AME), execution=TestMqlFullStateComparisonExecutionConfig( write_summary=False, write_comparison_csv=False, data_paths=("press@pn_c1_8",), solver=SolveIVPConfig(t_start=0.0, t_stop=0.0), t_eval=(0.0,), ), ) ) summary = format_test_mql_full_state_comparison_summary(run) self.assertIn("Model: test_mql", summary) self.assertIn("Output schema signals: 858", summary) self.assertIn("Compared signals: 1", summary) self.assertIn("Largest absolute error: press@pn_c1_8=", summary) self.assertIn("Largest final endpoint error: press@pn_c1_8=", summary) self.assertIn("Metrics by max absolute error:", summary) self.assertIn("Endpoint diagnostics by final absolute error:", summary) self.assertTrue(summary.endswith("\n")) if __name__ == "__main__": unittest.main()