from __future__ import annotations from collections.abc import Callable, Sequence from copy import copy from dataclasses import dataclass, replace from app.simulation.core.errors import RecoverableTrialStateError from app.simulation.core.ports import PortState _STREAM_CACHE_ATTRIBUTE_NAMES = frozenset( { "_connected_h", "temperature_reference_h", } ) def _is_stream_cache_attribute(name: str) -> bool: """Return whether an attribute belongs to the stream/temperature replay state. Catalog components currently use ``_connected_h`` and ``temperature_reference_h``. The name-based extension keeps conservative third-party caches recoverable without copying an entire component graph. Components with opaque cache names can provide the explicit hooks documented by :class:`ThermofluidTransactionPlan`. """ lowered = name.lower() return ( name in _STREAM_CACHE_ATTRIBUTE_NAMES or lowered.startswith("_stream_") or "connected_h" in lowered or "connected_enthalpy" in lowered or "temperature_reference" in lowered ) def _copy_cache_value(value: object) -> object: """Shallow-copy a stream cache without traversing the component graph.""" if isinstance(value, (dict, list, set, bytearray)): return copy(value) return value @dataclass(frozen=True) class ThermofluidWorstPort: component: str port: str value: float signed_delta: float def as_dict(self) -> dict[str, object]: return { "component": self.component, "port": self.port, "value": self.value, "signedDelta": self.signed_delta, } @dataclass(frozen=True) class ThermofluidIterationDelta: iteration: int max_delta: float scale: float tolerance: float worst_port: ThermofluidWorstPort | None def as_dict(self) -> dict[str, object]: return { "iteration": self.iteration, "maxDelta": self.max_delta, "scale": self.scale, "tolerance": self.tolerance, "worstPort": ( self.worst_port.as_dict() if self.worst_port is not None else None ), } @dataclass(frozen=True) class ThermofluidClosureSuccess: rhs_time: float iterations: int max_delta: float scale: float tolerance: float worst_port: ThermofluidWorstPort | None @classmethod def from_iteration( cls, rhs_time: float, delta: ThermofluidIterationDelta, ) -> ThermofluidClosureSuccess: return cls( rhs_time=float(rhs_time), iterations=delta.iteration, max_delta=delta.max_delta, scale=delta.scale, tolerance=delta.tolerance, worst_port=delta.worst_port, ) def as_dict(self) -> dict[str, object]: return { "rhsTime": self.rhs_time, "iterations": self.iterations, "maxDelta": self.max_delta, "scale": self.scale, "tolerance": self.tolerance, "worstPort": ( self.worst_port.as_dict() if self.worst_port is not None else None ), } @dataclass(frozen=True) class ThermofluidClosureFailure: failed_rhs_time: float iterations: int delta_tail: tuple[ThermofluidIterationDelta, ...] max_delta: float scale: float tolerance: float worst_port: ThermofluidWorstPort | None failure_count: int = 0 @classmethod def from_iterations( cls, failed_rhs_time: float, deltas: Sequence[ThermofluidIterationDelta], *, tail_limit: int = 8, ) -> ThermofluidClosureFailure: if not deltas: raise ValueError("A thermofluid failure requires iteration diagnostics.") final = deltas[-1] return cls( failed_rhs_time=float(failed_rhs_time), iterations=final.iteration, delta_tail=tuple(deltas[-tail_limit:]), max_delta=final.max_delta, scale=final.scale, tolerance=final.tolerance, worst_port=final.worst_port, ) def as_dict(self) -> dict[str, object]: return { "failedRhsTime": self.failed_rhs_time, "iterations": self.iterations, "deltaTail": [item.as_dict() for item in self.delta_tail], "maxDelta": self.max_delta, "scale": self.scale, "tolerance": self.tolerance, "worstPort": ( self.worst_port.as_dict() if self.worst_port is not None else None ), "failureCount": self.failure_count, } class ThermofluidClosureError(RecoverableTrialStateError): """Recoverable exhaustion of the stream/pressure-flow fixed point. Stream propagation failures and algebraic-solver failures intentionally retain their original exception types: rollback is still applied, but a smaller ODE step is not known to repair those structural/numerical errors. """ def __init__(self, diagnostics: ThermofluidClosureFailure) -> None: super().__init__( "Stream enthalpy and pressure-flow coupling did not converge " f"after {diagnostics.iterations} iterations at " f"t={diagnostics.failed_rhs_time:.17g}." ) self.diagnostics = diagnostics class ThermofluidClosureDiagnostics: """Run-level RHS outcomes; maintenance/postprocessing calls do not write it.""" def __init__(self) -> None: self.failure_count = 0 self.last_failure: ThermofluidClosureFailure | None = None self.last_success: ThermofluidClosureSuccess | None = None def record_success(self, success: ThermofluidClosureSuccess) -> None: self.last_success = success def record_failure( self, failure: ThermofluidClosureFailure, ) -> ThermofluidClosureFailure: self.failure_count += 1 recorded = replace(failure, failure_count=self.failure_count) self.last_failure = recorded return recorded def as_dict(self) -> dict[str, object]: return { "failureCount": self.failure_count, "lastFailure": ( self.last_failure.as_dict() if self.last_failure is not None else None ), "lastSuccess": ( self.last_success.as_dict() if self.last_success is not None else None ), } @dataclass(frozen=True) class _PortValueBinding: component_name: str port_name: str state: PortState variable: str @dataclass(frozen=True) class _PortFieldPlan: variable: str states: tuple[PortState, ...] @dataclass(frozen=True) class _FlowBinding: component_name: str port_name: str state: PortState @dataclass(frozen=True) class _ComponentCacheBinding: component: object attribute_names: tuple[str, ...] attribute_name_set: frozenset[str] snapshot_hook: Callable[[], object] | None restore_hook: Callable[[object], None] | None @dataclass class ThermofluidTransactionSnapshot: plan: ThermofluidTransactionPlan port_values: tuple[list[float], ...] component_cache_values: tuple[list[object], ...] custom_cache_values: list[object | None] diagnostic_values: list[object] def restore(self) -> None: plan = self.plan plan._restore_port_values(self.port_values) for binding, values, custom_value in zip( plan.component_cache_bindings, self.component_cache_values, self.custom_cache_values, ): component = binding.component for name in tuple(getattr(component, "__dict__", {})): if ( name.startswith("_causal_") or _is_stream_cache_attribute(name) ) and name not in binding.attribute_name_set: delattr(component, name) for name, value in zip(binding.attribute_names, values): setattr(component, name, _copy_cache_value(value)) if binding.restore_hook is not None: binding.restore_hook(custom_value) for owner, value in zip( plan.diagnostic_owners, self.diagnostic_values, ): owner.last_diagnostics = value class ThermofluidTransactionPlan: """Compiled, lightweight rollback boundary for one Generic RHS closure. It snapshots active physical-port values, catalog stream-temperature caches, component ``_causal_*`` seed fields, and resolver/solver last diagnostics. A custom stream-aware component with an opaque mutable cache can implement both ``snapshot_thermofluid_closure_cache()`` and ``restore_thermofluid_closure_cache(snapshot)``; these hooks are invoked in addition to the standard name-based cache capture. """ def __init__( self, *, port_value_bindings: tuple[_PortValueBinding, ...], port_field_plans: tuple[_PortFieldPlan, ...], flow_bindings: tuple[_FlowBinding, ...], component_cache_bindings: tuple[_ComponentCacheBinding, ...], component_count: int, diagnostic_owners: tuple[object, ...], ) -> None: self.port_value_bindings = port_value_bindings self.port_field_plans = port_field_plans self.flow_bindings = flow_bindings self.component_cache_bindings = component_cache_bindings self.component_count = component_count self.diagnostic_owners = diagnostic_owners self._snapshot = ThermofluidTransactionSnapshot( plan=self, port_values=tuple( [0.0] * len(field.states) for field in port_field_plans ), component_cache_values=tuple( [None] * len(binding.attribute_names) for binding in component_cache_bindings ), custom_cache_values=[None] * len(component_cache_bindings), diagnostic_values=[None] * len(diagnostic_owners), ) @classmethod def compile( cls, network: object, *, diagnostic_owners: Sequence[object] = (), ) -> ThermofluidTransactionPlan: components = tuple(getattr(network, "components").values()) port_value_bindings: list[_PortValueBinding] = [] port_states_by_variable: dict[str, list[PortState]] = {} flow_bindings: list[_FlowBinding] = [] component_cache_bindings: list[_ComponentCacheBinding] = [] for component in components: active_definitions = tuple( definition for definition in component.active_port_definitions if definition.kind == "physical" ) for definition in active_definitions: state = component.get_port(definition.name) flow_bindings.append( _FlowBinding(component.name, definition.name, state) ) for variable in definition.variables: port_states_by_variable.setdefault(variable.name, []).append(state) port_value_bindings.append( _PortValueBinding( component.name, definition.name, state, variable.name, ) ) attribute_names = tuple( name for name in getattr(component, "__dict__", {}) if name.startswith("_causal_") or _is_stream_cache_attribute(name) ) snapshot_hook = getattr( component, "snapshot_thermofluid_closure_cache", None, ) restore_hook = getattr( component, "restore_thermofluid_closure_cache", None, ) hooks_are_available = callable(snapshot_hook) and callable(restore_hook) if attribute_names or hooks_are_available: component_cache_bindings.append( _ComponentCacheBinding( component=component, attribute_names=attribute_names, attribute_name_set=frozenset(attribute_names), snapshot_hook=(snapshot_hook if hooks_are_available else None), restore_hook=(restore_hook if hooks_are_available else None), ) ) owners = tuple( dict.fromkeys( owner for owner in diagnostic_owners if hasattr(owner, "last_diagnostics") ) ) return cls( port_value_bindings=tuple(port_value_bindings), port_field_plans=tuple( _PortFieldPlan(variable, tuple(states)) for variable, states in port_states_by_variable.items() ), flow_bindings=tuple(flow_bindings), component_cache_bindings=tuple(component_cache_bindings), component_count=len(components), diagnostic_owners=owners, ) def capture(self) -> ThermofluidTransactionSnapshot: # GenericFluidSystem executes one RHS serially. Reuse one compiled # workspace rather than allocating a snapshot object and several outer # tuples at every successful trial point. snapshot = self._snapshot self._capture_port_values(snapshot.port_values) for binding, values in zip( self.component_cache_bindings, snapshot.component_cache_values, ): for position, name in enumerate(binding.attribute_names): values[position] = _copy_cache_value( getattr(binding.component, name) ) for position, binding in enumerate(self.component_cache_bindings): snapshot.custom_cache_values[position] = ( binding.snapshot_hook() if binding.snapshot_hook is not None else None ) for position, owner in enumerate(self.diagnostic_owners): snapshot.diagnostic_values[position] = owner.last_diagnostics return snapshot def _capture_port_values( self, workspaces: tuple[list[float], ...], ) -> None: for field, values in zip(self.port_field_plans, workspaces): variable = field.variable states = field.states if variable == "p": for position, state in enumerate(states): values[position] = state.p elif variable == "m_flow": for position, state in enumerate(states): values[position] = state.m_flow elif variable == "h_outflow": for position, state in enumerate(states): values[position] = state.h_outflow elif variable == "volume": for position, state in enumerate(states): values[position] = state.volume elif variable == "volume_flow": for position, state in enumerate(states): values[position] = state.volume_flow elif variable == "x": for position, state in enumerate(states): values[position] = state.x elif variable == "v": for position, state in enumerate(states): values[position] = state.v elif variable == "f": for position, state in enumerate(states): values[position] = state.f else: for position, state in enumerate(states): values[position] = getattr(state, variable) def _restore_port_values( self, workspaces: tuple[list[float], ...], ) -> None: for field, values in zip(self.port_field_plans, workspaces): variable = field.variable states = field.states if variable == "p": for state, value in zip(states, values): state.p = value elif variable == "m_flow": for state, value in zip(states, values): state.m_flow = value elif variable == "h_outflow": for state, value in zip(states, values): state.h_outflow = value elif variable == "volume": for state, value in zip(states, values): state.volume = value elif variable == "volume_flow": for state, value in zip(states, values): state.volume_flow = value elif variable == "x": for state, value in zip(states, values): state.x = value elif variable == "v": for state, value in zip(states, values): state.v = value elif variable == "f": for state, value in zip(states, values): state.f = value else: for state, value in zip(states, values): setattr(state, variable, value) def flow_values(self) -> tuple[float, ...]: return tuple(float(binding.state.m_flow) for binding in self.flow_bindings) def measure_flow_delta( self, previous: Sequence[float], *, iteration: int, relative_tolerance: float, ) -> ThermofluidIterationDelta: current = self.flow_values() scale = max( (abs(value) for value in (*previous, *current)), default=1.0, ) scale = max(scale, 1.0) worst_index = -1 worst_signed_delta = 0.0 max_delta = 0.0 for index, (old, new) in enumerate(zip(previous, current)): signed_delta = new - old magnitude = abs(signed_delta) if magnitude > max_delta: worst_index = index worst_signed_delta = signed_delta max_delta = magnitude worst_port = None if worst_index >= 0: binding = self.flow_bindings[worst_index] worst_port = ThermofluidWorstPort( component=binding.component_name, port=binding.port_name, value=current[worst_index], signed_delta=worst_signed_delta, ) return ThermofluidIterationDelta( iteration=int(iteration), max_delta=max_delta, scale=scale, tolerance=float(relative_tolerance) * scale, worst_port=worst_port, ) def diagnostics(self) -> dict[str, int]: stream_cache_slot_count = sum( len(binding.attribute_names) for binding in self.component_cache_bindings ) return { "physicalPortValueSlotCount": len(self.port_value_bindings), "physicalFlowPortCount": len(self.flow_bindings), "componentCount": self.component_count, "cacheBindingCount": len(self.component_cache_bindings), "streamAndCausalCacheSlotCount": stream_cache_slot_count, "customCacheHookCount": sum( binding.snapshot_hook is not None for binding in self.component_cache_bindings ), "diagnosticOwnerCount": len(self.diagnostic_owners), }