import abc
import itertools
import re
from collections.abc import Iterable, Sequence
from difflib import get_close_matches
from typing import Any, NotRequired, TypedDict
import networkx as nx
from qcodes.instrument import Instrument
from qcodes.parameters import Parameter
from semi_cr.core.lab.station.nodes import MultiSourceForwardingNode
from semi_cr.core.lab.station.graphing.models import Node, Edge
from semi_cr.core.lab.station.graphing.ids import (
connector_connection_node_id,
connector_node_id,
terminal_node_id,
)
from semi_cr.core.lab.station.graphing.operations import build_nx_graph
from semi_cr.core.lab.station.graphing.parsing import parse_endpoint_ref
from ..station.routing.base import Routable
[docs]
class Quellable(abc.ABC):
[docs]
@abc.abstractmethod
def quell(self) -> None:
pass
[docs]
class ConnectorRouteNode(MultiSourceForwardingNode):
def __init__(self, parameter: Parameter, node: Node):
super().__init__(name=node.name)
self.node = node
self._resistance = parameter
@property
def parameters(self) -> Iterable[Parameter]:
if not self._sources:
return ()
return itertools.chain(
*(source.parameters for source in self._sources),
[self._resistance],
)
[docs]
def substitute_non_identifier_characters(
node_name: str,
valid_character: str = "_",
) -> str:
return re.sub(r"[^0-9a-zA-Z_]", valid_character, node_name)
[docs]
class Connections(TypedDict):
name: NotRequired[str]
function: NotRequired[str]
endpoints: tuple[str, str]
[docs]
class Connector(Instrument, Routable):
def __init__(self,
name: str,
connections: Sequence[Connections],
**kwargs: Any):
super().__init__(name, **kwargs)
self._are_connection_names_unique(connections)
self._connections = connections
[docs]
def get_idn(self) -> dict[str, str | None]:
"""
Override Instrument.get_idn so QCoDeS never calls self.ask('*IDN?').
"""
return {"vendor": None,
"model": self.name,
"serial": None,
"firmware": None}
@staticmethod
def _are_connection_names_unique(connections: Sequence[Connections]) -> None:
names = {
connection.get("name", str(index))
for index, connection in enumerate(connections)
}
if len(names) != len(connections):
raise KeyError("...")
@staticmethod
def _check_similar_key(name: str, dictionary: Connections) -> None:
key_list = get_close_matches(name, dictionary.keys())
if len(key_list) != 0:
raise ValueError(f"{key_list[0]} key is not defined correctly...")
def _make_connector_node(
self,
name: str,
dictionary: Connections,
) -> ConnectorRouteNode:
value = dictionary.get("ohms", 0)
if "ohms" not in dictionary:
self._check_similar_key("ohms", dictionary)
parameter_name = substitute_non_identifier_characters(name)
self.add_parameter(
parameter_name,
initial_cache_value=value,
unit="ohms",
set_cmd=False,
get_cmd=None,
)
node = Node(
id=connector_connection_node_id(self.name, name),
kind="connector_connection",
name=name,
obj=None,
)
return ConnectorRouteNode(
parameter=getattr(self, parameter_name),
node=node,
)
[docs]
def quell(self) -> None:
pass
@property
def graph(self) -> nx.MultiDiGraph:
nodes: list[Node] = []
edges: list[Edge] = []
connector_id = connector_node_id(self.name)
nodes.append(
Node(
id=connector_id,
kind="connector",
name=self.name,
obj=self,
)
)
for index, connection in enumerate(self._connections):
connection_name = connection.get("name", str(index))
# connection_id = connector_connection_node_id(
# self.name,
# connection_name,
# )
terminal_id = terminal_node_id(self.name, connection_name)
# nodes.append(
# Node(
# id=connection_id,
# kind="connector_connection",
# name=connection_name,
# obj=None,
# )
# )
nodes.append(Node(
id=terminal_id,
kind="terminal",
name=f"{self.name}[{connection_name}]",
obj=None,
))
# edges.append(
# Edge(
# source=connection_id,
# target=terminal_id,
# kind="represents",
# )
# )
for endpoint in connection.get("endpoints", []):
endpoint_owner, endpoint_name = parse_endpoint_ref(endpoint)
endpoint_terminal_id = terminal_node_id(
endpoint_owner,
endpoint_name,
)
edges.append(
Edge(
source=endpoint_terminal_id,
target=terminal_id,
kind="dependency",
)
)
graph = build_nx_graph(nodes, edges)
return graph