Source code for semi_cr.core.lab.station.context.graph
from dataclasses import dataclass
from typing import Any
import networkx as nx
from semi_cr.core.lab.station.graphing.models import GraphRelation
from semi_cr.core.lab.station.graphing.ids import qcodes_node_id
from semi_cr.core.lab.station.graphing.overlay import (
DependencyOverlayBuilder,
)
from semi_cr.core.lab.station.graphing.resolver import GraphNodeResolver
from semi_cr.core.lab.station.graphing.runtime_builder import (
ContainmentGraphBuilder,
)
from semi_cr.core.lab.station.routing_v1.models import (
ConnectionRoutable,
)
from semi_cr.core.lab.station.runtime import iter_routables
from .bindings import RuntimeBindingRegistry
[docs]
@dataclass
class StationGraphBuilder:
station: Any
[docs]
def build(
self,
static_graph: nx.MultiDiGraph,
devices: tuple[Any, ...],
instruments: dict[str, Any],
bindings: RuntimeBindingRegistry,
) -> nx.MultiDiGraph:
canonical_graph = nx.compose_all([
static_graph,
*(device.graph for device in devices),
])
if not instruments:
return canonical_graph
runtime_containment = ContainmentGraphBuilder(
node_id=qcodes_node_id,
attach_objects=True,
).build(self.station)
base_graph = nx.compose(
canonical_graph,
runtime_containment,
)
runtime_graphs = self._collect_runtime_graphs(
instruments
)
graph = DependencyOverlayBuilder(
resolver=GraphNodeResolver(
qcodes_id=qcodes_node_id
)
).overlay(
base_graph,
runtime_graphs,
)
self._add_bindings(graph, bindings)
return graph
@staticmethod
def _collect_runtime_graphs(
instruments: dict[str, Any],
) -> list[nx.MultiDiGraph]:
graphs: list[nx.MultiDiGraph] = []
seen: set[int] = set()
for instrument in instruments.values():
for candidate in (
instrument,
*iter_routables(instrument),
):
candidate_id = id(candidate)
if candidate_id in seen:
continue
seen.add(candidate_id)
if isinstance(candidate, ConnectionRoutable):
graphs.append(candidate.graph)
return graphs
@staticmethod
def _add_bindings(
graph: nx.MultiDiGraph,
bindings: RuntimeBindingRegistry,
) -> None:
for static_id, binding in bindings.items():
if static_id not in graph:
raise KeyError(
f"Static node {static_id!r} is missing "
"from the station graph."
)
if binding.runtime_id not in graph:
raise KeyError(
f"Runtime node {binding.runtime_id!r} is "
"missing from the station graph."
)
graph.add_edge(
static_id,
binding.runtime_id,
relation=GraphRelation.REPRESENTS,
)