diff --git a/beamline_editor/BeamlineEditorWidgetAPI.py b/beamline_editor/BeamlineEditorWidgetAPI.py new file mode 100644 index 0000000..cd538a5 --- /dev/null +++ b/beamline_editor/BeamlineEditorWidgetAPI.py @@ -0,0 +1,35 @@ +# Lifecycle callbacks (set at construction) +on_element_added(type, id) +on_elements_deleted([id, ...]) +on_element_selected(id | None) +get_properties(id) → dict +on_property_changed(id, dict) +on_flatten(id, case_num) # always receives both args + +# Querying +editor.getElementType(id) # → "Quadrupole" etc. +editor.getElementName(id) # → display label +editor.getSecondaryConnection(id) # → id of node on out_port2, or None + +# Naming & properties +editor.setElementName(id, name) +editor.updatePropertyPane(id, data) → bool + +# Traversal +editor.flatten(start_id, case_num=1) → FlattenTree | None + tree.path # straight-line segment + tree.branch_id # id of the forking Line/Case node + tree.primary / .secondary / .tertiary # sub-trees + +# Templates +editor.get_templates() → dict # serialisable +editor.set_templates(dict) + +# Layout persistence +editor.saveLayout(path) +editor.loadLayout(path) +editor.importLayout(path) → {old_id: new_id} + +# Copy/paste +editor.copy_selected() +editor.paste() \ No newline at end of file diff --git a/beamline_editor/editor_widget.py b/beamline_editor/editor_widget.py index efc7ef3..6785ff3 100644 --- a/beamline_editor/editor_widget.py +++ b/beamline_editor/editor_widget.py @@ -61,7 +61,7 @@ class BeamlineEditorWidget(QWidget): on_element_selected: Callable[[str | None], None] | None = None, on_property_changed: Callable[[str, dict], None] | None = None, get_properties: Callable[[str], dict[str, Any]] | None = None, - on_flatten: Callable[[str], None] | None = None, + on_flatten: Callable | None = None, ): super().__init__(parent) @@ -217,35 +217,30 @@ class BeamlineEditorWidget(QWidget): self._templates.set_templates(data) # ── flatten API ─────────────────────────────────────────────────────────── - def flatten(self, start_id: str): + def flatten(self, start_id: str, case_num: int | None = None): """ - Build a branch tree starting at the element identified by - `start_id`, following the lattice forward. Only "Line" - containers may fork (every other element has a single - out-port). + Build a branch tree starting at `start_id`. + + Parameters + ---------- + start_id : str + node_id to start the walk from. + case_num : int | None + 1, 2, or 3 — when the walk encounters a Case node, follow + only the port for that case (with fallback to lower cases if + the requested port is unconnected). None follows all + connected branches of every Case node as sub-trees. Returns ------- FlattenTree | None - The root of a branch tree (see beamline_editor.flatten_tree - .FlattenTree). Each tree node holds a straight-line `path` - of node_ids, and — if that segment ended at a forking Line - — `branch_id` (that Line's node_id) plus `.primary` / - `.secondary` child sub-trees for its two out-ports. + Tree rooted at start_id, or None if start_id is not found. - Returns None if start_id is not found on the canvas. - - Example - ------- - tree = editor.flatten(start_id) - if tree: - print(tree.path) # first straight run - if tree.branch_id: - print("forks at", tree.branch_id) - print(tree.primary.path) # primary branch - print(tree.secondary.path) # secondary branch + At a "Line" fork: .primary / .secondary + At a "Case" fork: .primary / .secondary / .tertiary + (or just .primary when case_num selects a single branch) """ - return self.scene.flatten(start_id) + return self.scene.flatten(start_id, case_num) def getSecondaryConnection(self, node_id: str) -> str | None: """ diff --git a/beamline_editor/flatten_tree.py b/beamline_editor/flatten_tree.py index d3c879c..a794b04 100644 --- a/beamline_editor/flatten_tree.py +++ b/beamline_editor/flatten_tree.py @@ -35,76 +35,61 @@ class FlattenTree: Attributes ---------- node_id : str - node_id of the first element in this segment (usually the - element right after the branch that produced this sub-tree, - or the original start_id for the root). + node_id of the first element in this segment. element_type : str element_type of that first node. path : list[str] - Ordered node_ids of every element walked in this straight-line - segment, starting with node_id. Stops just before a fork, or at - a dead end / cycle (the terminating node IS included in path). + Ordered node_ids walked in this straight-line segment. + Stops just before a fork (the forking node IS included). branch_id : str | None - node_id of the "Line" element where this segment forked, or - None if the segment ended at a dead end / revisit instead. + node_id of the forking element ("Line" or "Case"), or None. primary : FlattenTree | None - Sub-tree continuing from the forking Line's primary out-port - (out_port, index 0). None if there was no fork, or that branch - was unconnected. + Sub-tree from out_port index 0 (case 1 / Line primary). secondary : FlattenTree | None - Sub-tree continuing from the forking Line's secondary out-port - (out_port2, index 1). None if there was no fork, or that branch - was unconnected. + Sub-tree from out_port index 1 (case 2 / Line secondary). + tertiary : FlattenTree | None + Sub-tree from out_port index 2 (case 3, Case nodes only). """ node_id: str element_type: str - path: list[str] = field(default_factory=list) - branch_id: str | None = None + path: list[str] = field(default_factory=list) + branch_id: str | None = None primary: "FlattenTree | None" = None secondary: "FlattenTree | None" = None + tertiary: "FlattenTree | None" = None # ── convenience ─────────────────────────────────────────────────────────── def is_leaf(self) -> bool: """True if this segment ends without a further fork.""" - return self.primary is None and self.secondary is None + return self.primary is None and self.secondary is None and self.tertiary is None def all_paths(self) -> list[list[str]]: - """ - Return every root-to-leaf path through the tree as a flat list - of node_id lists (each list is the concatenation of all - segments along that root-to-leaf route). - """ + """Return every root-to-leaf path as a flat list of node_id lists.""" if self.is_leaf(): return [list(self.path)] results: list[list[str]] = [] - if self.primary: - for sub in self.primary.all_paths(): - results.append(self.path + sub) - if self.secondary: - for sub in self.secondary.all_paths(): - results.append(self.path + sub) - if not results: - results = [list(self.path)] - return results + for child in (self.primary, self.secondary, self.tertiary): + if child: + for sub in child.all_paths(): + results.append(self.path + sub) + return results or [list(self.path)] def all_node_ids(self) -> list[str]: - """Return every node_id appearing anywhere in the tree (no duplicates, order not guaranteed).""" + """Every node_id in the tree (no duplicates).""" ids = set(self.path) - if self.primary: - ids.update(self.primary.all_node_ids()) - if self.secondary: - ids.update(self.secondary.all_node_ids()) + for child in (self.primary, self.secondary, self.tertiary): + if child: + ids.update(child.all_node_ids()) return list(ids) def find_branches(self) -> list["FlattenTree"]: - """Return every FlattenTree node in this tree that represents a fork (branch_id is not None).""" + """Every FlattenTree node where branch_id is not None.""" result = [] if self.branch_id is not None: result.append(self) - if self.primary: - result.extend(self.primary.find_branches()) - if self.secondary: - result.extend(self.secondary.find_branches()) + for child in (self.primary, self.secondary, self.tertiary): + if child: + result.extend(child.find_branches()) return result def to_dict(self) -> dict: @@ -116,11 +101,12 @@ class FlattenTree: "branch_id": self.branch_id, "primary": self.primary.to_dict() if self.primary else None, "secondary": self.secondary.to_dict() if self.secondary else None, + "tertiary": self.tertiary.to_dict() if self.tertiary else None, } @staticmethod def from_dict(data: dict | None) -> "FlattenTree | None": - """Reconstruct a FlattenTree from a dict produced by to_dict().""" + """Reconstruct a FlattenTree from to_dict() output.""" if data is None: return None return FlattenTree( @@ -130,4 +116,5 @@ class FlattenTree: branch_id = data.get("branch_id"), primary = FlattenTree.from_dict(data.get("primary")), secondary = FlattenTree.from_dict(data.get("secondary")), + tertiary = FlattenTree.from_dict(data.get("tertiary")), ) diff --git a/beamline_editor/nodes.py b/beamline_editor/nodes.py index 85dbfbe..bfe1a8d 100644 --- a/beamline_editor/nodes.py +++ b/beamline_editor/nodes.py @@ -4,20 +4,22 @@ nodes.py — Graphical node items for the beamline editor. Port layout ----------- Most elements : 1 in (left) + 1 out (right) - Line : 1 in (left) + 2 out (right primary + right offset) - — the ONLY element type that may branch. - Start : no in + 1 out (right) — beam source - End (Marker) : 2 in (left primary + left offset) + NO out — beam sink - Branch (Marker) : 1 in (left) + NO out — pure fork marker; - branching itself is expressed by which Line - containers are wired downstream, not by a second - out-port on Branch. + Line : 1 in (left) + 2 out (right primary + right offset) + — the only container type that may branch. + Start : no in + 1 out (right) — beam source + Case (Marker) : 1 in (left) + 3 out (right, evenly spaced top→bot) + out index 0 = bottom (case 1) + out index 1 = middle (case 2) + out index 2 = top (case 3) + Merge (Marker) : 3 in (left, evenly spaced top→bot) + 1 out (right) + in index 0 = bottom (case 1) + in index 1 = middle (case 2) + in index 2 = top (case 3) Display names ------------- Each node shows element_type by default. - Call node.set_display_name(name) to override what is painted on the canvas. - The element_type never changes (used for logic / templates / palette). + Call set_display_name(name) to override. element_type never changes. """ from __future__ import annotations @@ -49,31 +51,29 @@ _STYLES: dict[str, dict] = { "Vacuum": {"top": "#00838F", "bot": "#006064", "w": 64, "h": 36}, "Target": {"top": "#B71C1C", "bot": "#7B0000", "w": 36, "h": 36}, "Alignment": {"top": "#1565C0", "bot": "#0D3B8C", "w": 36, "h": 36}, - "Branch": {"top": "#F9A825", "bot": "#C17900", "w": 40, "h": 40}, "Start": {"top": "#00695C", "bot": "#004D40", "w": 54, "h": 40}, - "End": {"top": "#4E342E", "bot": "#3E2723", "w": 54, "h": 40}, + "Case": {"top": "#E65100", "bot": "#BF360C", "w": 60, "h": 72}, + "Merge": {"top": "#1A237E", "bot": "#0D1452", "w": 60, "h": 72}, } -# Elements whose in-port is suppressed (no left connection dot) -_NO_IN = {"Start"} -# Elements whose primary out-port is suppressed (no right connection dot) -_NO_OUT = {"End", "Branch"} -# Elements with a second exit port on the right. -# Per spec: ONLY "Line" may branch — every other element has at most -# one out-port. -_DUAL_EXIT = {"Line"} -# Elements with a second entry port on the left -_DUAL_ENTRY = {"End"} +# ── Port topology flags ─────────────────────────────────────────────────────── +_NO_IN = {"Start"} # no in-port at all +_NO_OUT = {"Merge"} # primary out only (set below), no suppression + # — Merge DOES have 1 out; this set is unused + # but kept for documentation clarity +_DUAL_EXIT = {"Line"} # exactly 2 out-ports +_TRIPLE_EXIT = {"Case"} # exactly 3 out-ports +_TRIPLE_ENTRY = {"Merge"} # exactly 3 in-ports class BeamlineNode(QGraphicsItem): - """Base graphical node. Subclasses override _draw_body() for shape variety.""" + """Base graphical node. Subclasses override _draw_body().""" def __init__(self, element_type: str, scene_pos: QPointF): super().__init__() self.element_type = element_type self.node_id = str(uuid.uuid4()) - self._display_name = element_type # overrideable via set_display_name() + self._display_name = element_type style = _STYLES.get(element_type, {"top": "#455A64", "bot": "#263238", "w": 64, "h": 40}) @@ -90,29 +90,38 @@ class BeamlineNode(QGraphicsItem): ) self.setZValue(2) - # ── ports ───────────────────────────────────────────────────────────── - # in_port: None for Start (no input) - self.in_port: Port | None = ( - None if element_type in _NO_IN - else Port(self, "in", 0) - ) - # second entry port (End only) - self.in_port2: Port | None = ( - Port(self, "in", 1) if element_type in _DUAL_ENTRY else None - ) - # out_port: None for End (no output) - self.out_port: Port | None = ( - None if element_type in _NO_OUT - else Port(self, "out", 0) - ) - # second exit port (Dipole, Overlap, Line) - self.out_port2: Port | None = ( - Port(self, "out", 1) if element_type in _DUAL_EXIT else None - ) + # ── in-ports ────────────────────────────────────────────────────────── + if element_type in _NO_IN: + self.in_port = None + self.in_port2 = None + self._extra_in_ports: list[Port] = [] + elif element_type in _TRIPLE_ENTRY: + # 3 in-ports: index 0=bottom(case1), 1=middle(case2), 2=top(case3) + self.in_port = Port(self, "in", 0) + self.in_port2 = Port(self, "in", 1) + self._extra_in_ports = [Port(self, "in", 2)] + else: + self.in_port = Port(self, "in", 0) + self.in_port2 = None + self._extra_in_ports = [] + + # ── out-ports ───────────────────────────────────────────────────────── + if element_type in _TRIPLE_EXIT: + # 3 out-ports: index 0=bottom(case1), 1=middle(case2), 2=top(case3) + self.out_port = Port(self, "out", 0) + self.out_port2 = Port(self, "out", 1) + self._extra_out_ports: list[Port] = [Port(self, "out", 2)] + elif element_type in _DUAL_EXIT: + self.out_port = Port(self, "out", 0) + self.out_port2 = Port(self, "out", 1) + self._extra_out_ports = [] + else: + self.out_port = Port(self, "out", 0) + self.out_port2 = None + self._extra_out_ports = [] self._place_ports() - # ── label ───────────────────────────────────────────────────────────── self._label = QGraphicsTextItem(self._display_name, self) self._label.setDefaultTextColor(TEXT_COLOR) self._label.setFont(QFont("Segoe UI", 7, QFont.Medium)) @@ -121,11 +130,6 @@ class BeamlineNode(QGraphicsItem): # ── display name API ────────────────────────────────────────────────────── def set_display_name(self, name: str): - """ - Override the text shown on the canvas node. - Does not affect element_type (used for logic/templates). - Pass an empty string or None to revert to the element_type default. - """ self._display_name = name if name else self.element_type self._label.setPlainText(self._display_name) self._center_label() @@ -136,17 +140,36 @@ class BeamlineNode(QGraphicsItem): # ── port layout ─────────────────────────────────────────────────────────── def _place_ports(self): + """Distribute ports evenly along the left/right edges.""" mid_y = self.H / 2 - offset = self.H * 0.35 - if self.in_port: - self.in_port.setPos(0, mid_y) - if self.in_port2: - self.in_port2.setPos(0, mid_y + offset) - if self.out_port: - self.out_port.setPos(self.W, mid_y) - if self.out_port2: - # Only "Line" reaches here — secondary exit sits above mid-plane - self.out_port2.setPos(self.W, mid_y - offset) + offset = self.H * 0.35 # used for 2-port elements + + # ── in-ports ────────────────────────────────────────────────────────── + all_in = self._all_in_ports() + if len(all_in) == 1: + all_in[0].setPos(0, mid_y) + elif len(all_in) == 2: + all_in[0].setPos(0, mid_y + offset) # bottom + all_in[1].setPos(0, mid_y - offset) # top + elif len(all_in) == 3: + # evenly spaced: top, middle, bottom + step = self.H / 4 + all_in[0].setPos(0, self.H - step) # index 0 → bottom (case 1) + all_in[1].setPos(0, mid_y) # index 1 → middle (case 2) + all_in[2].setPos(0, step) # index 2 → top (case 3) + + # ── out-ports ───────────────────────────────────────────────────────── + all_out = self._all_out_ports() + if len(all_out) == 1: + all_out[0].setPos(self.W, mid_y) + elif len(all_out) == 2: + all_out[0].setPos(self.W, mid_y + offset) # primary (bottom) + all_out[1].setPos(self.W, mid_y - offset) # secondary (top) + elif len(all_out) == 3: + step = self.H / 4 + all_out[0].setPos(self.W, self.H - step) # index 0 → bottom (case 1) + all_out[1].setPos(self.W, mid_y) # index 1 → middle (case 2) + all_out[2].setPos(self.W, step) # index 2 → top (case 3) def _center_label(self): br = self._label.boundingRect() @@ -188,15 +211,31 @@ class BeamlineNode(QGraphicsItem): # ── port helpers ────────────────────────────────────────────────────────── def _all_ports(self) -> list[Port]: - return [p for p in ( - self.in_port, self.in_port2, self.out_port, self.out_port2 - ) if p is not None] + return self._all_in_ports() + self._all_out_ports() def _all_in_ports(self) -> list[Port]: - return [p for p in (self.in_port, self.in_port2) if p is not None] + ports = [p for p in (self.in_port, self.in_port2) if p is not None] + ports.extend(self._extra_in_ports) + return ports def _all_out_ports(self) -> list[Port]: - return [p for p in (self.out_port, self.out_port2) if p is not None] + ports = [p for p in (self.out_port, self.out_port2) if p is not None] + ports.extend(self._extra_out_ports) + return ports + + def get_out_port(self, index: int) -> Port | None: + """Return the out-port with the given index, or None.""" + for p in self._all_out_ports(): + if p.index == index: + return p + return None + + def get_in_port(self, index: int) -> Port | None: + """Return the in-port with the given index, or None.""" + for p in self._all_in_ports(): + if p.index == index: + return p + return None # ── Specialised shapes ──────────────────────────────────────────────────────── @@ -282,61 +321,21 @@ class VariableLineNode(BeamlineNode): p.drawLine(QPointF(tip, mid_y), QPointF(tip + d * 6, mid_y - 4)) p.drawLine(QPointF(tip, mid_y), QPointF(tip + d * 6, mid_y + 4)) -class StartNode(BeamlineNode): - """Start marker: no in-port, 1 out. Beam source.""" - def __init__(self, pos): super().__init__("Start", pos) - def _draw_body(self, p): - p.setBrush(QBrush(self._gradient())) - p.setPen(QPen(QColor("#00897B"), 1.5)) - p.drawRoundedRect(QRectF(0, 0, self.W, self.H), 6, 6) - cx, cy = self.W / 2, self.H / 2 - pts = QPolygonF([ - QPointF(cx - 10, cy - 10), - QPointF(cx + 4, cy), - QPointF(cx - 10, cy + 10), - QPointF(cx - 6, cy), - ]) - p.setBrush(QBrush(QColor("#A5D6A7"))) - p.setPen(Qt.NoPen) - p.drawPolygon(pts) - -class EndNode(BeamlineNode): - """End marker: 2 in, no out-port. Beam sink.""" - def __init__(self, pos): super().__init__("End", pos) - def _draw_body(self, p): - p.setBrush(QBrush(self._gradient())) - p.setPen(QPen(QColor("#6D4C41"), 1.5)) - p.drawRoundedRect(QRectF(0, 0, self.W, self.H), 6, 6) - sq = 7 - p.setPen(Qt.NoPen) - for r in range(int(self.H // sq)): - for c in range(int(self.W // sq)): - col = QColor("#8D6E63") if (r + c) % 2 == 0 else QColor("#4E342E") - p.setBrush(QBrush(col)) - p.drawRect(QRectF(c * sq, r * sq, - min(sq, self.W - c * sq), - min(sq, self.H - r * sq))) - - class SolenoidNode(BeamlineNode): - """Solenoid coil — teal rounded rect with circular coil symbol.""" def __init__(self, pos): super().__init__("Solenoid", pos) def _draw_body(self, p): p.setBrush(QBrush(self._gradient())) p.setPen(QPen(QColor("#00838F"), 1.5)) p.drawRoundedRect(QRectF(0, 0, self.W, self.H), 6, 6) - # coil loops p.setPen(QPen(QColor("#80DEEA"), 1.2)) step = 10 - y1, y2 = 8.0, float(self.H - 8) x = 10.0 while x < self.W - 8: p.drawEllipse(QPointF(x + step / 2, self.H / 2), - step / 2, (y2 - y1) / 2) + step / 2, (self.H - 16) / 2) x += step class CorrectorNode(BeamlineNode): - """Corrector — purple square with crossed-arrows symbol.""" def __init__(self, pos): super().__init__("Corrector", pos) def _draw_body(self, p): p.setBrush(QBrush(self._gradient())) @@ -345,83 +344,79 @@ class CorrectorNode(BeamlineNode): cx, cy = self.W / 2, self.H / 2 r = min(cx, cy) - 5 p.setPen(QPen(QColor("#CE93D8"), 1.5)) - # vertical arrow p.drawLine(QPointF(cx, cy - r), QPointF(cx, cy + r)) p.drawLine(QPointF(cx, cy - r), QPointF(cx - 4, cy - r + 6)) p.drawLine(QPointF(cx, cy - r), QPointF(cx + 4, cy - r + 6)) - # horizontal arrow p.drawLine(QPointF(cx - r, cy), QPointF(cx + r, cy)) p.drawLine(QPointF(cx + r, cy), QPointF(cx + r - 6, cy - 4)) p.drawLine(QPointF(cx + r, cy), QPointF(cx + r - 6, cy + 4)) class RFNode(BeamlineNode): - """RF cavity — crimson rect with sine-wave decoration.""" def __init__(self, pos): super().__init__("RF", pos) def _draw_body(self, p): p.setBrush(QBrush(self._gradient())) p.setPen(QPen(QColor("#C2185B"), 1.5)) p.drawRoundedRect(QRectF(0, 0, self.W, self.H), 5, 5) - # sine wave path = QPainterPath() cy = self.H / 2 amp = self.H / 4 - 2 - pts = 40 x0, x1 = 8.0, float(self.W - 8) - for i in range(pts + 1): - t = i / pts + for i in range(41): + t = i / 40 x = x0 + t * (x1 - x0) y = cy - amp * math.sin(t * 2 * math.pi * 2) - if i == 0: - path.moveTo(x, y) - else: - path.lineTo(x, y) + if i == 0: path.moveTo(x, y) + else: path.lineTo(x, y) p.setPen(QPen(QColor("#F48FB1"), 1.5)) p.setBrush(Qt.NoBrush) p.drawPath(path) class UndulatorNode(BeamlineNode): - """Undulator / wiggler — olive-green rect with alternating pole marks.""" def __init__(self, pos): super().__init__("Undulator", pos) def _draw_body(self, p): p.setBrush(QBrush(self._gradient())) p.setPen(QPen(QColor("#558B2F"), 1.5)) p.drawRect(QRectF(0, 0, self.W, self.H)) - # alternating pole blocks pole_w = 10 n_poles = int(self.W // pole_w) for i in range(n_poles): - col = QColor("#AED581") if i % 2 == 0 else QColor("#33691E") - p.setBrush(QBrush(col)) - p.setPen(Qt.NoPen) + c1 = QColor("#AED581") if i % 2 == 0 else QColor("#33691E") + c2 = QColor("#33691E") if i % 2 == 0 else QColor("#AED581") x = i * pole_w - # top half pole + p.setBrush(QBrush(c1)); p.setPen(Qt.NoPen) p.drawRect(QRectF(x, 2, pole_w, self.H / 2 - 3)) - # bottom half pole (opposite polarity → alternate colour swapped) - col2 = QColor("#33691E") if i % 2 == 0 else QColor("#AED581") - p.setBrush(QBrush(col2)) + p.setBrush(QBrush(c2)) p.drawRect(QRectF(x, self.H / 2 + 1, pole_w, self.H / 2 - 3)) class DiagnosticsNode(BeamlineNode): - """Diagnostics — indigo rect with oscilloscope / eye symbol.""" def __init__(self, pos): super().__init__("Diagnostics", pos) def _draw_body(self, p): p.setBrush(QBrush(self._gradient())) p.setPen(QPen(QColor("#512DA8"), 1.5)) p.drawRoundedRect(QRectF(0, 0, self.W, self.H), 5, 5) cx, cy = self.W / 2, self.H / 2 - # eye outline rx, ry = self.W / 2 - 8, self.H / 2 - 6 p.setPen(QPen(QColor("#B39DDB"), 1.5)) p.setBrush(Qt.NoBrush) p.drawEllipse(QPointF(cx, cy), rx, ry) - # pupil p.setBrush(QBrush(QColor("#7E57C2"))) p.setPen(Qt.NoPen) p.drawEllipse(QPointF(cx, cy), 5, 5) +class VacuumNode(BeamlineNode): + def __init__(self, pos): super().__init__("Vacuum", pos) + def _draw_body(self, p): + p.setBrush(QBrush(self._gradient())) + p.setPen(QPen(QColor("#00838F"), 1.5)) + p.drawRoundedRect(QRectF(0, 0, self.W, self.H), 5, 5) + cx, cy = self.W / 2, self.H / 2 + r = min(cx, cy) - 6 + p.setBrush(Qt.NoBrush) + p.setPen(QPen(QColor("#80DEEA"), 1.5)) + p.drawEllipse(QPointF(cx, cy), r, r) + p.drawLine(QPointF(cx, cy - r - 4), QPointF(cx, cy + r + 4)) class TargetNode(BeamlineNode): - """Target marker — red circle with crosshair.""" def __init__(self, pos): super().__init__("Target", pos) def _draw_body(self, p): p.setBrush(QBrush(self._gradient())) @@ -437,13 +432,12 @@ class TargetNode(BeamlineNode): p.drawEllipse(QPointF(cx, cy), 4, 4) class AlignmentNode(BeamlineNode): - """Alignment marker — blue square with corner tick marks.""" def __init__(self, pos): super().__init__("Alignment", pos) def _draw_body(self, p): p.setBrush(QBrush(self._gradient())) p.setPen(QPen(QColor("#1565C0"), 1.5)) p.drawRect(QRectF(0, 0, self.W, self.H)) - t = 6.0 # tick length + t = 6.0 p.setPen(QPen(QColor("#90CAF9"), 1.5)) for cx, cy in [(0.0, 0.0), (float(self.W), 0.0), (0.0, float(self.H)), (float(self.W), float(self.H))]: @@ -452,48 +446,87 @@ class AlignmentNode(BeamlineNode): p.drawLine(QPointF(cx, cy), QPointF(cx + dx, cy)) p.drawLine(QPointF(cx, cy), QPointF(cx, cy + dy)) - -class BranchNode(BeamlineNode): - """ - Branch marker — 1 in, NO out-port. - - Pure topological marker indicating a fork point in the lattice - documentation/diagram; the actual branching is expressed by which - Line containers are wired downstream of this point (each Line may - have up to 2 out-ports). Rendered as an amber diamond with a - three-way fork glyph (no out-port dots are drawn since none exist). - """ - def __init__(self, pos): super().__init__("Branch", pos) +class StartNode(BeamlineNode): + """Start marker: no in-port, 1 out.""" + def __init__(self, pos): super().__init__("Start", pos) def _draw_body(self, p): + p.setBrush(QBrush(self._gradient())) + p.setPen(QPen(QColor("#00897B"), 1.5)) + p.drawRoundedRect(QRectF(0, 0, self.W, self.H), 6, 6) cx, cy = self.W / 2, self.H / 2 - diamond = QPolygonF([ - QPointF(cx, 2), QPointF(self.W - 2, cy), - QPointF(cx, self.H - 2), QPointF(2, cy), + pts = QPolygonF([ + QPointF(cx - 10, cy - 10), QPointF(cx + 4, cy), + QPointF(cx - 10, cy + 10), QPointF(cx - 6, cy), ]) - p.setBrush(QBrush(self._gradient())) - p.setPen(QPen(QColor("#C17900"), 1.5)) - p.drawPolygon(diamond) - # decorative fork glyph (visual only — no real second out-port) - p.setPen(QPen(QColor("#FFECB3"), 1.4)) - p.drawLine(QPointF(cx - 6, cy), QPointF(cx + 2, cy)) - p.drawLine(QPointF(cx + 2, cy), QPointF(cx + 9, cy - 7)) - p.drawLine(QPointF(cx + 2, cy), QPointF(cx + 9, cy + 7)) + p.setBrush(QBrush(QColor("#A5D6A7"))) + p.setPen(Qt.NoPen) + p.drawPolygon(pts) +class CaseNode(BeamlineNode): + """ + Case marker — 1 in, 3 out. -class VacuumNode(BeamlineNode): - """Vacuum element — teal rect with a pump/vessel symbol.""" - def __init__(self, pos): super().__init__("Vacuum", pos) + Out-port indices (and their position on the right edge): + 0 = bottom → Case 1 + 1 = middle → Case 2 + 2 = top → Case 3 + + Rendered as an orange pentagon with 1/2/3 labels next to the ports. + """ + def __init__(self, pos): super().__init__("Case", pos) def _draw_body(self, p): p.setBrush(QBrush(self._gradient())) - p.setPen(QPen(QColor("#00838F"), 1.5)) - p.drawRoundedRect(QRectF(0, 0, self.W, self.H), 5, 5) - # simple pump symbol: circle with a vertical line + p.setPen(QPen(QColor("#BF360C"), 1.5)) + # pentagon pointing right cx, cy = self.W / 2, self.H / 2 - r = min(cx, cy) - 6 - p.setBrush(Qt.NoBrush) - p.setPen(QPen(QColor("#80DEEA"), 1.5)) - p.drawEllipse(QPointF(cx, cy), r, r) - p.drawLine(QPointF(cx, cy - r - 4), QPointF(cx, cy + r + 4)) + pts = QPolygonF([ + QPointF(0, 4), + QPointF(self.W - 12, 4), + QPointF(self.W, cy), + QPointF(self.W - 12, self.H - 4), + QPointF(0, self.H - 4), + ]) + p.drawPolygon(pts) + # case number labels + p.setPen(QPen(QColor("#FFE0B2"), 1)) + p.setFont(QFont("Segoe UI", 6)) + step = self.H / 4 + for i, label in enumerate(["1", "2", "3"]): + y = self.H - step - i * step # bottom=1, middle=2, top=3 + p.drawText(QRectF(self.W - 16, y - 6, 12, 12), + Qt.AlignCenter, label) + +class MergeNode(BeamlineNode): + """ + Merge marker — 3 in, 1 out. + + In-port indices (and their position on the left edge): + 0 = bottom → Case 1 + 1 = middle → Case 2 + 2 = top → Case 3 + + Rendered as a deep-blue pentagon pointing right with 1/2/3 labels. + """ + def __init__(self, pos): super().__init__("Merge", pos) + def _draw_body(self, p): + p.setBrush(QBrush(self._gradient())) + p.setPen(QPen(QColor("#0D1452"), 1.5)) + cx, cy = self.W / 2, self.H / 2 + pts = QPolygonF([ + QPointF(12, 4), + QPointF(self.W, 4), + QPointF(self.W, self.H - 4), + QPointF(12, self.H - 4), + QPointF(0, cy), + ]) + p.drawPolygon(pts) + p.setPen(QPen(QColor("#C5CAE9"), 1)) + p.setFont(QFont("Segoe UI", 6)) + step = self.H / 4 + for i, label in enumerate(["1", "2", "3"]): + y = self.H - step - i * step + p.drawText(QRectF(4, y - 6, 12, 12), + Qt.AlignCenter, label) # ── Factory ─────────────────────────────────────────────────────────────────── @@ -511,10 +544,10 @@ _NODE_MAP: dict[str, type] = { "Reference": ReferenceNode, "Target": TargetNode, "Alignment": AlignmentNode, - "Branch": BranchNode, "Overlap": OverlapNode, "Start": StartNode, - "End": EndNode, + "Case": CaseNode, + "Merge": MergeNode, "Line": LineNode, "Variable Line": VariableLineNode, } diff --git a/beamline_editor/palette.py b/beamline_editor/palette.py index 15d595c..d49f972 100644 --- a/beamline_editor/palette.py +++ b/beamline_editor/palette.py @@ -18,7 +18,7 @@ ELEMENT_GROUPS: list[tuple[str, list[str]]] = [ ("Elements", ["Dipole", "Quadrupole", "Sextupole", "Solenoid", "Corrector", "RF", "Undulator", "Diagnostics", "Vacuum"]), - ("Markers", ["Reference", "Start", "End", "Target", "Alignment", "Branch"]), + ("Markers", ["Reference", "Start", "Target", "Alignment", "Case", "Merge"]), ("Containers", ["Line", "Variable Line"]), ] diff --git a/beamline_editor/scene.py b/beamline_editor/scene.py index d8179fe..218c071 100644 --- a/beamline_editor/scene.py +++ b/beamline_editor/scene.py @@ -62,6 +62,9 @@ class BeamlineScene(QGraphicsScene): self._temp_line: TempConnection | None = None self._clipboard: list[tuple[str, QPointF]] = [] + self._collapsed_starts: set[str] = set() # node_ids of collapsed Start markers + # {start_node_id: {downstream_node_id: QPointF offset from start}} + self._collapse_offsets: dict[str, dict[str, QPointF]] = {} self._prev_selected: set[str] = set() self.selectionChanged.connect(self._on_selection_changed) @@ -225,13 +228,31 @@ class BeamlineScene(QGraphicsScene): act_copy.triggered.connect(lambda: self._ctx_copy(node)) menu.addAction(act_copy) - # Flatten — calls on_flatten with this node's ID + # Collapse / Uncollapse — only for Start markers + if node.element_type == "Start": + menu.addSeparator() + is_collapsed = node.node_id in self._collapsed_starts + label = "Uncollapse downstream" if is_collapsed else "Collapse downstream" + act_collapse = QAction(label, menu) + act_collapse.triggered.connect( + lambda: self._ctx_toggle_collapse(node) + ) + menu.addAction(act_collapse) + + # Flatten — always a submenu with Case 1/2/3; Case 1 is the default. + # For non-Case elements the case_num is passed through to flatten() + # which simply follows the primary path (case_num has no effect on + # elements that are not Case markers). menu.addSeparator() - act_flatten = QAction("Flatten from here", menu) - act_flatten.triggered.connect( - lambda: self._ctx_flatten(node) - ) - menu.addAction(act_flatten) + flatten_menu = menu.addMenu("Flatten from here") + flatten_menu.setStyleSheet(menu.styleSheet()) + for case_num in (1, 2, 3): + act = QAction(f"Case {case_num}" + (" (default)" if case_num == 1 else ""), + flatten_menu) + act.triggered.connect( + lambda checked, n=node, c=case_num: self._ctx_flatten(n, c) + ) + flatten_menu.addAction(act) menu.addSeparator() @@ -263,16 +284,122 @@ class BeamlineScene(QGraphicsScene): node.setSelected(True) self.copy_selected() - def _ctx_flatten(self, node: BeamlineNode): - """Called when 'Flatten from here' is chosen on a Start marker.""" + def _ctx_flatten(self, node: BeamlineNode, case_num: int): + """ + Called when 'Flatten from here → Case N' is chosen. + case_num (1/2/3) is always forwarded so that any Case nodes + encountered during the downstream walk follow the correct branch. + """ if self._on_flatten: - self._on_flatten(node.node_id) + self._on_flatten(node.node_id, case_num) else: self.status_changed.emit( - f"Flatten triggered from Start [{node.node_id[:8]}] " - "(no on_flatten callback registered)." + f"Flatten triggered from {node.element_type} " + f"[{node.node_id[:8]}] Case {case_num}" + " (no on_flatten callback registered)." ) + def _ctx_toggle_collapse(self, start_node: BeamlineNode): + """ + Toggle the collapsed state of all downstream nodes and wires + that follow `start_node` (a Start marker). + + When collapsed: + - Every downstream BeamlineNode is hidden. + - Every Connection that touches at least one hidden node is hidden. + - The Start marker itself stays visible, but its label changes to + show "(collapsed)" as a visual cue. + - Each downstream node's position relative to the Start marker + is stored so that uncollapsing after a move keeps the layout. + + When uncollapsed: + - All hidden items are repositioned relative to the Start marker's + current scene position, then made visible again. + - The Start marker label reverts to its normal display name. + """ + is_collapsed = start_node.node_id in self._collapsed_starts + nodes, conns = self._downstream_items(start_node) + + if is_collapsed: + # ── uncollapse: restore positions relative to current Start pos ── + self._collapsed_starts.discard(start_node.node_id) + start_pos = start_node.pos() + offsets = self._collapse_offsets.pop(start_node.node_id, {}) + for n in nodes: + if n.node_id in offsets: + n.setPos(start_pos + offsets[n.node_id]) + n.setVisible(True) + for c in conns: + c.setVisible(True) + c.update_path() + # revert label (strip suffix added at collapse time) + start_node.set_display_name( + start_node.display_name().replace(" (collapsed)", "") + ) + self.status_changed.emit( + f"Uncollapsed downstream of {start_node.display_name()}" + ) + else: + # ── collapse: snapshot relative positions, then hide ────────────── + self._collapsed_starts.add(start_node.node_id) + start_pos = start_node.pos() + offsets: dict[str, QPointF] = {} + for n in nodes: + offsets[n.node_id] = n.pos() - start_pos + n.setVisible(False) + self._collapse_offsets[start_node.node_id] = offsets + for c in conns: + c.setVisible(False) + name = start_node.display_name() + if "(collapsed)" not in name: + start_node.set_display_name(name + " (collapsed)") + self.status_changed.emit(f"Collapsed downstream of {name}") + + def _downstream_items( + self, + start_node: BeamlineNode, + ) -> tuple[list[BeamlineNode], list]: + """ + Collect every node and connection reachable forward from + `start_node` via out-port connections (excluding start_node itself). + + Uses a breadth-first walk over all out-ports so it handles + Line forks, Case three-way forks, and Merge nodes correctly. + + Returns + ------- + (nodes, connections) + nodes : list[BeamlineNode] — all downstream nodes + connections : list[Connection] — all wires touching any of + those nodes (or connecting start_node to them) + """ + from .connections import Connection as Conn + + visited_nodes: set[str] = {start_node.node_id} + queue: list[BeamlineNode] = [start_node] + downstream_nodes: list[BeamlineNode] = [] + downstream_conns: set = set() + + while queue: + current = queue.pop(0) + for port in current._all_out_ports(): + for conn in list(port.connections): + downstream_conns.add(conn) + dst_node = conn.dst_port.node + if dst_node.node_id not in visited_nodes: + visited_nodes.add(dst_node.node_id) + downstream_nodes.append(dst_node) + queue.append(dst_node) + + # Also include wires that arrive INTO downstream nodes from outside + # (e.g. Merge inputs from other branches) so nothing is left dangling. + for dn in downstream_nodes: + for port in dn._all_in_ports(): + for conn in port.connections: + downstream_conns.add(conn) + + return downstream_nodes, list(downstream_conns) + def _show_canvas_context_menu(self, screen_pos): """Right-click on empty canvas — shows Paste action.""" menu = QMenu() @@ -389,54 +516,43 @@ class BeamlineScene(QGraphicsScene): for conn in list(port.connections): conn.remove() self.removeItem(item) + # clean up any collapse state for this node + self._collapsed_starts.discard(item.node_id) + self._collapse_offsets.pop(item.node_id, None) # ── Flatten traversal ───────────────────────────────────────────────────── def flatten( self, start_id: str, + case_num: int | None = None, ) -> "FlattenTree | None": """ - Build a branch tree by walking forward from the node identified - by `start_id`, following single out-port connections through - ordinary elements and Start, and treating "Line" containers as - the only elements that may fork (via their second out-port). + Build a branch tree starting at `start_id`. - Topology assumptions (see nodes.py) - ------------------------------------ - - Every element has at most ONE out-port, except "Line", - which may have a connected secondary out-port (a fork). - - "End" and "Branch" markers have NO out-port (dead ends). - - Walking simply follows out_port → in_port edges; when a - "Line" node's out_port2 is also connected, two independent - sub-trees are built — one from out_port's target, one from - out_port2's target — and attached as .primary / .secondary - on the returned tree node for that Line. + Branch points + ------------- + "Line" containers with a connected secondary out-port fork into + .primary (out_port index 0) and .secondary (out_port2 index 1). + + "Case" markers fork into three branches: + out_port index 0 → case 1 (bottom port) + out_port2 index 1 → case 2 (middle port) + extra index 2 → case 3 (top port) + + When `case_num` (1/2/3) is supplied, the walk at a Case node + attempts to follow the requested case port. If that port has no + connection, it falls back to lower-numbered connected ports + (e.g. case 3 falls to 2 then 1). If no port is connected the + Case node is a dead end. When case_num is None (or the node is + not a Case marker) all connected ports are followed as separate + sub-trees (.primary / .secondary / .tertiary on the FlattenTree). + + "Merge" markers have no out-port — they are always a dead end + for the forward walk (the path stops there). Returns ------- - FlattenTree | None - A tree of FlattenTree nodes (see class below), rooted at - `start_id`. Returns None if start_id is not found on the - canvas. - - Each FlattenTree node has: - node_id : str - element_type : str - path : list[str] — node_ids from this tree-node's - id up to (but not including) the next branch - point or a dead end / cycle - branch_id : str | None — node_id of the Line where this - segment's path stopped because it forks - (None if path ended at a dead end / cycle - instead) - primary : FlattenTree | None — sub-tree from the - forking Line's out_port (index 0) - secondary : FlattenTree | None — sub-tree from the - forking Line's out_port2 (index 1) - - Use .walk() or simply recurse over .primary / .secondary to - traverse the whole tree. See FlattenTree docstring for - convenience accessors. + FlattenTree | None — None if start_id not on canvas. """ all_nodes: dict[str, BeamlineNode] = { item.node_id: item @@ -447,30 +563,46 @@ class BeamlineScene(QGraphicsScene): if start is None: return None - def _next_via(port, seen: set): - """Return the node wired to `port`'s connection not yet seen.""" - if port is None: + def _connected_node(port) -> BeamlineNode | None: + """Return the node wired to port's first connection, ignoring seen.""" + if port is None or not port.connections: return None - for conn in port.connections: - dst = conn.dst_port.node - if dst.node_id not in seen: - return dst - return None + return port.connections[0].dst_port.node - def _build(node: BeamlineNode, seen: set) -> "FlattenTree": + def _resolve_case(node: BeamlineNode, requested: int | None): """ - Walk forward from `node` via primary out-ports, collecting a - straight-line path, until we hit: - (a) a dead end (no out-port, or out-port unconnected), or - (b) a node already in `seen` (cycle), or - (c) a "Line" node whose out_port2 is ALSO connected - (a genuine fork) — recurse into both branches. + Return a list of (port_index, next_node) pairs for a Case node. + + Port index mapping (0-based internally, 1-based to user): + index 0 → Case 1 (bottom port) + index 1 → Case 2 (middle port) + index 2 → Case 3 (top port) + + When requested is 1/2/3: try the exact port first, then fall + back to lower-numbered connected ports (3→2→1). + When requested is None: return all connected ports. """ + all_out = sorted(node._all_out_ports(), key=lambda p: p.index) + if requested is None: + return [(p.index, _connected_node(p)) + for p in all_out if p.connections] + # 1-based → 0-based index, then fall back downward + target_idx = requested - 1 # case 3 → idx 2, case 2 → idx 1, etc. + for idx in range(target_idx, -1, -1): + port = node.get_out_port(idx) + nxt = _connected_node(port) + if nxt is not None: + return [(idx, nxt)] + return [] + + def _build(node: BeamlineNode, seen: set, + case_hint: int | None) -> "FlattenTree": path: list[str] = [] current = node - branch_id: str | None = None - primary_tree: "FlattenTree | None" = None - secondary_tree: "FlattenTree | None" = None + branch_id = None + primary = None + secondary = None + tertiary = None while current is not None: if current.node_id in seen: @@ -478,35 +610,54 @@ class BeamlineScene(QGraphicsScene): seen.add(current.node_id) path.append(current.node_id) - is_fork = ( + # ── Case node: 3-way fork ───────────────────────────────── + if current.element_type == "Case": + branch_id = current.node_id + branches = _resolve_case(current, case_hint) + sub_trees = [] + for _idx, nxt in branches: + # skip if the downstream node was already visited on + # another branch (e.g. two ports wired to same element) + if nxt is not None and nxt.node_id not in seen: + sub_trees.append(_build(nxt, seen, case_hint)) + if len(sub_trees) >= 1: primary = sub_trees[0] + if len(sub_trees) >= 2: secondary = sub_trees[1] + if len(sub_trees) >= 3: tertiary = sub_trees[2] + break + + # ── Line node: 2-way fork ───────────────────────────────── + is_line_fork = ( current.element_type == "Line" - and getattr(current, "out_port2", None) is not None + and current.out_port2 is not None and current.out_port2.connections ) - - if is_fork: + if is_line_fork: branch_id = current.node_id - prim_next = _next_via(current.out_port, seen) - sec_next = _next_via(current.out_port2, seen) - if prim_next is not None: - primary_tree = _build(prim_next, seen) - if sec_next is not None: - secondary_tree = _build(sec_next, seen) + prim_nxt = _connected_node(current.out_port) + sec_nxt = _connected_node(current.out_port2) + if prim_nxt and prim_nxt.node_id not in seen: + primary = _build(prim_nxt, seen, case_hint) + if sec_nxt and sec_nxt.node_id not in seen: + secondary = _build(sec_nxt, seen, case_hint) break - # Not a fork — continue straight-line walk - nxt = _next_via(getattr(current, "out_port", None), seen) - if nxt is None: + # ── straight-line walk ──────────────────────────────────── + out_port = getattr(current, "out_port", None) + nxt = _connected_node(out_port) + if nxt is None or nxt.node_id in seen: break current = nxt + current = nxt return FlattenTree( node_id = path[0] if path else node.node_id, - element_type = all_nodes[path[0]].element_type if path else node.element_type, + element_type = all_nodes[path[0]].element_type if path + else node.element_type, path = path, branch_id = branch_id, - primary = primary_tree, - secondary = secondary_tree, + primary = primary, + secondary = secondary, + tertiary = tertiary, ) - return _build(start, set()) + return _build(start, set(), case_num) diff --git a/main.py b/main.py index aa7850d..a616e64 100644 --- a/main.py +++ b/main.py @@ -49,6 +49,7 @@ class BeamlineEditor(QtWidgets.QMainWindow, Ui_BeamlineGUI): self.PL=ProtoListe(0) self.loadFile('Layouts/SFTest2.json') + self.loadTemplates() def saveas(self): @@ -204,7 +205,8 @@ class BeamlineEditor(QtWidgets.QMainWindow, Ui_BeamlineGUI): return - def flatten(self, node_id ): + def flatten(self, node_id, case_num=1 ): + print('Flatten case:', case_num) self.lines.clear() type = self.editor.getElementType(node_id) name='XXX' @@ -216,16 +218,16 @@ class BeamlineEditor(QtWidgets.QMainWindow, Ui_BeamlineGUI): line=LineContainer(namein=name,Lin=0) line.clear() print('\nUnwrapping beamline...\n') - self.unwrap(name,node_id,line) + self.unwrap(name,node_id,line,case_num) self.PL.generateLayout(line) plrecs = [self.PL.info[ele] for ele in self.PL.order] populate_table(self.UIProtoList,plrecs) - def unwrap(self, name, node_id,line): + def unwrap(self, name, node_id,line,case_num): print('Unwrapping Line: %s' % name) - tree=self.editor.flatten(node_id) + tree=self.editor.flatten(node_id,case_num) for ele in tree.path: self.report(ele, name) type = self.editor.getElementType(ele) @@ -237,7 +239,7 @@ class BeamlineEditor(QtWidgets.QMainWindow, Ui_BeamlineGUI): name1 = subline.Name ref = subline.firstElementID if ref in self.elementDB.keys(): - self.unwrap(name + name1, ref, subline) + self.unwrap(name + name1, ref, subline,case_num) else: print('Undefined reference in Line Container:', name + name1) elif type == 'Variable Line': @@ -247,6 +249,8 @@ class BeamlineEditor(QtWidgets.QMainWindow, Ui_BeamlineGUI): else: self.elementDB[ele].ref=None line.append(self.elementDB[ele],sRef=0,Ref='relative') + elif type == 'Case' or type == 'Merge': + continue else: sRef=self.elementDB[ele].OffsetS Ref='absolute' @@ -255,7 +259,7 @@ class BeamlineEditor(QtWidgets.QMainWindow, Ui_BeamlineGUI): line.append(self.elementDB[ele],sRef=sRef,Ref=Ref) self.elementDB[ele].Prefix=name if tree.primary: - self.unwrap(name,tree.primary.node_id,line) + self.unwrap(name,tree.primary.node_id,line,case_num) @@ -265,6 +269,9 @@ class BeamlineEditor(QtWidgets.QMainWindow, Ui_BeamlineGUI): if type == 'Variable Line': print('FLATTEN: Adding Variable Line') return + if not node_id in self.elementDB.keys(): + print('Generic element of type:',type) + return name = self.elementDB[node_id].Name if type=='Start': name=prefix