from __future__ import annotations from collections.abc import Mapping import unittest from app.simulation.core.base import AlgebraicComponent, DynamicComponent from app.simulation.core.ports import PortDefinition from app.simulation.solvers.stream import StreamResolver from app.simulation.systems.network import SimulationNetwork class _CountingDynamicAnchor(DynamicComponent): PORTS = (PortDefinition.pneumatic("port"),) def __init__( self, name: str, *, enthalpy: float, temperature_reference_h: float, ) -> None: super().__init__(name) self.enthalpy = enthalpy self.temperature_reference_h = temperature_reference_h self.refresh_count = 0 self.port = self.register_declared_port("port") self.port.h_outflow = -1.0 def make_ports_current(self) -> None: self.port.h_outflow = self.enthalpy def get_state_vector(self) -> list[float]: return [0.0, 0.0] def set_state_vector(self, values: list[float]) -> None: if len(values) != self.state_size: raise ValueError("Unexpected test state size.") def refresh_thermodynamic_ports(self) -> None: self.refresh_count += 1 self.make_ports_current() def state_derivative_from_ports( self, connected_h: Mapping[str, float], ) -> list[float]: return [0.0, 0.0] class _PassThrough(AlgebraicComponent): PORTS = ( PortDefinition.pneumatic("left"), PortDefinition.pneumatic("right"), ) def __init__(self, name: str, update_log: list[str]) -> None: super().__init__(name) self.left = self.register_declared_port("left") self.right = self.register_declared_port("right") self.update_log = update_log def update_stream_outflows(self, connected_h: Mapping[str, float]) -> None: self.update_log.append(self.name) self.left.h_outflow = connected_h["right"] self.right.h_outflow = connected_h["left"] class _ReferenceAwarePassThrough(_PassThrough): def __init__( self, name: str, update_log: list[str], reference_log: list[str], ) -> None: super().__init__(name, update_log) self.reference_log = reference_log self.flow_temperature_references: dict[str, float] = {} def update_flow_temperature_references( self, connected_h: Mapping[str, float], ) -> None: self.reference_log.append(self.name) self.flow_temperature_references = dict(connected_h) def _build_chain() -> tuple[ SimulationNetwork, _CountingDynamicAnchor, _PassThrough, _PassThrough, _CountingDynamicAnchor, list[str], ]: update_log: list[str] = [] left = _CountingDynamicAnchor( "left_anchor", enthalpy=100.0, temperature_reference_h=1_100.0, ) first = _PassThrough("first", update_log) second = _PassThrough("second", update_log) right = _CountingDynamicAnchor( "right_anchor", enthalpy=400.0, temperature_reference_h=1_400.0, ) network = SimulationNetwork("stream-chain") for component in (left, first, second, right): network.add_component(component) network.connect("left_anchor", "port", "first", "left") network.connect("first", "right", "second", "left") network.connect("second", "right", "right_anchor", "port") return network, left, first, second, right, update_log class StreamResolverExecutionPlanTests(unittest.TestCase): def test_standalone_solve_refreshes_each_dynamic_exactly_once(self) -> None: network, left, _first, _second, right, _update_log = _build_chain() diagnostics, _connected = StreamResolver(network).solve() self.assertTrue(diagnostics.converged) self.assertEqual(left.refresh_count, 1) self.assertEqual(right.refresh_count, 1) def test_current_dynamic_ports_skip_refresh(self) -> None: network, left, _first, _second, right, _update_log = _build_chain() left.make_ports_current() right.make_ports_current() diagnostics, connected = StreamResolver(network).solve( dynamic_ports_are_current=True ) self.assertTrue(diagnostics.converged) self.assertEqual(left.refresh_count, 0) self.assertEqual(right.refresh_count, 0) self.assertEqual(connected["first"], {"left": 100.0, "right": 400.0}) def test_multiple_iterations_do_not_repeat_dynamic_refresh(self) -> None: network, left, _first, _second, right, update_log = _build_chain() diagnostics, _connected = StreamResolver(network).solve() self.assertGreater(diagnostics.iterations, 1) self.assertEqual(left.refresh_count, 1) self.assertEqual(right.refresh_count, 1) self.assertEqual( update_log, ["first", "second"] * diagnostics.iterations, ) def test_non_dynamic_flow_reference_hook_does_not_repeat_stream_update( self, ) -> None: update_log: list[str] = [] reference_log: list[str] = [] left = _CountingDynamicAnchor( "left_anchor", enthalpy=100.0, temperature_reference_h=1_100.0, ) middle = _ReferenceAwarePassThrough( "middle", update_log, reference_log, ) right = _CountingDynamicAnchor( "right_anchor", enthalpy=400.0, temperature_reference_h=1_400.0, ) network = SimulationNetwork("flow-temperature-reference") for component in (left, middle, right): network.add_component(component) network.connect("left_anchor", "port", "middle", "left") network.connect("middle", "right", "right_anchor", "port") resolver = StreamResolver(network) diagnostics, _connected = resolver.solve() stream_updates_before = tuple(update_log) outflows_before = (middle.left.h_outflow, middle.right.h_outflow) resolver.refresh_flow_temperature_references() self.assertEqual( stream_updates_before, ("middle",) * diagnostics.iterations, ) self.assertEqual(tuple(update_log), stream_updates_before) self.assertEqual(reference_log, ["middle"]) self.assertEqual( middle.flow_temperature_references, {"left": 1_100.0, "right": 1_400.0}, ) self.assertEqual( (middle.left.h_outflow, middle.right.h_outflow), outflows_before, ) def test_precompiled_bindings_preserve_outputs_and_references(self) -> None: default_network, default_left, default_first, default_second, default_right, _ = ( _build_chain() ) current_network, current_left, current_first, current_second, current_right, _ = ( _build_chain() ) current_left.make_ports_current() current_right.make_ports_current() default_resolver = StreamResolver(default_network) current_resolver = StreamResolver(current_network) default_diagnostics, default_connected = default_resolver.solve() current_diagnostics, current_connected = current_resolver.solve( dynamic_ports_are_current=True ) self.assertEqual(default_diagnostics, current_diagnostics) self.assertEqual(default_connected, current_connected) self.assertEqual( ( default_left.port.h_outflow, default_first.left.h_outflow, default_first.right.h_outflow, default_second.left.h_outflow, default_second.right.h_outflow, default_right.port.h_outflow, ), ( current_left.port.h_outflow, current_first.left.h_outflow, current_first.right.h_outflow, current_second.left.h_outflow, current_second.right.h_outflow, current_right.port.h_outflow, ), ) references = default_resolver.connected_temperature_reference_enthalpies() self.assertEqual(references["first"]["left"], 1_100.0) self.assertEqual(references["second"]["right"], 1_400.0) if __name__ == "__main__": unittest.main()