Merge model-development into main
This commit is contained in:
commit
127ec36a55
218 files changed
+65243
-47
No files matched your search
@@ -0,0 +1,161 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from math import isfinite
|
||||
|
||||
from PythonModels.reporting.amesim_results import AmesimResults
|
||||
from PythonModels.reporting.test_mql_comparison import (
|
||||
TestMqlComparisonResult,
|
||||
compare_test_mql_series,
|
||||
)
|
||||
from PythonModels.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
|
||||
Reference in new issue
Block a user