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, )