Implementing switching of beamline path with a case structure

This commit is contained in:
2026-07-09 14:56:16 +02:00
parent 7452f69e6d
commit ba3a514920
7 changed files with 530 additions and 322 deletions
@@ -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()
+18 -23
View File
@@ -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:
"""
+29 -42
View File
@@ -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")),
)
+200 -167
View File
@@ -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,
}
+1 -1
View File
@@ -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"]),
]
+234 -83
View File
@@ -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)
+13 -6
View File
@@ -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