Files
SystemSimulationApp/tests/test_system_numeric_ir_v2.py
T
2026-09-02 19:17:55 +08:00

768 lines
26 KiB
Python

from __future__ import annotations
from dataclasses import fields, is_dataclass, replace
import json
from pathlib import Path
import unittest
from app.main import compile_reactflow_network, compile_system_xml_network
from app.simulation.ir import schema as ir_schema
from app.simulation.ir import (
IRBufferKind,
IRCapabilityLevel,
IRDType,
IREntryPointKind,
IRKernelAvailability,
IRKernelCallOperation,
IRNativeBuildIdentity,
IRStageKind,
IRStepKind,
canonical_json_bytes,
compile_system_ir,
native_artifact_key,
)
from app.simulation.ir.validation import (
SystemIRValidationError,
require_valid_system_ir,
validate_system_ir,
)
from app.simulation.systems.generic import GenericFluidSystem
from app.system_xml import validate_system_xml_document
from tests.test_amesim_mechanical_xml import zero_force_mass_project
TARGET_XML = Path("tests/data/test-mql-8.xml")
HISTORICAL_XML = Path("tests/data/test_mql-full-branches-01-04.xml")
MACHINE_SCHEMA = Path("schemas/system-numeric-ir-v2.schema.json")
def _system_from_xml(path: Path) -> GenericFluidSystem:
report = validate_system_xml_document(path.read_bytes())
if not report.valid or report.document is None:
raise AssertionError(report.as_dict())
return GenericFluidSystem(compile_system_xml_network(report.document))
def _entry_stage_kinds(program, entry_kind: IREntryPointKind) -> set[IRStageKind]:
entry = next(item for item in program.entry_points if item.kind is entry_kind)
result: set[IRStageKind] = set()
visited_blocks: set[int] = set()
def visit(step) -> None:
if step.kind is IRStepKind.STAGE:
result.add(program.stages[step.index].kind)
return
if step.index in visited_blocks:
return
visited_blocks.add(step.index)
for nested in program.blocks[step.index].steps:
visit(nested)
for step in entry.steps:
visit(step)
return result
def _assert_callback_free(test: unittest.TestCase, value: object) -> None:
if is_dataclass(value) and not isinstance(value, type):
for item in fields(value):
_assert_callback_free(test, getattr(value, item.name))
return
if isinstance(value, tuple):
for item in value:
_assert_callback_free(test, item)
return
test.assertFalse(callable(value), type(value).__name__)
test.assertNotIsInstance(value, (dict, list, set))
class SystemNumericIRV2ContractTests(unittest.TestCase):
@classmethod
def setUpClass(cls) -> None:
cls.system = GenericFluidSystem(
compile_reactflow_network(zero_force_mass_project())
)
cls.program = compile_system_ir(cls.system)
def test_compiler_returns_a_statically_valid_callback_free_program(self) -> None:
report = validate_system_ir(self.program)
self.assertTrue(report.valid, report.issues)
self.assertIs(require_valid_system_ir(self.program), self.program)
_assert_callback_free(self, self.program)
with self.assertRaises(TypeError):
canonical_json_bytes(lambda: None)
def test_canonical_bytes_and_signature_are_deterministic(self) -> None:
second = compile_system_ir(
GenericFluidSystem(
compile_reactflow_network(zero_force_mass_project())
)
)
self.assertEqual(
self.program.canonical_json_bytes(),
second.canonical_json_bytes(),
)
self.assertEqual(
self.program.structural_signature,
second.structural_signature,
)
self.assertEqual(len(self.program.structural_signature), 64)
self.assertEqual(
self.program.structural_signature,
self.program.calculate_structural_signature(),
)
decomposed = replace(self.program, model_id="e\u0301")
composed = replace(self.program, model_id="é")
self.assertEqual(
decomposed.canonical_json_bytes(),
composed.canonical_json_bytes(),
)
def test_parameter_change_invalidates_the_program_signature(self) -> None:
changed_project = zero_force_mass_project()
changed_project.nodes[1].data.parameters["mass"] = 3.0
changed = compile_system_ir(
GenericFluidSystem(compile_reactflow_network(changed_project))
)
self.assertNotEqual(
self.program.structural_signature,
changed.structural_signature,
)
def test_native_artifact_key_is_separate_from_structural_signature(self) -> None:
windows = IRNativeBuildIdentity(
abi_version=1,
target_triple="x86_64-pc-windows-msvc",
compiler_id="msvc",
compiler_version="19.40",
compile_flags=("/O2", "/fp:precise"),
floating_point_policy="strict",
kernel_library_signature="a" * 64,
)
linux = replace(
windows,
target_triple="x86_64-unknown-linux-gnu",
compiler_id="gcc",
compiler_version="14.2",
compile_flags=("-O2", "-fno-fast-math"),
)
signature = self.program.structural_signature
self.assertNotEqual(
native_artifact_key(self.program, windows),
native_artifact_key(self.program, linux),
)
self.assertEqual(self.program.structural_signature, signature)
with self.assertRaises(ValueError):
native_artifact_key(
self.program,
replace(windows, abi_version=windows.abi_version + 1),
)
def test_four_entry_points_are_independent(self) -> None:
self.assertEqual(
{item.kind for item in self.program.entry_points},
set(IREntryPointKind),
)
rhs_kinds = _entry_stage_kinds(self.program, IREntryPointKind.RHS)
event_kinds = _entry_stage_kinds(
self.program, IREntryPointKind.EVENTS
)
self.assertIn(IRStageKind.DERIVATIVE_REDUCE, rhs_kinds)
self.assertNotIn(IRStageKind.EVENT, rhs_kinds)
self.assertNotIn(IRStageKind.JACOBIAN, rhs_kinds)
self.assertNotIn(IRStageKind.OUTPUT, rhs_kinds)
self.assertIn(IRStageKind.EVENT, event_kinds)
self.assertNotIn(IRStageKind.JACOBIAN, event_kinds)
self.assertNotIn(IRStageKind.OUTPUT, event_kinds)
def test_validator_rejects_a_missing_entry_point(self) -> None:
broken = replace(
self.program,
entry_points=self.program.entry_points[:-1],
)
report = validate_system_ir(broken)
self.assertFalse(report.valid)
self.assertIn(
"ENTRY_POINT_SET_INVALID",
{item.code for item in report.issues},
)
with self.assertRaises(SystemIRValidationError):
require_valid_system_ir(broken)
def test_validator_rejects_entry_contract_and_buffer_dtype_corruption(self) -> None:
rhs_index = next(
index
for index, entry in enumerate(self.program.entry_points)
if entry.kind is IREntryPointKind.RHS
)
events_entry = next(
entry
for entry in self.program.entry_points
if entry.kind is IREntryPointKind.EVENTS
)
corrupted_entries = list(self.program.entry_points)
corrupted_entries[rhs_index] = replace(
corrupted_entries[rhs_index],
steps=events_entry.steps,
output_slots=(),
)
entry_report = validate_system_ir(
replace(self.program, entry_points=tuple(corrupted_entries))
)
self.assertTrue(
{
"ENTRY_POINT_OUTPUT_COVERAGE",
"ENTRY_POINT_FINAL_STAGE_MISSING",
"ENTRY_POINT_STAGE_FORBIDDEN",
}.issubset({item.code for item in entry_report.issues})
)
state_buffer_index = next(
index
for index, buffer in enumerate(self.program.buffers)
if buffer.kind is IRBufferKind.STATE_INPUT
)
corrupted_buffers = list(self.program.buffers)
corrupted_buffers[state_buffer_index] = replace(
corrupted_buffers[state_buffer_index],
dtype=IRDType.INT32,
initial_float_values=(),
initial_int_values=tuple(
0 for _ in range(corrupted_buffers[state_buffer_index].size)
),
)
dtype_report = validate_system_ir(
replace(self.program, buffers=tuple(corrupted_buffers))
)
self.assertIn(
"BUFFER_DTYPE_INVALID",
{item.code for item in dtype_report.issues},
)
def test_validator_rejects_a_component_call_bound_to_another_kernel(self) -> None:
stage_index, operation_index, operation = next(
(stage_index, operation_index, operation)
for stage_index, stage in enumerate(self.program.stages)
for operation_index, operation in enumerate(stage.operations)
if isinstance(operation, IRKernelCallOperation)
and operation.component_index is not None
)
wrong_kernel_index = next(
index
for index in range(len(self.program.kernels))
if index != operation.kernel_index
)
broken_operations = list(self.program.stages[stage_index].operations)
broken_operations[operation_index] = replace(
operation,
kernel_index=wrong_kernel_index,
)
broken_stages = list(self.program.stages)
broken_stages[stage_index] = replace(
broken_stages[stage_index],
operations=tuple(broken_operations),
)
report = validate_system_ir(
replace(self.program, stages=tuple(broken_stages))
)
self.assertFalse(report.valid)
self.assertIn(
"KERNEL_COMPONENT_MISMATCH",
{item.code for item in report.issues},
)
def test_validator_rejects_native_component_with_missing_called_phases(self) -> None:
native_kernels = tuple(
replace(
kernel,
availability=IRKernelAvailability.NATIVE,
unavailable_reason=None,
)
for kernel in self.program.kernels
)
native_capabilities = tuple(
replace(
capability,
level=IRCapabilityLevel.NATIVE,
supported_phases=(),
missing_features=(),
)
for capability in self.program.capabilities.components
)
broken = replace(
self.program,
kernels=native_kernels,
required_features=tuple(
feature
for feature in self.program.required_features
if feature != "reference_kernel_dispatch"
),
transaction=replace(
self.program.transaction,
cache_attribute_ids=(),
),
capabilities=replace(
self.program.capabilities,
system_level=IRCapabilityLevel.NATIVE,
components=native_capabilities,
issues=(),
),
)
report = validate_system_ir(broken)
self.assertFalse(report.valid)
self.assertIn(
"CAPABILITY_NATIVE_PHASE_MISSING",
{item.code for item in report.issues},
)
def test_validator_rejects_runtime_types_that_break_the_wire_schema(self) -> None:
outputs = list(self.program.outputs)
outputs[0] = replace(outputs[0], scale=1)
report = validate_system_ir(
replace(self.program, outputs=tuple(outputs))
)
self.assertFalse(report.valid)
self.assertIn(
"RUNTIME_TYPE_MISMATCH",
{item.code for item in report.issues},
)
def test_validator_propagates_unsupported_component_to_system_level(self) -> None:
capabilities = list(self.program.capabilities.components)
capabilities[0] = replace(
capabilities[0],
level=IRCapabilityLevel.UNSUPPORTED,
)
report = validate_system_ir(
replace(
self.program,
capabilities=replace(
self.program.capabilities,
components=tuple(capabilities),
),
)
)
self.assertFalse(report.valid)
self.assertIn(
"CAPABILITY_LEVEL_CONFLICT",
{item.code for item in report.issues},
)
def test_machine_schema_has_no_dangling_local_references(self) -> None:
schema = json.loads(MACHINE_SCHEMA.read_text(encoding="utf-8"))
definitions = schema["$defs"]
references: list[str] = []
pending: list[object] = [schema]
while pending:
current = pending.pop()
if isinstance(current, dict):
references.extend(
value
for key, value in current.items()
if key == "$ref" and isinstance(value, str)
)
pending.extend(current.values())
elif isinstance(current, list):
pending.extend(current)
self.assertFalse(
{
reference
for reference in references
if reference.startswith("#/$defs/")
and reference.removeprefix("#/$defs/") not in definitions
}
)
payload = json.loads(self.program.canonical_json_bytes())
self.assertEqual(payload["$type"], "system_ir")
self.assertEqual(
set(payload),
set(definitions["system_ir"]["required"]),
)
self.assertEqual(
definitions["kernel_phase"]["required"],
["$type", "phase"],
)
def test_machine_schema_fields_match_every_serialized_dataclass(self) -> None:
definitions = json.loads(
MACHINE_SCHEMA.read_text(encoding="utf-8")
)["$defs"]
skipped_types = {"native_build", "native_artifact_key_input"}
for value_type, canonical_type in ir_schema._CANONICAL_TYPE_NAMES:
if canonical_type in skipped_types:
continue
definition_name = (
value_type.opcode.value
if canonical_type == "operation"
else canonical_type
)
definition = definitions[definition_name]
expected_fields = {"$type", *(item.name for item in fields(value_type))}
if canonical_type == "operation":
expected_fields.add("opcode")
self.assertEqual(
set(definition["required"]),
expected_fields,
definition_name,
)
self.assertEqual(
set(definition["properties"]),
expected_fields,
definition_name,
)
self.assertFalse(
definition["additionalProperties"],
definition_name,
)
class TargetSystemNumericIRV2Tests(unittest.TestCase):
@classmethod
def setUpClass(cls) -> None:
cls.system = _system_from_xml(TARGET_XML)
cls.program = compile_system_ir(cls.system)
def test_target_model_is_fully_described(self) -> None:
program = self.program
self.assertEqual(len(program.components), 156)
self.assertEqual(len(program.ports), 356)
self.assertEqual(len(program.connections), 178)
self.assertEqual(program.state_reducer.solver_state_count, 132)
self.assertEqual(len(program.pressure_flow.unknowns), 776)
self.assertEqual(len(program.pressure_flow.equations), 776)
self.assertEqual(len(program.outputs), 1784)
self.assertEqual(program.jacobian.pattern.row_count, 132)
self.assertEqual(program.jacobian.pattern.column_count, 132)
self.assertGreater(program.jacobian.pattern.nonzero_count, 132)
self.assertLessEqual(len(program.jacobian.color_groups), 132)
self.assertEqual(len(program.causal_plans), 1)
causal = program.causal_plans[0]
self.assertEqual(len(causal.canonical_slots), 452)
self.assertEqual(len(causal.compatibility_slots), 776)
self.assertIsNone(causal.fallback_reason)
self.assertEqual(len(program.modes), 2)
self.assertEqual(len(program.events), 14)
self.assertTrue(program.thermofluid.sensitive_component_indices)
self.assertEqual(
program.thermofluid.secondary_pressure_scope_indices,
program.pressure_flow.secondary_scope_indices,
)
def test_state_reducer_and_all_buffers_have_complete_index_contracts(self) -> None:
reducer = self.program.state_reducer
self.assertEqual(
(reducer.state_scatter.pattern.row_count,
reducer.state_scatter.pattern.column_count),
(len(reducer.local_state_slots), reducer.solver_state_count),
)
self.assertEqual(
(reducer.derivative_gather.pattern.row_count,
reducer.derivative_gather.pattern.column_count),
(reducer.solver_state_count, len(reducer.raw_derivative_slots)),
)
by_buffer: dict[IRBufferKind, list[int]] = {
buffer.kind: [] for buffer in self.program.buffers
}
for value in self.program.values:
by_buffer[value.slot.buffer].append(value.slot.index)
for buffer in self.program.buffers:
self.assertEqual(
sorted(by_buffer[buffer.kind]),
list(range(buffer.size)),
buffer.kind,
)
def test_transaction_tracks_exactly_the_active_pneumatic_flows(self) -> None:
pneumatic_flow_slots = {
variable.slot
for port in self.program.ports
if port.kind.value == "physical" and port.domain == "pneumatic"
for variable in port.variables
if variable.name == "m_flow"
}
self.assertEqual(
set(self.program.transaction.flow_slots),
pneumatic_flow_slots,
)
def test_native_support_is_not_claimed_before_c02(self) -> None:
self.assertIs(
self.program.capabilities.system_level,
IRCapabilityLevel.REFERENCE_ONLY,
)
self.assertTrue(self.program.capabilities.components)
self.assertTrue(
all(
item.level is IRCapabilityLevel.REFERENCE_ONLY
for item in self.program.capabilities.components
)
)
self.assertIn(
"IR_NATIVE_KERNELS_NOT_DECLARED",
{item.code for item in self.program.capabilities.issues},
)
def test_validator_rejects_cross_plan_and_sparse_contract_corruption(self) -> None:
program = self.program
variants: list[tuple[str, object, str]] = []
components = list(program.components)
foreign_port = next(
index
for index, port in enumerate(program.ports)
if port.component_index != 0
)
components[0] = replace(
components[0],
port_indices=(*components[0].port_indices, foreign_port),
)
variants.append(
(
"component port back-reference",
replace(program, components=tuple(components)),
"COMPONENT_PORT_COVERAGE",
)
)
variants.append(
(
"thermofluid scope mismatch",
replace(
program,
thermofluid=replace(
program.thermofluid,
secondary_pressure_scope_indices=(),
),
),
"THERMOFLUID_SCOPE_MISMATCH",
)
)
equations = list(program.pressure_flow.equations)
equations[1] = replace(
equations[1],
residual_slot=equations[0].residual_slot,
)
variants.append(
(
"pressure-flow residual slot alias",
replace(
program,
pressure_flow=replace(
program.pressure_flow,
equations=tuple(equations),
),
),
"PRESSURE_FLOW_RESIDUAL_SLOT_DUPLICATE",
)
)
buffers = list(program.buffers)
state_buffer_index = next(
index
for index, buffer in enumerate(buffers)
if buffer.kind is IRBufferKind.STATE_INPUT
)
state_values = list(buffers[state_buffer_index].initial_float_values)
state_values[0] += 1.0
buffers[state_buffer_index] = replace(
buffers[state_buffer_index],
initial_float_values=tuple(state_values),
)
variants.append(
(
"state initial value disagreement",
replace(program, buffers=tuple(buffers)),
"STATE_INITIAL_VALUE_MISMATCH",
)
)
aliased_values = (
program.jacobian.value_slots[0],
program.jacobian.value_slots[0],
*program.jacobian.value_slots[2:],
)
variants.append(
(
"Jacobian value slot alias",
replace(
program,
jacobian=replace(
program.jacobian,
value_slots=aliased_values,
),
),
"JACOBIAN_VALUE_SLOT_COVERAGE",
)
)
variants.append(
(
"Jacobian color conflict",
replace(
program,
jacobian=replace(
program.jacobian,
color_groups=(
tuple(range(program.state_reducer.solver_state_count)),
),
),
),
"JACOBIAN_COLOR_CONFLICT",
)
)
fd_columns = list(program.jacobian.local_finite_difference_columns)
fd_column = fd_columns[0]
wrong_value_index = next(
index
for index, column in enumerate(
program.jacobian.pattern.column_indices
)
if column != fd_column.column_index
)
fd_columns[0] = replace(
fd_column,
value_indices=(wrong_value_index, *fd_column.value_indices[1:]),
)
variants.append(
(
"Jacobian finite-difference column mismatch",
replace(
program,
jacobian=replace(
program.jacobian,
local_finite_difference_columns=tuple(fd_columns),
),
),
"JACOBIAN_FD_COLUMN_MISMATCH",
)
)
capabilities = list(program.capabilities.components)
capabilities[0] = replace(
capabilities[0],
level=IRCapabilityLevel.NATIVE,
missing_features=(),
)
variants.append(
(
"native capability overclaim",
replace(
program,
capabilities=replace(
program.capabilities,
components=tuple(capabilities),
),
),
"CAPABILITY_KERNEL_MISMATCH",
)
)
variants.append(
(
"duplicate transaction flow",
replace(
program,
transaction=replace(
program.transaction,
flow_slots=(
*program.transaction.flow_slots,
program.transaction.flow_slots[0],
),
),
),
"TRANSACTION_DUPLICATE_SLOT",
)
)
modes = list(program.modes)
modes[1] = replace(modes[1], slot=modes[0].slot)
variants.append(
(
"duplicate mode slot",
replace(program, modes=tuple(modes)),
"MODE_SLOT_DUPLICATE",
)
)
gather_pattern = program.state_reducer.derivative_gather.pattern
gather_pointers = list(gather_pattern.row_pointers)
gather_pointers[1] = gather_pointers[0]
variants.append(
(
"empty derivative row",
replace(
program,
state_reducer=replace(
program.state_reducer,
derivative_gather=replace(
program.state_reducer.derivative_gather,
pattern=replace(
gather_pattern,
row_pointers=tuple(gather_pointers),
),
),
),
),
"DERIVATIVE_GATHER_EMPTY_ROW",
)
)
for label, corrupted, expected_code in variants:
with self.subTest(label=label):
report = validate_system_ir(corrupted)
self.assertFalse(report.valid)
self.assertIn(
expected_code,
{item.code for item in report.issues},
)
def test_target_recompilation_is_byte_stable(self) -> None:
second = compile_system_ir(_system_from_xml(TARGET_XML))
self.assertEqual(
self.program.canonical_json_bytes(),
second.canonical_json_bytes(),
)
class HistoricalSystemNumericIRV2Tests(unittest.TestCase):
def test_historical_complex_model_also_compiles(self) -> None:
program = compile_system_ir(_system_from_xml(HISTORICAL_XML))
self.assertTrue(validate_system_ir(program).valid)
self.assertEqual(len(program.components), 98)
self.assertEqual(len(program.connections), 106)
self.assertEqual(program.state_reducer.solver_state_count, 74)
self.assertEqual(len(program.pressure_flow.unknowns), 472)
self.assertEqual(len(program.pressure_flow.equations), 472)
self.assertEqual(len(program.outputs), 1021)
if __name__ == "__main__":
unittest.main()