79 lines
2.5 KiB
Python
79 lines
2.5 KiB
Python
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass
|
|
|
|
from PythonModels.core.base import Component, DynamicComponent
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class Connection:
|
|
source_component: str
|
|
source_port: str
|
|
target_component: str
|
|
target_port: str
|
|
|
|
|
|
class SimulationNetwork:
|
|
"""Container for components, topology, and state-vector bookkeeping."""
|
|
|
|
def __init__(self, name: str) -> None:
|
|
self.name = name
|
|
self.components: dict[str, Component] = {}
|
|
self.connections: list[Connection] = []
|
|
|
|
def add_component(self, component: Component) -> None:
|
|
if component.name in self.components:
|
|
raise ValueError(f"Duplicate component name: {component.name}")
|
|
self.components[component.name] = component
|
|
|
|
def connect(
|
|
self,
|
|
source_component: str,
|
|
source_port: str,
|
|
target_component: str,
|
|
target_port: str,
|
|
) -> None:
|
|
self.connections.append(
|
|
Connection(
|
|
source_component=source_component,
|
|
source_port=source_port,
|
|
target_component=target_component,
|
|
target_port=target_port,
|
|
)
|
|
)
|
|
|
|
def dynamic_components(self) -> list[DynamicComponent]:
|
|
return [
|
|
component
|
|
for component in self.components.values()
|
|
if isinstance(component, DynamicComponent)
|
|
]
|
|
|
|
def initial_state_vector(self) -> list[float]:
|
|
values: list[float] = []
|
|
for component in self.dynamic_components():
|
|
values.extend(component.get_state_vector())
|
|
return values
|
|
|
|
def apply_state_vector(self, values: list[float]) -> None:
|
|
cursor = 0
|
|
for component in self.dynamic_components():
|
|
next_cursor = cursor + component.state_size
|
|
component.set_state_vector(values[cursor:next_cursor])
|
|
cursor = next_cursor
|
|
if cursor != len(values):
|
|
raise ValueError("State vector length does not match dynamic components.")
|
|
|
|
def summary(self) -> str:
|
|
lines = [f"Network: {self.name}", "Components:"]
|
|
for name, component in self.components.items():
|
|
lines.append(f" - {name}: {component.__class__.__name__}")
|
|
lines.append("Connections:")
|
|
for conn in self.connections:
|
|
lines.append(
|
|
f" - {conn.source_component}.{conn.source_port}"
|
|
f" -> {conn.target_component}.{conn.target_port}"
|
|
)
|
|
return "\n".join(lines)
|
|
|