from typing import Any
from collections.abc import Mapping
import networkx as nx
from semi_cr.core.lab.station.graphing.ids import (
qualified_device_name,
resolve_terminal_id,
instrument_node_id,
terminal_id,
pin_node_id,
)
from semi_cr.core.lab.station.graphing.parsing import (
normalize_channel,
parse_endpoint,
)
from semi_cr.core.lab.station.graphing.operations import (
add_connection,
ensure_instrument_terminal,
)
[docs]
def build_device_registry(instruments):
device_registry: dict[
str,
tuple[str, str],
] = {}
for chip_name, chip_config in instruments.items():
devices = (
chip_config
.get("init", {})
.get("devices", {})
)
for local_device_name in devices:
compact_name = qualified_device_name(
chip_name=chip_name,
device_name=str(local_device_name),
)
device_registry[compact_name] = (
chip_name,
str(local_device_name),
)
return device_registry
[docs]
def add_static_instruments(
graph: nx.MultiDiGraph,
instruments: Mapping[str, Mapping[str, Any]],
) -> None:
for instrument_name, instrument_config in instruments.items():
init = instrument_config.get("init", {})
declared_name = init.get("name")
if (
declared_name is not None
and declared_name != instrument_name
):
raise ValueError(
f"Instrument {instrument_name!r} declares runtime name "
f"{declared_name!r}. Config and runtime names must match."
)
graph.add_node(
instrument_node_id(instrument_name),
kind="instrument",
name=instrument_name,
type=instrument_config.get("type"),
)
[docs]
def ensure_static_terminal(
graph: nx.MultiDiGraph,
instrument_name: str,
terminal_name: str,
) -> str:
normalized_name = normalize_channel(
terminal_name
)
node_id = terminal_id(
instrument_name,
normalized_name,
)
graph.add_node(
node_id,
kind="terminal",
instrument_name=instrument_name,
terminal_name=normalized_name,
)
ensure_instrument_terminal(
graph,
instrument_name=instrument_name,
terminal_name=normalized_name,
node_id=node_id,
)
return node_id
[docs]
def add_instrument_wiring(
graph: nx.MultiDiGraph,
instruments: Mapping[str, Mapping[str, Any]],
) -> None:
for instrument_name, instrument_config in instruments.items():
connections = (
instrument_config
.get("init", {})
.get("connections", [])
)
for connection in connections:
local_name = connection["name"]
local_node = ensure_static_terminal(
graph,
instrument_name=instrument_name,
terminal_name=local_name,
)
for endpoint in connection.get("endpoints", []):
(
endpoint_instrument,
endpoint_channel,
) = parse_endpoint(endpoint)
endpoint_node = ensure_static_terminal(
graph,
instrument_name=endpoint_instrument,
terminal_name=endpoint_channel,
)
add_connection(
graph,
endpoint_node,
local_node,
kind="physical",
source_instrument=endpoint_instrument,
source_channel=normalize_channel(
endpoint_channel
),
target_instrument=instrument_name,
target_channel=normalize_channel(
local_name
),
)
[docs]
def add_device_pad_wiring(
graph: nx.MultiDiGraph,
instruments: Mapping[str, Mapping[str, Any]],
device_registry: Mapping[str, tuple[str, str]],
) -> None:
for chip_name, chip_config in instruments.items():
devices = (
chip_config
.get("init", {})
.get("devices", {})
)
for local_device_name, device_config in devices.items():
local_device_name = str(local_device_name)
canonical_device_name = qualified_device_name(
chip_name=chip_name,
device_name=local_device_name,
)
pins = device_config.get("pins", {})
for pin_name, pin_config in pins.items():
pin_name = str(pin_name)
pin_id = pin_node_id(
chip_name=chip_name,
device_name=local_device_name,
pin_name=pin_name,
)
graph.add_node(
pin_id,
kind="device_pin",
name=pin_name,
chip_name=chip_name,
device_name=local_device_name,
qualified_device_name=canonical_device_name,
)
for pad_name in pin_config.get("pads", []):
(
pad_instrument,
pad_channel,
) = parse_endpoint(pad_name)
external_terminal_id = resolve_terminal_id(
instrument_name=pad_instrument,
channel=pad_channel,
device_registry=device_registry,
)
add_connection(
graph,
external_terminal_id,
pin_id,
kind="device",
source_instrument=pad_instrument,
source_channel=normalize_channel(pad_channel),
target_instrument=canonical_device_name,
target_channel=pin_name,
pad_name=pad_name,
)
[docs]
def build_wiring_graph(
config: Mapping[str, Any],
) -> nx.MultiDiGraph:
"""
Build the static station topology.
This graph contains configuration-level instruments, physical
terminals, device pads, and device terminals. It contains no
live QCoDeS objects.
"""
graph = nx.MultiDiGraph()
instruments: Mapping[str, Mapping[str, Any]] = (
config.get("instruments", {})
)
# ---------------------------------------------------------
# 1. Instrument nodes
# ---------------------------------------------------------
add_static_instruments(
graph,
instruments,
)
# ---------------------------------------------------------
# 2. Device registry
# ---------------------------------------------------------
device_registry = build_device_registry(
instruments
)
# ---------------------------------------------------------
# 3. Physical instrument wiring
# ---------------------------------------------------------
add_instrument_wiring(
graph,
instruments,
)
# ---------------------------------------------------------
# 4. Device pin / pad wiring
# ---------------------------------------------------------
add_device_pad_wiring(
graph,
instruments,
device_registry=device_registry,
)
return graph