from __future__ import annotations from dataclasses import dataclass from math import isfinite from app.simulation.reporting.amesim_results import AmesimResults from app.simulation.reporting.test_mql_comparison import ( TestMqlComparisonResult, compare_test_mql_series, ) from app.simulation.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