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()