146 lines
5.1 KiB
Python
146 lines
5.1 KiB
Python
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass
|
|
from math import isfinite
|
|
from typing import Callable, Protocol
|
|
|
|
from app.simulation.core.base import Component
|
|
from app.simulation.core.ports import PortState
|
|
from app.simulation.performance import profile_phase
|
|
from app.simulation.systems.network import Endpoint, SimulationNetwork
|
|
|
|
|
|
class SignalOutputComponent(Protocol):
|
|
name: str
|
|
|
|
def signal_output_values(self, time: float) -> dict[str, float]:
|
|
...
|
|
|
|
|
|
class SignalEventSource(Protocol):
|
|
"""Optional contract for signal sources with known time discontinuities."""
|
|
|
|
name: str
|
|
|
|
def signal_event_times(
|
|
self,
|
|
start_time: float,
|
|
stop_time: float,
|
|
) -> tuple[float, ...]:
|
|
"""Return event times strictly inside ``(start_time, stop_time)``."""
|
|
|
|
...
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class SignalSolveDiagnostics:
|
|
propagated: int
|
|
|
|
def as_dict(self) -> dict[str, object]:
|
|
return {"propagated": self.propagated}
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class _SignalOutputBinding:
|
|
component: Component
|
|
evaluate: Callable[[float], dict[str, float]]
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class _SignalConnectionBinding:
|
|
source: PortState
|
|
target: PortState
|
|
|
|
|
|
class SignalResolver:
|
|
"""Propagate scalar signal connections from output ports to input ports."""
|
|
|
|
def __init__(self, network: SimulationNetwork) -> None:
|
|
self.network = network
|
|
self._output_bindings = tuple(
|
|
_SignalOutputBinding(component=component, evaluate=evaluate)
|
|
for component in network.components.values()
|
|
if (evaluate := getattr(component, "signal_output_values", None)) is not None
|
|
)
|
|
self._event_sources = tuple(
|
|
(component.name, source_event_times)
|
|
for component in network.components.values()
|
|
if (
|
|
source_event_times := getattr(
|
|
component,
|
|
"signal_event_times",
|
|
None,
|
|
)
|
|
)
|
|
is not None
|
|
)
|
|
self._connections = tuple(
|
|
self._connection_binding(connection.endpoints)
|
|
for connection in network.connections
|
|
if connection.kind == "signal"
|
|
)
|
|
self.last_diagnostics: SignalSolveDiagnostics | None = None
|
|
|
|
@profile_phase("simulation.signal", minimum_mode="audit")
|
|
def solve(self, time: float) -> SignalSolveDiagnostics:
|
|
for binding in self._output_bindings:
|
|
for port_name, value in binding.evaluate(time).items():
|
|
binding.component.get_port(port_name).signal = float(value)
|
|
|
|
propagated = 0
|
|
for binding in self._connections:
|
|
binding.target.signal = binding.source.signal
|
|
propagated += 1
|
|
|
|
diagnostics = SignalSolveDiagnostics(propagated=propagated)
|
|
self.last_diagnostics = diagnostics
|
|
return diagnostics
|
|
|
|
def event_times(self, start_time: float, stop_time: float) -> tuple[float, ...]:
|
|
"""Collect optional source events that can be used as integration splits.
|
|
|
|
Event discovery is deliberately duck typed so existing signal-output
|
|
components remain valid without implementing ``signal_event_times``.
|
|
"""
|
|
|
|
start = float(start_time)
|
|
stop = float(stop_time)
|
|
if not isfinite(start) or not isfinite(stop):
|
|
raise ValueError("Signal event interval must be finite.")
|
|
if stop < start:
|
|
raise ValueError("Signal event interval stop must not precede start.")
|
|
if stop == start:
|
|
return ()
|
|
|
|
events: set[float] = set()
|
|
for component_name, source_event_times in self._event_sources:
|
|
for raw_time in source_event_times(start, stop):
|
|
event_time = float(raw_time)
|
|
if not isfinite(event_time):
|
|
raise ValueError(
|
|
f"Signal event time from component '{component_name}' must be finite."
|
|
)
|
|
if start < event_time < stop:
|
|
events.add(event_time)
|
|
return tuple(sorted(events))
|
|
|
|
def _source_target(self, endpoints: tuple[Endpoint, Endpoint]) -> tuple[Endpoint, Endpoint]:
|
|
first, second = endpoints
|
|
first_port = self.network.components[first.component].get_port(first.port)
|
|
second_port = self.network.components[second.component].get_port(second.port)
|
|
if first_port.definition is not None and first_port.definition.nominal_role == "output":
|
|
return first, second
|
|
if second_port.definition is not None and second_port.definition.nominal_role == "output":
|
|
return second, first
|
|
raise ValueError("Signal connection must contain one output endpoint.")
|
|
|
|
def _connection_binding(
|
|
self,
|
|
endpoints: tuple[Endpoint, Endpoint],
|
|
) -> _SignalConnectionBinding:
|
|
source, target = self._source_target(endpoints)
|
|
return _SignalConnectionBinding(
|
|
source=self.network.components[source.component].get_port(source.port),
|
|
target=self.network.components[target.component].get_port(target.port),
|
|
)
|