"""DiagramDocument: the in-memory model of a diagram plus change signals. Uses ``QObject`` signals so the view can react to model changes, but does not require a running ``QApplication`` for construction or signal emission (direct connections work without an event loop), keeping it unit-test friendly. """ from __future__ import annotations import json from typing import Iterable, Iterator, Optional from PyQt6.QtCore import QObject, pyqtSignal from .catalogue import CatalogueRegistry, default_registry as default_catalogues from .edge import EdgeModel from .node import NodeModel from .relations import RelationRegistry, default_registry as default_relations class DiagramDocument(QObject): """Holds nodes and edges, assigns ids, and emits change notifications.""" node_added = pyqtSignal(str) node_removed = pyqtSignal(str) node_changed = pyqtSignal(str) edge_added = pyqtSignal(str) edge_removed = pyqtSignal(str) edge_changed = pyqtSignal(str) cleared = pyqtSignal() modified = pyqtSignal() def __init__( self, *, relations: Optional[RelationRegistry] = None, catalogues: Optional[CatalogueRegistry] = None, parent: Optional[QObject] = None, ) -> None: super().__init__(parent) self.relations = relations if relations is not None else default_relations() self.catalogues = catalogues if catalogues is not None else default_catalogues() self._nodes: dict[str, NodeModel] = {} self._edges: dict[str, EdgeModel] = {} self._id_counter = 0 self.dirty = False # -- id generation ---------------------------------------------------- def next_id(self, prefix: str) -> str: self._id_counter += 1 return f"{prefix}{self._id_counter}" # -- nodes ------------------------------------------------------------ def add_node(self, node: NodeModel) -> NodeModel: if node.node_id in self._nodes: raise ValueError(f"duplicate node id: {node.node_id}") self._nodes[node.node_id] = node self._mark_dirty() self.node_added.emit(node.node_id) return node def remove_node(self, node_id: str) -> list[str]: """Remove a node and any edges attached to it. Returns removed edge ids.""" if node_id not in self._nodes: return [] removed_edges = [e.edge_id for e in self._edges.values() if e.source_node == node_id or e.target_node == node_id] for eid in removed_edges: self.remove_edge(eid) del self._nodes[node_id] self._mark_dirty() self.node_removed.emit(node_id) return removed_edges def node(self, node_id: str) -> Optional[NodeModel]: return self._nodes.get(node_id) def nodes(self) -> Iterator[NodeModel]: return iter(self._nodes.values()) def notify_node_changed(self, node_id: str) -> None: if node_id in self._nodes: self._mark_dirty() self.node_changed.emit(node_id) # -- edges ------------------------------------------------------------ def can_connect(self, src_node: str, src_port: str, dst_node: str, dst_port: str) -> bool: """Whether two ports may be connected under the relation rules.""" sn, dn = self._nodes.get(src_node), self._nodes.get(dst_node) if sn is None or dn is None: return False sp, dp = sn.port(src_port), dn.port(dst_port) if sp is None or dp is None: return False if src_node == dst_node and src_port == dst_port: return False return self.relations.can_connect(sp.relation, dp.relation) def add_edge(self, edge: EdgeModel) -> EdgeModel: if edge.edge_id in self._edges: raise ValueError(f"duplicate edge id: {edge.edge_id}") if not edge.relation: sn = self._nodes.get(edge.source_node) sp = sn.port(edge.source_port) if sn else None if sp is not None: edge.relation = sp.relation self._edges[edge.edge_id] = edge self._mark_dirty() self.edge_added.emit(edge.edge_id) return edge def remove_edge(self, edge_id: str) -> bool: if edge_id not in self._edges: return False del self._edges[edge_id] self._mark_dirty() self.edge_removed.emit(edge_id) return True def edge(self, edge_id: str) -> Optional[EdgeModel]: return self._edges.get(edge_id) def edges(self) -> Iterator[EdgeModel]: return iter(self._edges.values()) def edges_for_node(self, node_id: str) -> list[EdgeModel]: return [e for e in self._edges.values() if e.source_node == node_id or e.target_node == node_id] def notify_edge_changed(self, edge_id: str) -> None: if edge_id in self._edges: self._mark_dirty() self.edge_changed.emit(edge_id) # -- bulk ------------------------------------------------------------- def clear(self) -> None: self._nodes.clear() self._edges.clear() self._id_counter = 0 self._mark_dirty() self.cleared.emit() def __len__(self) -> int: return len(self._nodes) + len(self._edges) def _mark_dirty(self) -> None: self.dirty = True self.modified.emit() # -- serialization ---------------------------------------------------- def to_dict(self) -> dict: return { "version": 1, "id_counter": self._id_counter, "nodes": [n.to_dict() for n in self._nodes.values()], "edges": [e.to_dict() for e in self._edges.values()], "relations": [r.to_dict() for r in self.relations], "active_scheme": self.relations.active_scheme, "catalogues": self.catalogues.to_dict()["catalogues"], } def to_json(self, indent: int = 2) -> str: return json.dumps(self.to_dict(), indent=indent) def load_dict(self, data: dict) -> None: """Replace document contents from a serialized dict (no signals per item).""" self._nodes.clear() self._edges.clear() for nd in data.get("nodes", []): n = NodeModel.from_dict(nd) self._nodes[n.node_id] = n for ed in data.get("edges", []): e = EdgeModel.from_dict(ed) self._edges[e.edge_id] = e self._id_counter = data.get("id_counter", len(self._nodes) + len(self._edges)) scheme = data.get("active_scheme") if scheme: self.relations.set_active_scheme(scheme) self.dirty = False self.cleared.emit() # tells the view to rebuild from scratch @classmethod def from_json(cls, text: str, **kwargs) -> "DiagramDocument": doc = cls(**kwargs) doc.load_dict(json.loads(text)) return doc