Merge system-optimization IR contracts and retire legacy engine adapter
This commit is contained in:
commit
743663e3a6
14 files changed
+3817
-161
No files matched your search
Binary file not shown.
@@ -0,0 +1,35 @@
|
||||
"""Read the archived IR schema examples without a Python numerical engine."""
|
||||
from dataclasses import is_dataclass
|
||||
from enum import Enum
|
||||
from functools import lru_cache
|
||||
import gzip
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
from app.simulation.ir import schema
|
||||
|
||||
|
||||
def _decode(value):
|
||||
if isinstance(value, list):
|
||||
return tuple(_decode(item) for item in value)
|
||||
if isinstance(value, dict):
|
||||
if 'enum' in value:
|
||||
cls = getattr(schema, value['enum'])
|
||||
if not isinstance(cls, type) or not issubclass(cls, Enum):
|
||||
raise ValueError('Invalid IR enum in fixture')
|
||||
return cls(value['value'])
|
||||
cls = getattr(schema, value['type'])
|
||||
if not isinstance(cls, type) or not is_dataclass(cls):
|
||||
raise ValueError('Invalid IR record in fixture')
|
||||
return cls(**{key: _decode(item) for key, item in value['fields'].items()})
|
||||
return value
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _references():
|
||||
path = Path(__file__).parent / 'data/system-ir-v2-reference.json.gz'
|
||||
return json.loads(gzip.decompress(path.read_bytes()))['programs']
|
||||
|
||||
|
||||
def reference_ir(name):
|
||||
return _decode(_references()[name])
|
||||
@@ -423,6 +423,8 @@ class GenericSystemXmlSimulationTests(unittest.TestCase):
|
||||
self.assertEqual(snapshot["result"]["status"], "stopped")
|
||||
|
||||
def test_streaming_endpoint_keeps_quiet_solver_connection_alive(self) -> None:
|
||||
heartbeat_observed = threading.Event()
|
||||
|
||||
def delayed_simulation(
|
||||
_xml_bytes,
|
||||
progress_callback,
|
||||
@@ -433,8 +435,10 @@ class GenericSystemXmlSimulationTests(unittest.TestCase):
|
||||
self.assertIsNotNone(activity_tracker)
|
||||
activity_tracker.start_integration(0.0)
|
||||
for trial_time in (0.0487, 0.0488, 0.0489, 0.0490):
|
||||
heartbeat_observed.clear()
|
||||
activity_tracker.record_rhs(trial_time)
|
||||
time.sleep(0.008)
|
||||
# Wait for the consumer, independent of machine scheduling.
|
||||
self.assertTrue(heartbeat_observed.wait(timeout=2.0))
|
||||
return {
|
||||
"success": True,
|
||||
"status": "completed",
|
||||
@@ -449,10 +453,12 @@ class GenericSystemXmlSimulationTests(unittest.TestCase):
|
||||
side_effect=delayed_simulation,
|
||||
),
|
||||
):
|
||||
events = [
|
||||
json.loads(line)
|
||||
for line in simulation_event_stream(b"<System />")
|
||||
]
|
||||
events = []
|
||||
for line in simulation_event_stream(b"<System />"):
|
||||
event = json.loads(line)
|
||||
events.append(event)
|
||||
if event.get("heartbeat"):
|
||||
heartbeat_observed.set()
|
||||
|
||||
heartbeats = [event for event in events if event.get("heartbeat") is True]
|
||||
self.assertGreaterEqual(len(heartbeats), 1)
|
||||
|
||||
@@ -0,0 +1,748 @@
|
||||
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 tests.ir_reference import reference_ir
|
||||
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 _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.program = reference_ir('mechanical')
|
||||
|
||||
def test_reference_schema_is_statically_valid_and_callback_free(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 = reference_ir('mechanical')
|
||||
|
||||
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_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.program = reference_ir('target')
|
||||
|
||||
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_reference_archive_does_not_claim_native_execution(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_reference_reload_is_byte_stable(self) -> None:
|
||||
second = reference_ir('target')
|
||||
|
||||
self.assertEqual(
|
||||
self.program.canonical_json_bytes(),
|
||||
second.canonical_json_bytes(),
|
||||
)
|
||||
|
||||
|
||||
class HistoricalSystemNumericIRV2Tests(unittest.TestCase):
|
||||
def test_historical_complex_model_schema_remains_valid(self) -> None:
|
||||
program = reference_ir('historical')
|
||||
|
||||
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()
|
||||
|
||||
class RetiredIRExporterTests(unittest.TestCase):
|
||||
def test_legacy_exporter_fails_with_a_migration_message(self):
|
||||
with self.assertRaisesRegex(NotImplementedError, 'retired'):
|
||||
compile_system_ir(object())
|
||||
Reference in new issue
Block a user