89 lines
3.2 KiB
Python
89 lines
3.2 KiB
Python
from __future__ import annotations
|
|
|
|
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
from PythonModels.scripts.run_test_mql_baseline import (
|
|
TestMqlBaselineExecutionConfig,
|
|
TestMqlBaselinePathConfig,
|
|
TestMqlBaselineScriptConfig,
|
|
format_test_mql_baseline_summary,
|
|
run_test_mql_baseline,
|
|
)
|
|
|
|
|
|
REPO_ROOT = Path(__file__).resolve().parents[1]
|
|
TEST_MQL_AME = REPO_ROOT / "AmesimModels" / "test_mql.ame"
|
|
|
|
|
|
class RunTestMqlBaselineScriptTests(unittest.TestCase):
|
|
def test_baseline_script_writes_summary_for_selected_paths(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmpdir:
|
|
output_dir = Path(tmpdir)
|
|
run, returned_output_dir = run_test_mql_baseline(
|
|
TestMqlBaselineScriptConfig(
|
|
paths=TestMqlBaselinePathConfig(
|
|
archive_path=TEST_MQL_AME,
|
|
output_dir=output_dir,
|
|
),
|
|
execution=TestMqlBaselineExecutionConfig(
|
|
data_paths=("press@pn_c1_8", "vol1@pn_brp2_8"),
|
|
),
|
|
)
|
|
)
|
|
|
|
summary_path = output_dir / "test_mql_baseline_summary.txt"
|
|
summary = summary_path.read_text(encoding="utf-8")
|
|
|
|
self.assertEqual(returned_output_dir, output_dir)
|
|
self.assertEqual(run.sample_count, 1002)
|
|
self.assertEqual(run.signal_count, 2)
|
|
self.assertEqual(run.comparison.max_abs_error, 0.0)
|
|
self.assertIn("Mode: AMESim baseline passthrough", summary)
|
|
self.assertIn("Samples: 1002", summary)
|
|
self.assertIn("Compared signals: 2", summary)
|
|
self.assertIn("Max absolute error: 0.0", summary)
|
|
|
|
def test_baseline_script_can_skip_summary_file(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmpdir:
|
|
output_dir = Path(tmpdir)
|
|
run, returned_output_dir = run_test_mql_baseline(
|
|
TestMqlBaselineScriptConfig(
|
|
paths=TestMqlBaselinePathConfig(
|
|
archive_path=TEST_MQL_AME,
|
|
output_dir=output_dir,
|
|
),
|
|
execution=TestMqlBaselineExecutionConfig(
|
|
write_summary=False,
|
|
data_paths=("press@pn_c1_8",),
|
|
),
|
|
)
|
|
)
|
|
|
|
self.assertEqual(returned_output_dir, output_dir)
|
|
self.assertEqual(run.signal_count, 1)
|
|
self.assertFalse((output_dir / "test_mql_baseline_summary.txt").exists())
|
|
|
|
def test_format_test_mql_baseline_summary(self) -> None:
|
|
run, _ = run_test_mql_baseline(
|
|
TestMqlBaselineScriptConfig(
|
|
paths=TestMqlBaselinePathConfig(archive_path=TEST_MQL_AME),
|
|
execution=TestMqlBaselineExecutionConfig(
|
|
write_summary=False,
|
|
data_paths=("press@pn_c1_8",),
|
|
),
|
|
)
|
|
)
|
|
|
|
summary = format_test_mql_baseline_summary(run)
|
|
|
|
self.assertIn("Model: test_mql", summary)
|
|
self.assertIn("Output schema signals: 858", summary)
|
|
self.assertIn("Compared signals: 1", summary)
|
|
self.assertTrue(summary.endswith("\n"))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|