diff --git a/pyproject.toml b/pyproject.toml
index 044818a6..e4529c7e 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -18,7 +18,7 @@ dependencies = [
"python-redis-lock==4.0.0",
"fastapi==0.115.13",
"uvicorn==0.34.2",
- "aaredb==0.1.1a42",
+ "aaredb==0.1.1a43",
"python_multipart==0.0.20",
"websocket-client==1.8.0",
"sseclient-py==1.8.0",
diff --git a/src/aare/common/aerotech_models.py b/src/aare/common/aerotech_models.py
new file mode 100644
index 00000000..786ea1c3
--- /dev/null
+++ b/src/aare/common/aerotech_models.py
@@ -0,0 +1,155 @@
+from enum import Enum
+from typing import Optional
+
+from pydantic import BaseModel, ConfigDict, Field
+
+from aare.common.rotation_scan import RotationScanRequest
+
+
+class TaskEnum(Enum):
+ TASK_0 = 0
+ TASK_1 = 1
+ TASK_2 = 2
+ TASK_3 = 3
+ TASK_4 = 4
+ TASK_5 = 5
+ TASK_6 = 6
+ TASK_7 = 7
+ TASK_8 = 8
+ TASK_9 = 9
+
+
+class AxisEnum(Enum):
+ X = "x"
+ Y = "y"
+ Z = "z"
+ OMEGA = "u"
+
+
+class AerotechRunEnum(Enum):
+ STOP = 0
+ START = 1
+ RUN = 2
+ LOAD = 3
+ PAUSE = 4
+ RESET = 5
+
+
+class VariableTypeEnum(Enum):
+ INT = 0
+ REAL = 1
+ STRING = 2
+
+
+class AerotechAxisStatus(BaseModel):
+ enabled: bool
+ fault: int
+ homed: bool
+ is_fault: bool
+ moving: bool
+ position: float
+ status: int
+ velocity: float
+
+
+class AerotechStatus(BaseModel):
+ state: str
+ x: Optional[AerotechAxisStatus] = None
+ y: Optional[AerotechAxisStatus] = None
+ z: Optional[AerotechAxisStatus] = None
+ u: Optional[AerotechAxisStatus] = None
+
+ model_config = ConfigDict(extra="allow")
+
+ def __str__(self) -> str:
+ return self.to_pretty_string()
+
+ def to_pretty_string(self) -> str:
+ def axis_line(name: str, axis: AerotechAxisStatus | None) -> str:
+ if axis is None:
+ return f"{name.upper()}: unavailable"
+ return (
+ f"{name.upper():>2} | pos={axis.position:>12.6f} | "
+ f"homed={axis.homed!s:<5} | moving={axis.moving!s:<5} | "
+ f"enabled={axis.enabled!s:<5} | fault={axis.fault} | "
+ f"faulted={axis.is_fault!s:<5} | vel={axis.velocity:>10.6f}"
+ )
+
+ return "\n".join(
+ [
+ f"STATE: {self.state}",
+ axis_line("x", self.x),
+ axis_line("y", self.y),
+ axis_line("z", self.z),
+ axis_line("u", self.u),
+ ]
+ )
+
+ def to_compact_string(self) -> str:
+ axes = []
+ for name in ("x", "y", "z", "u"):
+ axis = getattr(self, name)
+ if axis is not None:
+ axes.append(
+ f"{name}={axis.position:.4f} "
+ f"({'H' if axis.homed else 'NH'}, {'M' if axis.moving else '-'})"
+ )
+ return f"state={self.state} | " + " | ".join(axes)
+
+ def to_colored_string(self) -> str:
+ # ANSI colors for terminal use
+ RESET = "\033[0m"
+ BOLD = "\033[1m"
+ CYAN = "\033[36m"
+ GREEN = "\033[32m"
+ YELLOW = "\033[33m"
+ RED = "\033[31m"
+
+ def color_bool(value: bool) -> str:
+ return f"{GREEN}True{RESET}" if value else f"{RED}False{RESET}"
+
+ def axis_line(name: str, axis: AerotechAxisStatus | None) -> str:
+ if axis is None:
+ return f"{YELLOW}{name.upper()}: unavailable{RESET}"
+ return (
+ f"{BOLD}{name.upper()}{RESET} | "
+ f"pos={CYAN}{axis.position:>12.6f}{RESET} | "
+ f"homed={color_bool(axis.homed)} | "
+ f"moving={color_bool(axis.moving)} | "
+ f"enabled={color_bool(axis.enabled)} | "
+ f"fault={axis.fault} | "
+ f"faulted={color_bool(axis.is_fault)} | "
+ f"vel={axis.velocity:>10.6f}"
+ )
+
+ return "\n".join(
+ [
+ f"{BOLD}STATE:{RESET} {CYAN}{self.state}{RESET}",
+ axis_line("x", self.x),
+ axis_line("y", self.y),
+ axis_line("z", self.z),
+ axis_line("u", self.u),
+ ]
+ )
+
+
+class AerotechTarget(BaseModel):
+ x: Optional[float] = None
+ y: Optional[float] = None
+ z: Optional[float] = None
+ u: Optional[float] = None
+
+ def to_payload(self) -> dict:
+ return self.model_dump(exclude_none=True)
+
+
+class AerotechRotationScanRequest(RotationScanRequest):
+ rotation_deg: float
+ time_sec: float
+ start_pos_deg: float
+ async_move: bool = Field(default=False, alias="async")
+
+ model_config = ConfigDict(populate_by_name=True)
+
+ def to_payload(self) -> dict:
+ return self.model_dump(by_alias=True, exclude_none=True)
\ No newline at end of file
diff --git a/src/aare/common/auth_models.py b/src/aare/common/auth_models.py
new file mode 100644
index 00000000..9930e7cd
--- /dev/null
+++ b/src/aare/common/auth_models.py
@@ -0,0 +1,52 @@
+from enum import Enum
+from pydantic import BaseModel
+from datetime import datetime
+
+class BatonRequestStatus(Enum):
+ PENDING = "pending"
+ ACCEPTED = "accepted"
+ REFUSED = "refused"
+ TIMEOUT = "timeout"
+ CANCELLED = "cancelled"
+
+
+class BatonHolderInfo(BaseModel):
+ """Information about the current baton holder."""
+ username: str
+ session: int
+ is_staff: bool
+ pgroup: str | None = None
+
+
+class BatonRequest(BaseModel):
+ """A request from one user to take the baton from another."""
+ request_id: str
+ requester_username: str
+ requester_session: int
+ requester_is_staff: bool
+ holder_username: str | None = None
+ holder_session: int | None = None
+ created_at: float # Unix timestamp
+ timeout_seconds: int = 30
+ status: BatonRequestStatus = BatonRequestStatus.PENDING
+
+
+class BatonTransferQueue(BaseModel):
+ """Queued baton transfer waiting for beamline to be available."""
+ target_session: int
+ target_username: str
+ target_is_staff: bool
+ target_pgroup: str | None = None
+ queued_at: float # Unix timestamp
+ reason: str = "beamline_busy"
+
+
+class BatonStatus(BaseModel):
+ """Full baton status for GUI display."""
+ holder: BatonHolderInfo | None = None
+ pending_request: BatonRequest | None = None
+ queued_transfer: BatonTransferQueue | None = None
+ you_are_holder: bool = False
+ you_have_pending_request: bool = False
+ incoming_request: bool = False
+ allow_non_staff_request: bool = False
\ No newline at end of file
diff --git a/src/aare/common/automation_models.py b/src/aare/common/automation_models.py
new file mode 100644
index 00000000..7431e794
--- /dev/null
+++ b/src/aare/common/automation_models.py
@@ -0,0 +1,182 @@
+from __future__ import annotations
+
+from dataclasses import dataclass, field
+from enum import Enum
+from typing import Any
+import time
+import uuid
+
+from pydantic import BaseModel, Field
+
+class WorkflowMode(str, Enum):
+ FLEXIBLE_MANUAL = "flexible_manual"
+ GUIDED_MANUAL = "guided_manual"
+ AUTOMATION = "automation"
+
+
+class WorkflowStateKind(str, Enum):
+ MOUNT = "mount"
+ LOOP_CENTRE = "loop_centre"
+ RASTER = "raster"
+ DATA_COLLECTION = "data_collection"
+
+
+class StepStatus(str, Enum):
+ PENDING = "pending"
+ RUNNING = "running"
+ SUCCESS = "success"
+ FAILED = "failed"
+ SKIPPED = "skipped"
+ PAUSED = "paused"
+
+
+@dataclass(frozen=True)
+class TransitionRule:
+ to_state: WorkflowStateKind
+ allowed_modes: frozenset[WorkflowMode] = frozenset()
+ optional: bool = False
+ condition_name: str | None = None
+
+
+@dataclass(frozen=True)
+class StateDefinition:
+ kind: WorkflowStateKind
+ transitions: tuple[TransitionRule, ...]
+ description: str = ""
+
+
+@dataclass
+class StateResult:
+ state: WorkflowStateKind
+ status: StepStatus
+ message: str = ""
+ payload: dict[str, Any] = field(default_factory=dict)
+
+
+@dataclass
+class WorkflowContext:
+ mode: WorkflowMode
+ queue_id: str
+ item_id: str
+ sample_id: int | None = None
+ current_state: WorkflowStateKind | None = None
+ current_step_index: int = 0
+ paused: bool = False
+ abort_requested: bool = False
+ last_message: str = ""
+ metadata: dict[str, Any] = field(default_factory=dict)
+
+
+class QueueItemStatus(str, Enum):
+ PENDING = "pending"
+ RUNNING = "running"
+ COMPLETED = "completed"
+ FAILED = "failed"
+ SKIPPED = "skipped"
+ ABORTED = "aborted"
+
+
+class WorkflowStepRecord(BaseModel):
+ kind: str
+ status: str = "pending"
+ message: str = ""
+ started_at: float | None = None
+ completed_at: float | None = None
+ error_detail: str | None = None
+
+
+class QueueItem(BaseModel):
+ item_id: str
+ beamline: str
+ sample_id: int | None = None
+ sample_name: str = ""
+ owner_pgroup: str = ""
+ created_by: str = ""
+ created_at: float = Field(default_factory=time.time)
+ priority: int = 100
+ order_index: int = 0
+ status: QueueItemStatus = QueueItemStatus.PENDING
+ steps: list[WorkflowStepRecord] = Field(default_factory=list)
+ current_step_index: int = 0
+ recipe: dict[str, Any] = Field(default_factory=dict)
+ metadata: dict[str, Any] = Field(default_factory=dict)
+
+
+class RuntimeState(BaseModel):
+ running: bool = False
+ paused: bool = False
+ current_queue_id: str = ""
+ current_item_id: str | None = None
+ current_state: str | None = None
+ current_step_index: int = 0
+ last_error: str | None = None
+ last_update: float = Field(default_factory=time.time)
+
+
+class ControlState(BaseModel):
+ pause_requested: bool = False
+ resume_requested: bool = False
+ abort_requested: bool = False
+ skip_requested: bool = False
+ next_sample_requested: bool = False
+ requested_by: str | None = None
+ requested_at: float | None = None
+
+
+class WorkflowEvent(BaseModel):
+ event_id: str = Field(default_factory=lambda: str(uuid.uuid4()))
+ beamline: str = ""
+ item_id: str | None = None
+ step: str | None = None
+ event_type: str = ""
+ timestamp: float = Field(default_factory=time.time)
+ actor: str = ""
+ message: str = ""
+ payload: dict[str, Any] = Field(default_factory=dict)
+
+
+class CreateQueueItemRequest(BaseModel):
+ sample_id: int | None = None
+ sample_name: str = ""
+ priority: int = 100
+ recipe: dict = Field(default_factory=dict)
+ steps: list[str] | None = None # If None, use default steps
+
+
+class MoveItemRequest(BaseModel):
+ new_order_index: int
+
+
+class QueueListResponse(BaseModel):
+ items: list[QueueItem]
+ total: int
+
+
+class RuntimeResponse(BaseModel):
+ runtime: RuntimeState
+ control: ControlState
+
+
+class ControlActionResponse(BaseModel):
+ ok: bool
+ control: ControlState
+ message: str = ""
+
+
+class StepActionResponse(BaseModel):
+ ok: bool
+ item: QueueItem | None = None
+ step: str | None = None
+ status: str = ""
+ message: str = ""
+
+
+class EventListResponse(BaseModel):
+ events: list[WorkflowEvent]
+
+
+class AutomationStatusResponse(BaseModel):
+ enabled: bool
+ running: bool
+ runtime: RuntimeState
+ control: ControlState
\ No newline at end of file
diff --git a/src/aare/common/automation_queue_manager.py b/src/aare/common/automation_queue_manager.py
new file mode 100644
index 00000000..4b63acdd
--- /dev/null
+++ b/src/aare/common/automation_queue_manager.py
@@ -0,0 +1,245 @@
+from __future__ import annotations
+
+import time
+from typing import Any
+
+import redis
+
+from aare.common.automation_models import (
+ QueueItem,
+ WorkflowEvent,
+ ControlState,
+ RuntimeState,
+ QueueItemStatus,
+ WorkflowStepRecord,
+ WorkflowStateKind,
+)
+
+CURRENT_QUEUE_LIMIT = 576
+
+def build_default_steps() -> list[WorkflowStepRecord]:
+ return [
+ WorkflowStepRecord(kind=WorkflowStateKind.MOUNT.value),
+ WorkflowStepRecord(kind=WorkflowStateKind.LOOP_CENTRE.value),
+ WorkflowStepRecord(kind=WorkflowStateKind.RASTER.value),
+ WorkflowStepRecord(kind=WorkflowStateKind.DATA_COLLECTION.value),
+ ]
+
+
+class WorkflowRedisManager:
+ def __init__(self, client: redis.Redis, beamline: str):
+ self._client = client
+ self._bl = beamline.lower()
+
+ def _key(self, suffix: str) -> str:
+ return f"{self._bl}:workflow:{suffix}"
+
+ def _item_key(self, item_id: str) -> str:
+ return self._key(f"item:{item_id}")
+
+ def _queue_key(self) -> str:
+ return self._key("queue")
+
+ def _runtime_key(self) -> str:
+ return self._key("runtime")
+
+ def _control_key(self) -> str:
+ return self._key("control")
+
+ def _events_key(self) -> str:
+ return self._key("events")
+
+ def _next_item_id(self) -> str:
+ n = int(self._client.incr(self._key("item_seq")))
+ return f"wf_{int(time.time())}_{n:06d}"
+
+ # ─────────────────────────────────────────────
+ # Queue item operations
+ # ─────────────────────────────────────────────
+
+ def create_item(self, item: QueueItem) -> QueueItem:
+ if not item.item_id:
+ item.item_id = self._next_item_id()
+
+ current_count = self._client.zcard(self._queue_key())
+ if current_count >= CURRENT_QUEUE_LIMIT:
+ raise ValueError(f"Queue size limit reached ({CURRENT_QUEUE_LIMIT} items). Please clear the queue.")
+
+ pipe = self._client.pipeline(transaction=True)
+ pipe.set(self._item_key(item.item_id), item.model_dump_json())
+ pipe.zadd(self._queue_key(), {item.item_id: float(item.order_index)})
+ pipe.execute()
+
+ self.append_event(WorkflowEvent(
+ beamline=self._bl,
+ item_id=item.item_id,
+ event_type="item_created",
+ message="Queue item created",
+ payload={"status": item.status.value},
+ ))
+
+ return item
+
+ def get_item(self, item_id: str) -> QueueItem | None:
+ raw = self._client.get(self._item_key(item_id))
+ if raw is None:
+ return None
+ return QueueItem.model_validate_json(raw)
+
+ def update_item(self, item_id: str, patch: dict[str, Any]) -> QueueItem:
+ item = self.get_item(item_id)
+ if item is None:
+ raise KeyError(f"Queue item not found: {item_id}")
+
+ updated = item.model_copy(update=patch)
+ self._client.set(self._item_key(item_id), updated.model_dump_json())
+ return updated
+
+ def delete_item(self, item_id: str) -> None:
+ pipe = self._client.pipeline(transaction=True)
+ pipe.delete(self._item_key(item_id))
+ pipe.zrem(self._queue_key(), item_id)
+ pipe.execute()
+
+ def list_queue_order(self) -> list[str]:
+ return [str(x) for x in self._client.zrange(self._queue_key(), 0, -1)]
+
+ def list_items(self, *, include_finished: bool = True) -> list[QueueItem]:
+ out: list[QueueItem] = []
+ for item_id in self.list_queue_order():
+ item = self.get_item(item_id)
+ if item is None:
+ continue
+ if not include_finished and item.status in {
+ QueueItemStatus.COMPLETED,
+ QueueItemStatus.FAILED,
+ QueueItemStatus.SKIPPED,
+ QueueItemStatus.ABORTED,
+ }:
+ continue
+ out.append(item)
+ return out
+
+ def get_next_pending_item(self) -> QueueItem | None:
+ for item_id in self.list_queue_order():
+ item = self.get_item(item_id)
+ if item is not None and item.status == QueueItemStatus.PENDING:
+ return item
+ return None
+
+ def move_item(self, item_id: str, new_order_index: int) -> None:
+ if self.get_item(item_id) is None:
+ raise KeyError(f"Queue item not found: {item_id}")
+ self._client.zadd(self._queue_key(), {item_id: float(new_order_index)})
+ self.update_item(item_id, {"order_index": new_order_index})
+
+ # ─────────────────────────────────────────────
+ # Runtime state
+ # ─────────────────────────────────────────────
+
+ def get_runtime(self) -> RuntimeState:
+ raw = self._client.get(self._runtime_key())
+ if raw is None:
+ return RuntimeState()
+ return RuntimeState.model_validate_json(raw)
+
+ def set_runtime(self, runtime: RuntimeState) -> RuntimeState:
+ runtime.last_update = time.time()
+ self._client.set(self._runtime_key(), runtime.model_dump_json())
+ return runtime
+
+ def patch_runtime(self, patch: dict[str, Any]) -> RuntimeState:
+ runtime = self.get_runtime()
+ updated = runtime.model_copy(update=patch)
+ return self.set_runtime(updated)
+
+ # ─────────────────────────────────────────────
+ # Control state
+ # ─────────────────────────────────────────────
+
+ def get_control(self) -> ControlState:
+ raw = self._client.get(self._control_key())
+ if raw is None:
+ return ControlState()
+ return ControlState.model_validate_json(raw)
+
+ def request_control(self, patch: dict[str, Any], *, requested_by: str) -> ControlState:
+ control = self.get_control()
+ updated = control.model_copy(update={
+ **patch,
+ "requested_by": requested_by,
+ "requested_at": time.time(),
+ })
+ self._client.set(self._control_key(), updated.model_dump_json())
+
+ self.append_event(WorkflowEvent(
+ beamline=self._bl,
+ item_id=self.get_runtime().current_item_id,
+ event_type="control_requested",
+ actor=requested_by,
+ payload=patch,
+ ))
+
+ return updated
+
+ def clear_control(self) -> ControlState:
+ control = ControlState()
+ self._client.set(self._control_key(), control.model_dump_json())
+ return control
+
+ # ─────────────────────────────────────────────
+ # Event log
+ # ─────────────────────────────────────────────
+
+ def append_event(self, event: WorkflowEvent) -> None:
+ payload = event.model_dump_json()
+ self._client.xadd(self._events_key(), {"json": payload}, maxlen=5000, approximate=True)
+
+ def read_events(self, *, limit: int = 200) -> list[WorkflowEvent]:
+ rows = self._client.xrevrange(self._events_key(), count=limit)
+ out: list[WorkflowEvent] = []
+ for _, fields in reversed(rows):
+ raw = fields.get("json")
+ if raw:
+ out.append(WorkflowEvent.model_validate_json(raw))
+ return out
+
+ # ─────────────────────────────────────────────
+ # Step status updates
+ # ─────────────────────────────────────────────
+
+ def update_step(
+ self,
+ item_id: str,
+ step_index: int,
+ *,
+ status: str | None = None,
+ message: str | None = None,
+ error_detail: str | None = None,
+ ) -> QueueItem:
+ item = self.get_item(item_id)
+ if item is None:
+ raise KeyError(f"Queue item not found: {item_id}")
+
+ if not (0 <= step_index < len(item.steps)):
+ raise IndexError(f"Step index out of range: {step_index}")
+
+ step = item.steps[step_index]
+
+ if status is not None:
+ step.status = status
+ if status == "running" and step.started_at is None:
+ step.started_at = time.time()
+ elif status in ("completed", "failed", "skipped"):
+ step.completed_at = time.time()
+
+ if message is not None:
+ step.message = message
+
+ if error_detail is not None:
+ step.error_detail = error_detail
+
+ item.steps[step_index] = step
+ self._client.set(self._item_key(item_id), item.model_dump_json())
+
+ return item
\ No newline at end of file
diff --git a/src/aare/common/automation_workflow.py b/src/aare/common/automation_workflow.py
new file mode 100644
index 00000000..dacd48d0
--- /dev/null
+++ b/src/aare/common/automation_workflow.py
@@ -0,0 +1,351 @@
+from __future__ import annotations
+
+from abc import ABC, abstractmethod
+from aare.common.automation_models import (
+ WorkflowStateKind,
+ WorkflowMode,
+ StepStatus,
+ StateResult,
+ WorkflowContext,
+ TransitionRule,
+ StateDefinition,
+ )
+
+
+STATE_REGISTRY: dict[WorkflowStateKind, StateDefinition] = {
+ WorkflowStateKind.MOUNT: StateDefinition(
+ kind=WorkflowStateKind.MOUNT,
+ description="Mount the sample",
+ transitions=(
+ TransitionRule(
+ to_state=WorkflowStateKind.LOOP_CENTRE,
+ allowed_modes=frozenset({
+ WorkflowMode.FLEXIBLE_MANUAL,
+ WorkflowMode.GUIDED_MANUAL,
+ WorkflowMode.AUTOMATION,
+ }),
+ ),
+ ),
+ ),
+ WorkflowStateKind.LOOP_CENTRE: StateDefinition(
+ kind=WorkflowStateKind.LOOP_CENTRE,
+ description="Centre the loop",
+ transitions=(
+ TransitionRule(
+ to_state=WorkflowStateKind.RASTER,
+ allowed_modes=frozenset({
+ WorkflowMode.FLEXIBLE_MANUAL,
+ WorkflowMode.GUIDED_MANUAL,
+ WorkflowMode.AUTOMATION,
+ }),
+ ),
+ TransitionRule(
+ to_state=WorkflowStateKind.DATA_COLLECTION,
+ allowed_modes=frozenset({
+ WorkflowMode.FLEXIBLE_MANUAL,
+ WorkflowMode.GUIDED_MANUAL,
+ }),
+ optional=True,
+ ),
+ ),
+ ),
+ WorkflowStateKind.RASTER: StateDefinition(
+ kind=WorkflowStateKind.RASTER,
+ description="Run raster scan",
+ transitions=(
+ TransitionRule(
+ to_state=WorkflowStateKind.DATA_COLLECTION,
+ allowed_modes=frozenset({
+ WorkflowMode.FLEXIBLE_MANUAL,
+ WorkflowMode.GUIDED_MANUAL,
+ WorkflowMode.AUTOMATION,
+ }),
+ ),
+ ),
+ ),
+ WorkflowStateKind.DATA_COLLECTION: StateDefinition(
+ kind=WorkflowStateKind.DATA_COLLECTION,
+ description="Collect diffraction data",
+ transitions=(),
+ ),
+}
+
+
+def get_state_definition(kind: WorkflowStateKind) -> StateDefinition:
+ try:
+ return STATE_REGISTRY[kind]
+ except KeyError as exc:
+ raise KeyError(f"Unknown workflow state: {kind}") from exc
+
+
+def get_allowed_next_states(
+ kind: WorkflowStateKind,
+ mode: WorkflowMode | None = None,
+) -> list[WorkflowStateKind]:
+ definition = get_state_definition(kind)
+ out: list[WorkflowStateKind] = []
+
+ for transition in definition.transitions:
+ if mode is None:
+ out.append(transition.to_state)
+ continue
+
+ if not transition.allowed_modes or mode in transition.allowed_modes:
+ out.append(transition.to_state)
+
+ return out
+
+
+def can_transition(
+ from_state: WorkflowStateKind,
+ to_state: WorkflowStateKind,
+ mode: WorkflowMode | None = None,
+) -> bool:
+ return to_state in get_allowed_next_states(from_state, mode=mode)
+
+
+class StateHandler(ABC):
+ state_kind: WorkflowStateKind
+
+ def __init__(self, registry: dict[WorkflowStateKind, StateDefinition] | None = None):
+ self._registry = registry or STATE_REGISTRY
+
+ def definition(self) -> StateDefinition:
+ return get_state_definition(self.state_kind)
+
+ def can_run(self, context: WorkflowContext) -> bool:
+ return context.current_state in (None, self.state_kind)
+
+ @abstractmethod
+ def validate(self, context: WorkflowContext) -> None:
+ pass
+
+ @abstractmethod
+ def execute(self, context: WorkflowContext) -> StateResult:
+ pass
+
+
+class MountHandler(StateHandler):
+ state_kind = WorkflowStateKind.MOUNT
+
+ def validate(self, context: WorkflowContext) -> None:
+ if context.abort_requested:
+ raise RuntimeError("Abort requested; cannot mount.")
+
+ def execute(self, context: WorkflowContext) -> StateResult:
+ self.validate(context)
+ context.current_state = WorkflowStateKind.MOUNT
+ context.last_message = "Sample mounted"
+ return StateResult(
+ state=self.state_kind,
+ status=StepStatus.SUCCESS,
+ message="Sample mounted successfully.",
+ payload={"mounted": True},
+ )
+
+
+class LoopCentreHandler(StateHandler):
+ state_kind = WorkflowStateKind.LOOP_CENTRE
+
+ def validate(self, context: WorkflowContext) -> None:
+ if context.abort_requested:
+ raise RuntimeError("Abort requested; cannot loop-centre.")
+
+ def execute(self, context: WorkflowContext) -> StateResult:
+ self.validate(context)
+ context.current_state = WorkflowStateKind.LOOP_CENTRE
+ context.last_message = "Loop centred"
+ return StateResult(
+ state=self.state_kind,
+ status=StepStatus.SUCCESS,
+ message="Loop centring completed.",
+ payload={"centred": True},
+ )
+
+
+class RasterHandler(StateHandler):
+ state_kind = WorkflowStateKind.RASTER
+
+ def validate(self, context: WorkflowContext) -> None:
+ if context.abort_requested:
+ raise RuntimeError("Abort requested; cannot raster.")
+
+ def execute(self, context: WorkflowContext) -> StateResult:
+ self.validate(context)
+ context.current_state = WorkflowStateKind.RASTER
+ context.last_message = "Raster completed"
+ return StateResult(
+ state=self.state_kind,
+ status=StepStatus.SUCCESS,
+ message="Raster scan completed.",
+ payload={"best_spot_found": True},
+ )
+
+
+class DataCollectionHandler(StateHandler):
+ state_kind = WorkflowStateKind.DATA_COLLECTION
+
+ def validate(self, context: WorkflowContext) -> None:
+ if context.abort_requested:
+ raise RuntimeError("Abort requested; cannot collect data.")
+
+ def execute(self, context: WorkflowContext) -> StateResult:
+ self.validate(context)
+ context.current_state = WorkflowStateKind.DATA_COLLECTION
+ context.last_message = "Data collected"
+ return StateResult(
+ state=self.state_kind,
+ status=StepStatus.SUCCESS,
+ message="Data collection completed.",
+ payload={"frames_collected": 1},
+ )
+
+
+class WorkflowRunner:
+ """Simple in-memory runner (no persistence)."""
+
+ def __init__(
+ self,
+ registry: dict[WorkflowStateKind, StateDefinition] | None = None,
+ handlers: dict[WorkflowStateKind, StateHandler] | None = None,
+ ):
+ self._registry = registry or STATE_REGISTRY
+ self._handlers = handlers or HANDLER_REGISTRY
+
+ def get_handler(self, state: WorkflowStateKind) -> StateHandler:
+ try:
+ return self._handlers[state]
+ except KeyError as exc:
+ raise KeyError(f"No handler registered for state: {state}") from exc
+
+ def can_move_to(
+ self,
+ current: WorkflowStateKind,
+ next_state: WorkflowStateKind,
+ mode: WorkflowMode,
+ ) -> bool:
+ return can_transition(current, next_state, mode)
+
+ def run_state(
+ self,
+ context: WorkflowContext,
+ state: WorkflowStateKind,
+ ) -> StateResult:
+ if context.current_state is not None:
+ if not self.can_move_to(context.current_state, state, context.mode):
+ raise RuntimeError(
+ f"Transition not allowed: {context.current_state} -> {state}"
+ )
+
+ handler = self.get_handler(state)
+ result = handler.execute(context)
+
+ context.current_state = state
+ context.current_step_index += 1
+ context.last_message = result.message
+
+ return result
+
+
+class SimulatedMountHandler(StateHandler):
+ """Simulated mount handler for testing - doesn't actually mount."""
+ state_kind = WorkflowStateKind.MOUNT
+
+ def validate(self, context: WorkflowContext) -> None:
+ if context.abort_requested:
+ raise RuntimeError("Abort requested; cannot mount.")
+
+ def execute(self, context: WorkflowContext) -> StateResult:
+ self.validate(context)
+ import time
+ time.sleep(10) # Simulate some work
+ context.current_state = WorkflowStateKind.MOUNT
+ context.last_message = "[SIMULATION] Sample would be mounted"
+ return StateResult(
+ state=self.state_kind,
+ status=StepStatus.SUCCESS,
+ message=f"[SIMULATION] Would mount sample_id={context.sample_id}",
+ payload={"mounted": True, "simulated": True},
+ )
+
+
+class SimulatedLoopCentreHandler(StateHandler):
+ """Simulated loop centre handler for testing."""
+ state_kind = WorkflowStateKind.LOOP_CENTRE
+
+ def validate(self, context: WorkflowContext) -> None:
+ if context.abort_requested:
+ raise RuntimeError("Abort requested; cannot loop-centre.")
+
+ def execute(self, context: WorkflowContext) -> StateResult:
+ self.validate(context)
+ import time
+ time.sleep(10)
+ context.current_state = WorkflowStateKind.LOOP_CENTRE
+ context.last_message = "[SIMULATION] Loop would be centred"
+ return StateResult(
+ state=self.state_kind,
+ status=StepStatus.SUCCESS,
+ message="[SIMULATION] Would run loop centering algorithm",
+ payload={"centred": True, "simulated": True},
+ )
+
+
+class SimulatedRasterHandler(StateHandler):
+ """Simulated raster handler for testing."""
+ state_kind = WorkflowStateKind.RASTER
+
+ def validate(self, context: WorkflowContext) -> None:
+ if context.abort_requested:
+ raise RuntimeError("Abort requested; cannot raster.")
+
+ def execute(self, context: WorkflowContext) -> StateResult:
+ self.validate(context)
+ import time
+ time.sleep(10)
+ context.current_state = WorkflowStateKind.RASTER
+ context.last_message = "[SIMULATION] Raster scan would be performed"
+ return StateResult(
+ state=self.state_kind,
+ status=StepStatus.SUCCESS,
+ message="[SIMULATION] Would run raster scan, find best diffraction spot",
+ payload={"best_spot_found": True, "simulated": True},
+ )
+
+
+class SimulatedDataCollectionHandler(StateHandler):
+ """Simulated data collection handler for testing."""
+ state_kind = WorkflowStateKind.DATA_COLLECTION
+
+ def validate(self, context: WorkflowContext) -> None:
+ if context.abort_requested:
+ raise RuntimeError("Abort requested; cannot collect data.")
+
+ def execute(self, context: WorkflowContext) -> StateResult:
+ self.validate(context)
+ import time
+ time.sleep(10)
+ context.current_state = WorkflowStateKind.DATA_COLLECTION
+ context.last_message = "[SIMULATION] Data collection would be performed"
+ return StateResult(
+ state=self.state_kind,
+ status=StepStatus.SUCCESS,
+ message="[SIMULATION] Would collect 1800 frames at 0.2° oscillation",
+ payload={"frames_collected": 1800, "simulated": True},
+ )
+
+
+# Simulated handler registry for testing
+SIMULATED_HANDLER_REGISTRY: dict[WorkflowStateKind, StateHandler] = {
+ WorkflowStateKind.MOUNT: SimulatedMountHandler(STATE_REGISTRY),
+ WorkflowStateKind.LOOP_CENTRE: SimulatedLoopCentreHandler(STATE_REGISTRY),
+ WorkflowStateKind.RASTER: SimulatedRasterHandler(STATE_REGISTRY),
+ WorkflowStateKind.DATA_COLLECTION: SimulatedDataCollectionHandler(STATE_REGISTRY),
+}
+
+HANDLER_REGISTRY: dict[WorkflowStateKind, StateHandler] = {
+ WorkflowStateKind.MOUNT: MountHandler(STATE_REGISTRY),
+ WorkflowStateKind.LOOP_CENTRE: LoopCentreHandler(STATE_REGISTRY),
+ WorkflowStateKind.RASTER: RasterHandler(STATE_REGISTRY),
+ WorkflowStateKind.DATA_COLLECTION: DataCollectionHandler(STATE_REGISTRY),
+}
\ No newline at end of file
diff --git a/src/aare/common/exception_handler.py b/src/aare/common/exception_handler.py
index 7098c362..7012fc78 100644
--- a/src/aare/common/exception_handler.py
+++ b/src/aare/common/exception_handler.py
@@ -1,5 +1,7 @@
from __future__ import annotations
+import time
+
from aare.common.logger_config import setup_logger
from aare.common.error_codes import AuthErrorCode, DAQErrorCode
@@ -104,6 +106,8 @@ class AuthenticationException(Exception):
class UserRightsException(Exception):
+ _last_log_ts_by_message: dict[str, float] = {}
+ _throttle_window_s = 30.0
def __init__(self,
message: str = "User does not have rights to perform this action.",
*,
@@ -115,7 +119,12 @@ class UserRightsException(Exception):
self.status_code = status_code
self.headers = headers
self.code = code
- logger.error(message, extra={"exception:": Exception})
+
+ now = time.monotonic()
+ last_ts = self._last_log_ts_by_message.get(message, 0.0)
+ if (now - last_ts) >= self._throttle_window_s:
+ self._last_log_ts_by_message[message] = now
+ logger.warning(message, extra={"exception:": Exception})
def __str__(self) -> str:
return self.message
@@ -186,5 +195,106 @@ class TellCommunicationError(Exception):
},
)
+ def __str__(self) -> str:
+ return self.message
+
+class JFJochCommunicationError(Exception):
+ """
+ Raised when JFJoch HTTP/API communication fails.
+ Intended for scan-time fallbacks and GUI-visible alerts.
+ """
+
+ def __init__(
+ self,
+ message: str = "JFJoch communication error",
+ *,
+ operation: str | None = None,
+ endpoint: str | None = None,
+ base_url: str | None = None,
+ status_code: int | None = None,
+ ):
+ super().__init__(message)
+ self.message = message
+ self.operation = operation
+ self.endpoint = endpoint
+ self.base_url = base_url
+ self.status_code = status_code
+ logger.error(
+ message,
+ extra={
+ "device": "jfjjoch",
+ "operation": operation,
+ "endpoint": endpoint,
+ "base_url": base_url,
+ "status_code": status_code,
+ },
+ )
+
+ def __str__(self) -> str:
+ return self.message
+
+class AareDBCommunicationError(Exception):
+ """Raised when AareDB HTTPS communication fails."""
+ def __init__(
+ self,
+ message: str = "AareDB communication error",
+ *,
+ operation: str | None = None,
+ endpoint: str | None = None,
+ base_url: str | None = None,
+ status_code: int | None = None,
+ ):
+ super().__init__(message)
+ self.message = message
+ self.operation = operation
+ self.endpoint = endpoint
+ self.base_url = base_url
+ self.status_code = status_code
+ logger.error(
+ message,
+ extra={
+ "device": "aaredb",
+ "operation": operation,
+ "endpoint": endpoint,
+ "base_url": base_url,
+ "status_code": status_code,
+ },
+ )
+
+ def __str__(self) -> str:
+ return self.message
+
+class AerotechCommunicationError(Exception):
+ """
+ Raised when Aerotech HTTP/API communication fails (connection refused, timeout, bad HTTP status, etc).
+ Keep the original exception in `__cause__` by using `raise ... from e`.
+ """
+
+ def __init__(
+ self,
+ message: str = "Aerotech communication error",
+ *,
+ endpoint: str | None = None,
+ base_url: str | None = None,
+ operation: str | None = None,
+ status_code: int | None = None,
+ ):
+ super().__init__(message)
+ self.message = message
+ self.endpoint = endpoint
+ self.base_url = base_url
+ self.operation = operation
+ self.status_code = status_code
+ logger.error(
+ message,
+ extra={
+ "device": "aerotech",
+ "operation": operation,
+ "endpoint": endpoint,
+ "base_url": base_url,
+ "status_code": status_code,
+ },
+ )
+
def __str__(self) -> str:
return self.message
\ No newline at end of file
diff --git a/src/aare/common/models.py b/src/aare/common/models.py
index f90ad14d..1db6f879 100644
--- a/src/aare/common/models.py
+++ b/src/aare/common/models.py
@@ -725,6 +725,9 @@ class DAQStatusModel(BaseModel):
smargon_connected: bool = True
smargon_error: str | None = None
+ aerotech_connected: bool = True
+ aerotech_error: str | None = None
+
class BeamlineSettingsModel(BaseModel):
dtz_max: float | None = 1600.0
dtz_min: float | None = 120.0
diff --git a/src/aare/daq/aaredb.py b/src/aare/daq/aaredb.py
index e165f7b1..2d1b9785 100644
--- a/src/aare/daq/aaredb.py
+++ b/src/aare/daq/aaredb.py
@@ -42,17 +42,28 @@ class AareWrapper:
def __init__(
self,
bl: MXBeamline,
- host: str = "https://mx-db-01.psi.ch/dispatcher",
+ host: str = "https://mx-aaredb-dmz-01.psi.ch/dispatcher",
):
configuration = aareDB.Configuration(host=host)
- configuration.verify_ssl = False # Disable SSL verification
+ # --- mTLS & SSL CONFIGURATION ---
+ # 1. Trust the Server (CA that signed mx-aaredb-dmz-01)
+ configuration.verify_ssl = True
+ configuration.ssl_ca_cert = "/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem"
+
+ # 2. Present Machine Identity (The certs that worked in curl)
+ configuration.cert_file = "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_dmz-01.crt"
+ configuration.key_file = "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_dmz-01.key"
+
+ # 3. Initialize the Client with this config
self.client = aareDB.ApiClient(configuration)
+
+ # Identity Forwarding (Optional now that mTLS is active, but safe to keep)
self.client.default_headers["X-Shared-Password"] = os.getenv("AAREDB_SHARED_PASSWORD")
self.__host = host
self.__tell_api = aareDB.TellsRunnerApi(self.client)
self.__sample_api = aareDB.SamplesRunnerApi(self.client)
- self.__proc_api = aareDB.ProcessingsRunnerApi(self.client)
+ self.__proc_api = aareDB.ProcessingsRunnerApi(self.client)
self.__raster_api = aareDB.GridscanRunnerApi(self.client)
self.__bl = bl
@@ -70,7 +81,7 @@ class AareWrapper:
ret = self.__tell_api.set_tell_positions(
set_tell_position_request=payload,
)
- print(ret)
+ logger.debug(ret)
def create_manual_sample(self, s: SampleShortInfo):
from aareDB.models import ManualSampleCreate
@@ -84,7 +95,7 @@ class AareWrapper:
try:
s.db_id = self.__sample_api.insert_sample(manual_sample).id
except Exception as e:
- print(f"Error inserting sample: {e}")
+ logger.error(f"Error inserting sample: {e}")
def sample_mounted(self, s: Optional[SampleShortInfo]):
if s is not None:
@@ -94,7 +105,7 @@ class AareWrapper:
sample_event_create=SampleEventCreate(event_type=SampleEventType("Mounted")),
)
except Exception as e:
- print(e)
+ logger.error(e)
def sample_unmounted(self, s: Optional[SampleShortInfo]):
if s is not None:
@@ -104,7 +115,7 @@ class AareWrapper:
sample_event_create=SampleEventCreate(event_type=SampleEventType("Unmounted")),
)
except Exception as e:
- print(e)
+ logger.error(e)
def sample_centered(self, s: Optional[SampleShortInfo]):
if s is not None:
@@ -114,7 +125,7 @@ class AareWrapper:
sample_event_create=SampleEventCreate(event_type=SampleEventType("Centered")),
)
except Exception as e:
- print(e)
+ logger.error(e)
def sample_collected(self, s: Optional[SampleShortInfo]):
if s is None:
@@ -125,7 +136,7 @@ class AareWrapper:
sample_event_create=SampleEventCreate(event_type=SampleEventType("Collected")),
)
except Exception as e:
- print(e)
+ logger.error(e)
def sample_failed(self, s: Optional[SampleShortInfo], failed_comment: Optional[str] = None):
if s is None:
@@ -136,7 +147,7 @@ class AareWrapper:
sample_event_create=SampleEventCreate(event_type=SampleEventType("Failed"), comment=failed_comment),
)
except Exception as e:
- print(e)
+ logger.error(e)
def axc_failed(self, s: Optional[SampleShortInfo]):
if s is None:
@@ -147,7 +158,7 @@ class AareWrapper:
sample_event_create=SampleEventCreate(event_type=SampleEventType("AXCFailed")),
)
except Exception as e:
- print(e)
+ logger.error(e)
def alc_failed(self, s: Optional[SampleShortInfo], alc_comment: Optional[str] = None):
if s is None:
@@ -158,7 +169,7 @@ class AareWrapper:
sample_event_create=SampleEventCreate(event_type=SampleEventType("ALCFailed"), comment=alc_comment),
)
except Exception as e:
- print(e)
+ logger.error(e)
def sample_lost(self, s: Optional[SampleShortInfo]):
if s is None:
@@ -288,9 +299,9 @@ class AareWrapper:
sample_id=s.db_id,
experiment_parameters_create=experiment_params_payload
)
- print("Experiment parameters created:", response)
+ logger.debug("Experiment parameters created:", response)
except Exception as e:
- print(e)
+ logger.error(e)
def create_gridscan_run(self, s: Optional[SampleShortInfo], r:RasterGridRequest, d:DAQStatusModel):
if s is None:
@@ -352,9 +363,9 @@ class AareWrapper:
sample_id=s.db_id,
experiment_parameters_create=experiment_params_payload
)
- print("Experiment parameters created:", response)
+ logger.info("Experiment parameters created:", response)
except Exception as e:
- print(e)
+ logger.debug(e)
def ingest_gridscan(self, sample: Optional[SampleShortInfo], raster_result: ScanResult,
raster_request: RasterGridRequest, geom: SampleGeometryModel,
@@ -378,7 +389,7 @@ class AareWrapper:
headers=headers, data=json.dumps(payload), timeout=30, verify=False)
response.raise_for_status()
- print(f"Response status code: {response.status_code}")
+ logger.info(f"Response status code: {response.status_code}")
def format_gridscan_payload(self, sample: Optional[SampleShortInfo], raster_result:ScanResult,
@@ -421,7 +432,7 @@ class AareWrapper:
return payload
except Exception as e:
- print(e)
+ logger.error(e)
raise e
def ingest_scan(self, sample: Optional[SampleShortInfo], result: ScanResult,
@@ -445,7 +456,7 @@ class AareWrapper:
headers=headers, data=json.dumps(payload), timeout=30, verify=False)
response.raise_for_status()
- print(f"Response status code: {response.status_code}")
+ logger.info(f"Response status code: {response.status_code}")
def format_scan_payload(self, sample: Optional[SampleShortInfo], result:ScanResult,
geom:SampleGeometryModel,
@@ -463,5 +474,5 @@ class AareWrapper:
return payload
except Exception as e:
- print(e)
+ logger.error(e)
raise e
diff --git a/src/aare/daq/auth.py b/src/aare/daq/auth.py
index 3af7f247..22704334 100644
--- a/src/aare/daq/auth.py
+++ b/src/aare/daq/auth.py
@@ -1,14 +1,18 @@
import grp
import os
import pwd
+import uuid
from datetime import datetime, timedelta, UTC
from typing import List
+import time
import jwt
-from fastapi import HTTPException, status
-from fastapi.security import OAuth2PasswordRequestForm
+from fastapi import Depends, HTTPException, status
+from fastapi.security import OAuth2PasswordRequestForm, OAuth2PasswordBearer
from pydantic import BaseModel
+from aare.common.auth_models import BatonRequestStatus, BatonTransferQueue, BatonRequest, BatonStatus
+from aare.common.models import SessionsStateEnum
from aare.daq.config import BeamlineConfig
from aare.common.exception_handler import AuthenticationException, UserRightsException, AuthErrorCode
@@ -21,10 +25,13 @@ SECRET_KEY = os.environ.get("JWT_AAREDAQ_KEY")
ALGORITHM = "HS256"
ACCESS_TOKEN_EXPIRE_MINUTES = 24 * 60 * 7 # 1 week
SESSION_EXPIRE_SECONDS = 60 * 10
+BATON_REQUEST_TIMEOUT_SECONDS = 30
STAFF_GROUP = "unx-MXgroup"
SUPER_USERS = ["e10019", "e11206", "e18147"]
+oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token")
+
class TokenData(BaseModel):
sub: str # Username
pgroups: List[str]
@@ -54,7 +61,7 @@ def authenticate_user(cfg: BeamlineConfig, form_data: OAuth2PasswordRequestForm)
return create_access_token(token)
-def parse_token(token: str) -> TokenData:
+def parse_token(token: str = Depends(oauth2_scheme)) -> TokenData:
try:
payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
token = TokenData(**payload)
@@ -112,4 +119,267 @@ def check_jwt_staff(cfg: BeamlineConfig, data: TokenData) -> None:
) from e
def force_current_sesion(cfg: BeamlineConfig, data: TokenData) -> None:
- cfg.force_set_active_session(data.session, SESSION_EXPIRE_SECONDS)
+ cfg.execute_baton_transfer(
+ to_session=data.session,
+ to_username=data.sub,
+ to_is_staff=data.staff,
+ to_pgroup=cfg.pgroup,
+ expiry_sec=SESSION_EXPIRE_SECONDS
+ )
+
+def _finalize_expired_baton_request(cfg: BeamlineConfig, pending: BatonRequest) -> BatonStatus:
+ """
+ Resolve an expired baton request in one place.
+
+ Rules:
+ - staff can override immediately if beamline is free
+ - otherwise the transfer is queued
+ - pending request is cleared once it is no longer pending
+ """
+ requester_is_staff = bool(pending.requester_is_staff)
+
+ if cfg.can_transfer_baton_now():
+ cfg.execute_baton_transfer(
+ to_session=pending.requester_session,
+ to_username=pending.requester_username,
+ to_is_staff=requester_is_staff,
+ to_pgroup=cfg.pgroup,
+ expiry_sec=SESSION_EXPIRE_SECONDS
+ )
+ else:
+ cfg.queued_baton_transfer = BatonTransferQueue(
+ target_session=pending.requester_session,
+ target_username=pending.requester_username,
+ target_is_staff=requester_is_staff,
+ target_pgroup=cfg.pgroup,
+ queued_at=time.time(),
+ reason="timeout_beamline_busy"
+ )
+
+ cfg.clear_pending_baton_request()
+ return get_baton_status(cfg, TokenData(
+ sub=pending.requester_username,
+ pgroups=[],
+ session=pending.requester_session,
+ staff=requester_is_staff
+ ))
+
+def resolve_baton_timeout_if_needed(cfg: BeamlineConfig) -> BatonStatus | None:
+ """
+ Check the current pending request and resolve it if expired.
+ Returns the updated BatonStatus when a timeout was processed, else None.
+ """
+ pending = cfg.pending_baton_request
+ if pending is None or pending.status != BatonRequestStatus.PENDING:
+ return None
+
+ elapsed = time.time() - pending.created_at
+ if elapsed < pending.timeout_seconds:
+ return None
+
+ return _finalize_expired_baton_request(cfg, pending)
+
+def get_baton_status(cfg: BeamlineConfig, data: TokenData) -> BatonStatus:
+ """
+ Build baton status scoped to the requesting session.
+
+ Important:
+ - requester sees you_have_pending_request
+ - holder sees incoming_request
+ - nobody else sees the request as actionable
+ """
+ holder = cfg.baton_holder
+ pending = cfg.pending_baton_request
+
+ is_requester = bool(
+ pending
+ and pending.status in (BatonRequestStatus.PENDING, BatonRequestStatus.REFUSED)
+ and pending.requester_session == data.session
+ )
+
+ is_holder = bool(
+ pending
+ and pending.status == BatonRequestStatus.PENDING
+ and pending.holder_session == data.session
+ )
+
+ # Only expose the pending request object to the two relevant sessions
+ scoped_pending = pending if (is_requester or is_holder) else None
+
+ return BatonStatus(
+ holder=holder,
+ pending_request=scoped_pending,
+ queued_transfer=cfg.queued_baton_transfer,
+ you_are_holder=bool(holder and holder.session == data.session),
+ you_have_pending_request=is_requester,
+ incoming_request=is_holder,
+ allow_non_staff_request=cfg.allow_non_staff_request_from_staff,
+ )
+
+def request_baton(cfg: BeamlineConfig, data: TokenData) -> dict:
+ resolve_baton_timeout_if_needed(cfg)
+
+ session_state = cfg.session_state(data.session)
+
+ if session_state == SessionsStateEnum.Vacant:
+ cfg.execute_baton_transfer(
+ to_session=data.session,
+ to_username=data.sub,
+ to_is_staff=data.staff,
+ to_pgroup=cfg.pgroup,
+ expiry_sec=SESSION_EXPIRE_SECONDS,
+ )
+ return {"granted": True, "message": "Baton acquired (beamline was vacant)"}
+
+ if session_state == SessionsStateEnum.OwnedByYou:
+ cfg.try_set_active_session(data.session, SESSION_EXPIRE_SECONDS)
+ return {"already_holder": True, "message": "You already hold the baton"}
+
+ holder = cfg.baton_holder
+ print(cfg.allow_non_staff_request_from_staff)
+ if holder and holder.is_staff and not data.staff and not cfg.allow_non_staff_request_from_staff:
+ return {
+ "error": True,
+ "message": "Requesting baton from staff is disabled by backend policy.",
+ }
+
+ if data.staff:
+ if not cfg.can_transfer_baton_now():
+ cfg.queued_baton_transfer = BatonTransferQueue(
+ target_session=data.session,
+ target_username=data.sub,
+ target_is_staff=data.staff,
+ target_pgroup=cfg.pgroup,
+ queued_at=time.time(),
+ reason="beamline_busy_staff_override",
+ )
+ return {
+ "queued": True,
+ "message": "Staff override queued - will transfer when beamline is available",
+ }
+
+ cfg.execute_baton_transfer(
+ to_session=data.session,
+ to_username=data.sub,
+ to_is_staff=data.staff,
+ to_pgroup=cfg.pgroup,
+ expiry_sec=SESSION_EXPIRE_SECONDS,
+ )
+ return {"granted": True, "override": True, "message": "Staff override - baton acquired"}
+
+ existing_request = cfg.pending_baton_request
+ if existing_request and existing_request.status == BatonRequestStatus.PENDING:
+ if existing_request.requester_session == data.session:
+ elapsed = time.time() - existing_request.created_at
+ if elapsed >= existing_request.timeout_seconds:
+ return {"timeout": True, "message": "Request timed out"}
+ remaining = existing_request.timeout_seconds - elapsed
+ return {
+ "pending": True,
+ "existing": True,
+ "remaining_seconds": max(0, remaining),
+ "message": f"Request already pending ({remaining:.0f}s remaining)",
+ }
+ return {
+ "error": True,
+ "message": "Another user already has a pending request",
+ }
+
+ request = BatonRequest(
+ request_id=str(uuid.uuid4()),
+ requester_username=data.sub,
+ requester_session=data.session,
+ requester_is_staff=data.staff,
+ holder_username=holder.username if holder else None,
+ holder_session=holder.session if holder else None,
+ created_at=time.time(),
+ timeout_seconds=BATON_REQUEST_TIMEOUT_SECONDS,
+ status=BatonRequestStatus.PENDING,
+ )
+ cfg.set_pending_baton_request(request, timeout_sec=BATON_REQUEST_TIMEOUT_SECONDS)
+
+ return {
+ "pending": True,
+ "request_id": request.request_id,
+ "timeout_seconds": BATON_REQUEST_TIMEOUT_SECONDS,
+ "message": f"Request sent to {holder.username if holder else 'current holder'}",
+ }
+
+def respond_to_baton_request(cfg: BeamlineConfig, data: TokenData, accept: bool) -> dict:
+ """
+ Current baton holder responds to a pending request.
+ """
+ resolve_baton_timeout_if_needed(cfg)
+
+ holder = cfg.baton_holder
+ if holder is None or holder.session != data.session:
+ return {"error": True, "message": "You are not the current baton holder"}
+
+ pending = cfg.pending_baton_request
+ if pending is None or pending.status != BatonRequestStatus.PENDING:
+ return {"error": True, "message": "No pending request to respond to"}
+
+ if accept:
+ if cfg.can_transfer_baton_now():
+ cfg.execute_baton_transfer(
+ to_session=pending.requester_session,
+ to_username=pending.requester_username,
+ to_is_staff=pending.requester_is_staff,
+ to_pgroup=cfg.pgroup,
+ expiry_sec=SESSION_EXPIRE_SECONDS
+ )
+ return {"accepted": True, "transferred": True, "message": "Baton transferred"}
+ else:
+ cfg.queued_baton_transfer = BatonTransferQueue(
+ target_session=pending.requester_session,
+ target_username=pending.requester_username,
+ target_is_staff=pending.requester_is_staff,
+ target_pgroup=cfg.pgroup,
+ queued_at=time.time(),
+ reason="accepted_beamline_busy"
+ )
+ cfg.clear_pending_baton_request()
+ return {
+ "accepted": True,
+ "queued": True,
+ "message": "Request accepted - will transfer when beamline is available"
+ }
+ else:
+ pending.status = BatonRequestStatus.REFUSED
+ cfg.set_pending_baton_request(pending, timeout_sec=5)
+ return {"refused": True, "message": "Request refused"}
+
+def release_baton(cfg: BeamlineConfig, data: TokenData) -> dict:
+ """
+ Voluntarily release the baton (set session to free).
+ """
+ resolve_baton_timeout_if_needed(cfg)
+
+ holder = cfg.baton_holder
+ if holder is None:
+ return {"released": True, "message": "Baton was already vacant"}
+
+ if holder.session != data.session:
+ return {"info": True, "message": "You don't hold the baton"}
+
+ cfg.end_active_session(data.session)
+ cfg.baton_holder = None
+ cfg.clear_pending_baton_request()
+
+ return {"released": True, "message": "Baton released - beamline is now vacant"}
+
+def cancel_baton_request(cfg: BeamlineConfig, data: TokenData) -> dict:
+ """
+ Cancel your own pending baton request.
+ """
+ resolve_baton_timeout_if_needed(cfg)
+
+ pending = cfg.pending_baton_request
+ if pending is None:
+ return {"error": True, "message": "No pending request to cancel"}
+
+ if pending.requester_session != data.session:
+ return {"error": True, "message": "You can only cancel your own request"}
+
+ cfg.clear_pending_baton_request()
+ return {"cancelled": True, "message": "Request cancelled"}
\ No newline at end of file
diff --git a/src/aare/daq/automation_api_router.py b/src/aare/daq/automation_api_router.py
new file mode 100644
index 00000000..bc4f67eb
--- /dev/null
+++ b/src/aare/daq/automation_api_router.py
@@ -0,0 +1,761 @@
+from __future__ import annotations
+
+import asyncio
+import json
+from typing import AsyncGenerator, TYPE_CHECKING
+
+from fastapi import APIRouter, Depends, HTTPException, status
+from pydantic import BaseModel, Field
+from starlette.responses import StreamingResponse
+
+from aare.common.automation_models import (
+ QueueItem,
+ QueueItemStatus,
+ WorkflowContext,
+ WorkflowEvent,
+ WorkflowMode,
+ WorkflowStepRecord,
+ RuntimeState,
+ ControlState,
+ WorkflowStateKind, EventListResponse, StepActionResponse, ControlActionResponse, RuntimeResponse,
+ CreateQueueItemRequest, QueueListResponse, MoveItemRequest, AutomationStatusResponse,
+)
+from aare.common.automation_queue_manager import (
+ WorkflowRedisManager,
+ build_default_steps,
+)
+from aare.common.automation_workflow import STATE_REGISTRY, HANDLER_REGISTRY
+from aare.daq.automation_runner import PersistentWorkflowRunner
+from aare.daq.auth import parse_token, check_jwt_rw, check_jwt_ro, oauth2_scheme
+from aare.common.models import TokenData
+from aare.daq.config import BeamlineConfig
+
+from aare.daq.automation_runner import AutomationLoop
+
+router = APIRouter(prefix="/workflow", tags=["workflow"])
+
+# ─────────────────────────────────────────────
+# Dependency: get managers
+# ─────────────────────────────────────────────
+
+# These will be set up when the router is included
+_redis_manager: WorkflowRedisManager | None = None
+_runner: PersistentWorkflowRunner | None = None
+_cfg: BeamlineConfig | None = None
+_automation_loop: AutomationLoop | None = None
+
+
+def set_workflow_dependencies(
+ redis_manager: WorkflowRedisManager,
+ runner: PersistentWorkflowRunner,
+ cfg: BeamlineConfig,
+) -> None:
+ global _redis_manager, _runner, _cfg, _automation_loop
+ _redis_manager = redis_manager
+ _runner = runner
+ _cfg = cfg
+ _automation_loop = AutomationLoop(runner, redis_manager)
+
+
+def get_automation_loop() -> AutomationLoop:
+ if _automation_loop is None:
+ raise HTTPException(
+ status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
+ detail="Automation loop not initialized",
+ )
+ return _automation_loop
+
+
+def get_cfg() -> BeamlineConfig:
+ if _cfg is None:
+ raise HTTPException(
+ status_code=503,
+ detail="Workflow system not initialized",
+ )
+ return _cfg
+
+
+def get_redis_manager() -> WorkflowRedisManager:
+ if _redis_manager is None:
+ raise HTTPException(
+ status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
+ detail="Workflow system not initialized",
+ )
+ return _redis_manager
+
+
+def get_runner() -> PersistentWorkflowRunner:
+ if _runner is None:
+ raise HTTPException(
+ status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
+ detail="Workflow runner not initialized",
+ )
+ return _runner
+
+
+# ─────────────────────────────────────────────
+# Queue management endpoints
+# ─────────────────────────────────────────────
+
+
+@router.get("/queue", response_model=QueueListResponse)
+async def list_queue(
+ include_finished: bool = False,
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+):
+ """List all items in the workflow queue."""
+ data = parse_token(token)
+ check_jwt_ro(get_cfg(), data)
+ items = redis_mgr.list_items(include_finished=include_finished)
+
+ # Filter by pgroup if not staff
+ if not data.staff:
+ items = [i for i in items if i.owner_pgroup in data.pgroups]
+
+ return QueueListResponse(items=items, total=len(items))
+
+
+@router.post("/queue", response_model=QueueItem)
+async def create_queue_item(
+ request: CreateQueueItemRequest,
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+):
+ """Add a new item to the workflow queue."""
+ data = parse_token(token)
+ check_jwt_rw(get_cfg(), data)
+
+ # Build steps
+ if request.steps:
+ steps = [
+ WorkflowStepRecord(kind=s)
+ for s in request.steps
+ if s in [sk.value for sk in WorkflowStateKind]
+ ]
+ else:
+ steps = build_default_steps()
+
+ item = QueueItem(
+ item_id="",
+ beamline=redis_mgr._bl,
+ sample_id=request.sample_id,
+ sample_name=request.sample_name,
+ owner_pgroup=data.pgroups[0] if data.pgroups else "",
+ created_by=data.sub,
+ priority=request.priority,
+ order_index=int(asyncio.get_event_loop().time() * 1000),
+ steps=steps,
+ recipe=request.recipe,
+ )
+
+ try:
+ item = redis_mgr.create_item(item)
+ except ValueError as e:
+ raise HTTPException(status_code=400, detail=str(e))
+
+ return item
+
+
+@router.get("/queue/{item_id}", response_model=QueueItem)
+async def get_queue_item(
+ item_id: str,
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+):
+ """Get a single queue item by ID."""
+ data = parse_token(token)
+ check_jwt_ro(get_cfg(), data)
+
+ item = redis_mgr.get_item(item_id)
+ if item is None:
+ raise HTTPException(status_code=404, detail="Item not found")
+
+ # Check access
+ if not data.staff and item.owner_pgroup not in data.pgroups:
+ raise HTTPException(status_code=403, detail="Access denied")
+
+ return item
+
+
+@router.delete("/queue/{item_id}")
+async def delete_queue_item(
+ item_id: str,
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+):
+ """Remove an item from the queue."""
+ data = parse_token(token)
+ check_jwt_rw(get_cfg(), data)
+
+ item = redis_mgr.get_item(item_id)
+ if item is None:
+ raise HTTPException(status_code=404, detail="Item not found")
+
+ # Check access
+ if not data.staff and item.owner_pgroup not in data.pgroups:
+ raise HTTPException(status_code=403, detail="Access denied")
+
+ # Don't allow deleting running items
+ if item.status == QueueItemStatus.RUNNING:
+ raise HTTPException(status_code=409, detail="Cannot delete running item")
+
+ redis_mgr.delete_item(item_id)
+ return {"ok": True, "message": f"Item {item_id} deleted"}
+
+@router.delete("/queue/clear")
+async def clear_queue(
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+):
+ """Clear all non-running items from the queue."""
+ data = parse_token(token)
+ check_jwt_rw(get_cfg(), data)
+
+ items = redis_mgr.list_items(include_finished=True)
+ deleted_count = 0
+
+ for item in items:
+ # Skip running items
+ if item.status == QueueItemStatus.RUNNING:
+ continue
+
+ # Check access
+ if not data.staff and item.owner_pgroup not in data.pgroups:
+ continue
+
+ redis_mgr.delete_item(item.item_id)
+ deleted_count += 1
+
+ return {"ok": True, "message": f"Cleared {deleted_count} items from queue"}
+
+
+@router.post("/control/skip_sample", response_model=ControlActionResponse)
+async def request_skip_sample(
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+ runner: PersistentWorkflowRunner = Depends(get_runner),
+):
+ """Skip the current sample entirely (abort current sample and move to next)."""
+ data = parse_token(token)
+ check_jwt_rw(get_cfg(), data)
+
+ runtime = redis_mgr.get_runtime()
+
+ if not runtime.running or not runtime.current_item_id:
+ return ControlActionResponse(
+ ok=False,
+ control=redis_mgr.get_control(),
+ message="No sample currently running",
+ )
+
+ # Mark current item as skipped and complete it
+ runner.complete_item(runtime.current_item_id, QueueItemStatus.SKIPPED)
+
+ redis_mgr.append_event(WorkflowEvent(
+ beamline=runner.beamline,
+ item_id=runtime.current_item_id,
+ event_type="sample_skipped",
+ actor=data.sub,
+ message="Sample skipped by user",
+ ))
+
+ return ControlActionResponse(
+ ok=True,
+ control=redis_mgr.get_control(),
+ message="Sample skipped",
+ )
+
+@router.post("/queue/{item_id}/move")
+async def move_queue_item(
+ item_id: str,
+ request: MoveItemRequest,
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+):
+ """Reorder an item in the queue."""
+ data = parse_token(token)
+ check_jwt_rw(get_cfg(), data)
+
+ item = redis_mgr.get_item(item_id)
+ if item is None:
+ raise HTTPException(status_code=404, detail="Item not found")
+
+ if not data.staff and item.owner_pgroup not in data.pgroups:
+ raise HTTPException(status_code=403, detail="Access denied")
+
+ redis_mgr.move_item(item_id, request.new_order_index)
+ return {"ok": True, "message": f"Item {item_id} moved"}
+
+
+# ─────────────────────────────────────────────
+# Runtime and control endpoints
+# ─────────────────────────────────────────────
+
+
+@router.get("/runtime", response_model=RuntimeResponse)
+async def get_runtime(
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+):
+ """Get current runtime and control state."""
+ data = parse_token(token)
+ check_jwt_ro(get_cfg(), data)
+
+ return RuntimeResponse(
+ runtime=redis_mgr.get_runtime(),
+ control=redis_mgr.get_control(),
+ )
+
+
+@router.get("/control", response_model=ControlState)
+async def get_control(
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+):
+ """Get current control state."""
+ data = parse_token(token)
+ check_jwt_ro(get_cfg(), data)
+ return redis_mgr.get_control()
+
+
+@router.post("/control/pause", response_model=ControlActionResponse)
+async def request_pause(
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+):
+ """Request workflow pause after current step."""
+ data = parse_token(token)
+ check_jwt_rw(get_cfg(), data)
+
+ control = redis_mgr.request_control(
+ {"pause_requested": True, "resume_requested": False},
+ requested_by=data.sub,
+ )
+
+ return ControlActionResponse(
+ ok=True,
+ control=control,
+ message="Pause requested",
+ )
+
+
+@router.post("/control/resume", response_model=ControlActionResponse)
+async def request_resume(
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+):
+ """Resume paused workflow."""
+ data = parse_token(token)
+ check_jwt_rw(get_cfg(), data)
+
+ control = redis_mgr.request_control(
+ {"pause_requested": False, "resume_requested": True},
+ requested_by=data.sub,
+ )
+
+ return ControlActionResponse(
+ ok=True,
+ control=control,
+ message="Resume requested",
+ )
+
+
+@router.post("/control/abort", response_model=ControlActionResponse)
+async def request_abort(
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+):
+ """Abort current workflow execution."""
+ data = parse_token(token)
+ check_jwt_rw(get_cfg(), data)
+
+ # Clear pause_requested when aborting
+ control = redis_mgr.request_control(
+ {"abort_requested": True, "pause_requested": False, "resume_requested": False},
+ requested_by=data.sub,
+ )
+
+ return ControlActionResponse(
+ ok=True,
+ control=control,
+ message="Abort requested",
+ )
+
+
+@router.post("/control/skip", response_model=ControlActionResponse)
+async def request_skip(
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+):
+ """Skip current step."""
+ data = parse_token(token)
+ check_jwt_rw(get_cfg(), data)
+
+ control = redis_mgr.request_control(
+ {"skip_requested": True},
+ requested_by=data.sub,
+ )
+
+ return ControlActionResponse(
+ ok=True,
+ control=control,
+ message="Skip requested",
+ )
+
+
+@router.post("/control/clear", response_model=ControlActionResponse)
+async def clear_control(
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+):
+ """Clear all control requests."""
+ data = parse_token(token)
+ check_jwt_rw(get_cfg(), data)
+
+ control = redis_mgr.clear_control()
+
+ return ControlActionResponse(
+ ok=True,
+ control=control,
+ message="Control state cleared",
+ )
+
+
+# ─────────────────────────────────────────────
+# Execution endpoints
+# ─────────────────────────────────────────────
+
+
+@router.post("/start/{item_id}", response_model=StepActionResponse)
+async def start_item(
+ item_id: str,
+ mode: WorkflowMode = WorkflowMode.GUIDED_MANUAL,
+ token: str = Depends(oauth2_scheme),
+ runner: PersistentWorkflowRunner = Depends(get_runner),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+):
+ """Start processing a queue item."""
+ data = parse_token(token)
+ check_jwt_rw(get_cfg(), data)
+
+ item = redis_mgr.get_item(item_id)
+ if item is None:
+ raise HTTPException(status_code=404, detail="Item not found")
+
+ if not data.staff and item.owner_pgroup not in data.pgroups:
+ raise HTTPException(status_code=403, detail="Access denied")
+
+ if item.status == QueueItemStatus.RUNNING:
+ raise HTTPException(status_code=409, detail="Item already running")
+
+ # Check if another item is running
+ runtime = redis_mgr.get_runtime()
+ if runtime.running:
+ raise HTTPException(
+ status_code=409,
+ detail=f"Another item is running: {runtime.current_item_id}",
+ )
+
+ # Clear any stale control requests
+ redis_mgr.clear_control()
+
+ # Start the item
+ item = runner.start_item(item_id)
+
+ return StepActionResponse(
+ ok=True,
+ item=item,
+ step=item.steps[0].kind if item.steps else None,
+ status="started",
+ message=f"Started processing {item.sample_name}",
+ )
+
+
+@router.post("/next", response_model=StepActionResponse)
+async def run_next_step(
+ token: str = Depends(oauth2_scheme),
+ runner: PersistentWorkflowRunner = Depends(get_runner),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+):
+ """
+ Run the next step in guided mode.
+
+ This is the main endpoint for guided manual operation.
+ User clicks "Next" and this executes one step.
+ """
+ data = parse_token(token)
+ check_jwt_rw(get_cfg(), data)
+
+ runtime = redis_mgr.get_runtime()
+
+ if not runtime.running or not runtime.current_item_id:
+ raise HTTPException(status_code=409, detail="No item is currently running")
+
+ item = redis_mgr.get_item(runtime.current_item_id)
+ if item is None:
+ raise HTTPException(status_code=404, detail="Running item not found")
+
+ if not data.staff and item.owner_pgroup not in data.pgroups:
+ raise HTTPException(status_code=403, detail="Access denied")
+
+ # Check if all steps completed
+ if item.current_step_index >= len(item.steps):
+ item = runner.complete_item(item.item_id, QueueItemStatus.COMPLETED)
+ return StepActionResponse(
+ ok=True,
+ item=item,
+ status="completed",
+ message="All steps completed",
+ )
+
+ # Check for abort
+ if runner.should_abort():
+ item = runner.complete_item(item.item_id, QueueItemStatus.ABORTED)
+ redis_mgr.clear_control()
+ return StepActionResponse(
+ ok=True,
+ item=item,
+ status="aborted",
+ message="Workflow aborted by user",
+ )
+
+ # Check for skip
+ control = runner.check_control()
+ if control.skip_requested:
+ step_index = item.current_step_index
+ step_kind = item.steps[step_index].kind
+
+ redis_mgr.update_step(item.item_id, step_index, status="skipped")
+ redis_mgr.update_item(item.item_id, {"current_step_index": step_index + 1})
+ redis_mgr.request_control({"skip_requested": False}, requested_by=data.sub)
+
+ redis_mgr.append_event(WorkflowEvent(
+ beamline=runner.beamline,
+ item_id=item.item_id,
+ step=step_kind,
+ event_type="step_skipped",
+ actor=data.sub,
+ message=f"Step {step_kind} skipped by user",
+ ))
+
+ item = redis_mgr.get_item(item.item_id)
+ return StepActionResponse(
+ ok=True,
+ item=item,
+ step=step_kind,
+ status="skipped",
+ message=f"Skipped {step_kind}",
+ )
+
+ # Build context
+ context = WorkflowContext(
+ mode=WorkflowMode.GUIDED_MANUAL,
+ queue_id="default",
+ item_id=item.item_id,
+ sample_id=item.sample_id,
+ current_state=WorkflowStateKind(runtime.current_state) if runtime.current_state else None,
+ current_step_index=item.current_step_index,
+ )
+
+ # Run the step
+ try:
+ result = runner.run_current_step(context)
+ except Exception as e:
+ return StepActionResponse(
+ ok=False,
+ item=redis_mgr.get_item(item.item_id),
+ step=item.steps[item.current_step_index].kind,
+ status="failed",
+ message=str(e),
+ )
+
+ # Get updated item
+ item = redis_mgr.get_item(item.item_id)
+
+ # Check if completed
+ if item.current_step_index >= len(item.steps):
+ item = runner.complete_item(item.item_id, QueueItemStatus.COMPLETED)
+ return StepActionResponse(
+ ok=True,
+ item=item,
+ step=result.state.value,
+ status="completed",
+ message="All steps completed",
+ )
+
+ return StepActionResponse(
+ ok=True,
+ item=item,
+ step=result.state.value,
+ status=result.status.value,
+ message=result.message,
+ )
+
+
+# ─────────────────────────────────────────────
+# Events endpoints
+# ─────────────────────────────────────────────
+
+
+@router.get("/events", response_model=EventListResponse)
+async def get_events(
+ limit: int = 100,
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+):
+ """Get recent workflow events."""
+ data = parse_token(token)
+ check_jwt_ro(get_cfg(), data)
+
+ events = redis_mgr.read_events(limit=limit)
+ return EventListResponse(events=events)
+
+
+# ─────────────────────────────────────────────
+# SSE stream
+# ─────────────────────────────────────────────
+
+
+async def workflow_event_stream(
+ redis_mgr: WorkflowRedisManager,
+) -> AsyncGenerator[str, None]:
+ """
+ Server-Sent Events stream for live workflow updates.
+
+ Polls runtime state and streams changes.
+ """
+ last_runtime_json = ""
+ last_control_json = ""
+ last_event_id = "0-0"
+
+ try:
+ while True:
+ # Check runtime state
+ runtime = redis_mgr.get_runtime()
+ runtime_json = runtime.model_dump_json()
+
+ if runtime_json != last_runtime_json:
+ last_runtime_json = runtime_json
+ yield f"event: runtime\ndata: {runtime_json}\n\n"
+
+ # Check control state
+ control = redis_mgr.get_control()
+ control_json = control.model_dump_json()
+
+ if control_json != last_control_json:
+ last_control_json = control_json
+ yield f"event: control\ndata: {control_json}\n\n"
+
+ # Check for new events (using Redis streams)
+ try:
+ events_key = redis_mgr._events_key()
+ new_events = redis_mgr._client.xread(
+ {events_key: last_event_id},
+ count=10,
+ block=0,
+ )
+
+ if new_events:
+ for _, messages in new_events:
+ for msg_id, fields in messages:
+ last_event_id = msg_id
+ raw = fields.get("json")
+ if raw:
+ yield f"event: workflow_event\ndata: {raw}\n\n"
+ except Exception:
+ pass # Redis stream read failed, continue polling
+
+ await asyncio.sleep(0.2)
+
+ except asyncio.CancelledError:
+ return
+
+# ─────────────────────────────────────────────
+# Automation mode endpoints
+# ─────────────────────────────────────────────
+
+
+@router.get("/automation/status", response_model=AutomationStatusResponse)
+async def get_automation_status(
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+ loop: "AutomationLoop" = Depends(get_automation_loop),
+):
+ """Get automation loop status."""
+ data = parse_token(token)
+ check_jwt_ro(get_cfg(), data)
+
+ return AutomationStatusResponse(
+ enabled=loop.is_enabled,
+ running=loop.is_running,
+ runtime=redis_mgr.get_runtime(),
+ control=redis_mgr.get_control(),
+ )
+
+
+@router.post("/automation/start", response_model=AutomationStatusResponse)
+async def start_automation(
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+ loop: "AutomationLoop" = Depends(get_automation_loop),
+):
+ """Start automation mode - processes queue automatically."""
+ data = parse_token(token)
+ check_jwt_rw(get_cfg(), data)
+
+ loop.start()
+
+ return AutomationStatusResponse(
+ enabled=loop.is_enabled,
+ running=loop.is_running,
+ runtime=redis_mgr.get_runtime(),
+ control=redis_mgr.get_control(),
+ )
+
+
+@router.post("/automation/stop", response_model=AutomationStatusResponse)
+async def stop_automation(
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+ loop: "AutomationLoop" = Depends(get_automation_loop),
+):
+ """Stop automation mode - completes current step then stops."""
+ data = parse_token(token)
+ check_jwt_rw(get_cfg(), data)
+
+ loop.stop()
+
+ return AutomationStatusResponse(
+ enabled=loop.is_enabled,
+ running=loop.is_running,
+ runtime=redis_mgr.get_runtime(),
+ control=redis_mgr.get_control(),
+ )
+
+@router.get("/sse")
+async def workflow_sse(
+ token: str = Depends(oauth2_scheme),
+ redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
+):
+ """
+ SSE endpoint for live workflow updates.
+
+ Events:
+ - runtime: RuntimeState changes
+ - control: ControlState changes
+ - workflow_event: Individual workflow events
+ """
+ data = parse_token(token)
+ check_jwt_ro(get_cfg(), data)
+
+ return StreamingResponse(
+ workflow_event_stream(redis_mgr),
+ media_type="text/event-stream",
+ headers={
+ "Cache-Control": "no-cache",
+ "Connection": "keep-alive",
+ "Access-Control-Allow-Origin": "*",
+ "Access-Control-Allow-Headers": "Cache-Control",
+ },
+ )
\ No newline at end of file
diff --git a/src/aare/daq/automation_runner.py b/src/aare/daq/automation_runner.py
new file mode 100644
index 00000000..8b672235
--- /dev/null
+++ b/src/aare/daq/automation_runner.py
@@ -0,0 +1,461 @@
+import asyncio
+from typing import Callable
+
+import redis
+
+from aare.common.automation_models import (
+ WorkflowEvent,
+ QueueItemStatus,
+ QueueItem,
+ ControlState,
+ WorkflowStateKind,
+ WorkflowContext,
+ StateResult,
+ StateDefinition,
+ WorkflowMode,
+)
+from aare.common.automation_queue_manager import (
+ WorkflowRedisManager,
+ build_default_steps
+)
+
+from aare.common.automation_workflow import (
+ can_transition,
+ StateHandler,
+ STATE_REGISTRY,
+ HANDLER_REGISTRY
+)
+
+
+class PersistentWorkflowRunner:
+ """
+ Orchestrates workflow execution with Redis persistence.
+
+ Combines:
+ - State handlers (from automation_workflow)
+ - Redis persistence (from automation_queue_manager)
+ """
+
+ def __init__(
+ self,
+ redis_manager: WorkflowRedisManager,
+ registry: dict[WorkflowStateKind, StateDefinition] | None = None,
+ handlers: dict[WorkflowStateKind, StateHandler] | None = None,
+ ):
+ self._redis = redis_manager
+ self._registry = registry or STATE_REGISTRY
+ self._handlers = handlers or HANDLER_REGISTRY
+
+ @property
+ def beamline(self) -> str:
+ return self._redis._bl
+
+ def get_handler(self, state: WorkflowStateKind) -> StateHandler:
+ try:
+ return self._handlers[state]
+ except KeyError as exc:
+ raise KeyError(f"No handler registered for state: {state}") from exc
+
+ def start_item(self, item_id: str) -> QueueItem:
+ item = self._redis.get_item(item_id)
+ if item is None:
+ raise KeyError(f"Item not found: {item_id}")
+
+ item = self._redis.update_item(item_id, {
+ "status": QueueItemStatus.RUNNING,
+ "current_step_index": 0,
+ })
+
+ self._redis.patch_runtime({
+ "running": True,
+ "paused": False,
+ "current_item_id": item_id,
+ "current_step_index": 0,
+ "current_state": None,
+ "last_error": None,
+ })
+
+ self._redis.append_event(WorkflowEvent(
+ beamline=self.beamline,
+ item_id=item_id,
+ event_type="item_started",
+ message=f"Started processing {item.sample_name}",
+ ))
+
+ return item
+
+ def run_current_step(self, context: WorkflowContext) -> StateResult:
+ item = self._redis.get_item(context.item_id)
+ if item is None:
+ raise KeyError(f"Item not found: {context.item_id}")
+
+ step_index = item.current_step_index
+ if step_index >= len(item.steps):
+ raise RuntimeError("No more steps to run")
+
+ step_record = item.steps[step_index]
+ state_kind = WorkflowStateKind(step_record.kind)
+
+ # Check transition is allowed
+ if context.current_state is not None:
+ if not can_transition(context.current_state, state_kind, context.mode):
+ raise RuntimeError(
+ f"Transition not allowed: {context.current_state} -> {state_kind}"
+ )
+
+ # Mark step as running
+ self._redis.update_step(context.item_id, step_index, status="running")
+ self._redis.patch_runtime({
+ "current_state": state_kind.value,
+ "current_step_index": step_index,
+ })
+
+ self._redis.append_event(WorkflowEvent(
+ beamline=self.beamline,
+ item_id=context.item_id,
+ step=state_kind.value,
+ event_type="step_started",
+ message=f"Starting {state_kind.value}",
+ ))
+
+ # Execute handler
+ handler = self.get_handler(state_kind)
+
+ try:
+ result = handler.execute(context)
+ except Exception as e:
+ self._redis.update_step(
+ context.item_id,
+ step_index,
+ status="failed",
+ error_detail=str(e),
+ )
+ self._redis.patch_runtime({"last_error": str(e)})
+ self._redis.append_event(WorkflowEvent(
+ beamline=self.beamline,
+ item_id=context.item_id,
+ step=state_kind.value,
+ event_type="step_failed",
+ message=str(e),
+ ))
+ raise
+
+ # Mark step as completed
+ self._redis.update_step(
+ context.item_id,
+ step_index,
+ status=result.status.value,
+ message=result.message,
+ )
+
+ # Advance step index
+ self._redis.update_item(context.item_id, {
+ "current_step_index": step_index + 1,
+ })
+
+ context.current_state = state_kind
+ context.current_step_index = step_index + 1
+
+ self._redis.append_event(WorkflowEvent(
+ beamline=self.beamline,
+ item_id=context.item_id,
+ step=state_kind.value,
+ event_type="step_completed",
+ message=result.message,
+ payload=result.payload,
+ ))
+
+ return result
+
+ def complete_item(self, item_id: str, status: QueueItemStatus) -> QueueItem:
+ item = self._redis.update_item(item_id, {"status": status})
+
+ self._redis.patch_runtime({
+ "running": False,
+ "current_item_id": None,
+ "current_state": None,
+ })
+
+ self._redis.append_event(WorkflowEvent(
+ beamline=self.beamline,
+ item_id=item_id,
+ event_type="item_completed",
+ message=f"Item finished with status {status.value}",
+ ))
+
+ if status in (QueueItemStatus.COMPLETED, QueueItemStatus.ABORTED, QueueItemStatus.SKIPPED):
+ self._redis.delete_item(item_id)
+ return item
+
+ return item
+
+ def check_control(self) -> ControlState:
+ return self._redis.get_control()
+
+ def should_pause(self) -> bool:
+ return self.check_control().pause_requested
+
+ def should_abort(self) -> bool:
+ return self.check_control().abort_requested
+
+ def should_skip(self) -> bool:
+ return self.check_control().skip_requested
+
+
+class AutomationLoop:
+ """
+ Background task that drives fully automated workflow execution.
+
+ Polls control state and processes the queue automatically.
+
+ IMPORTANT: This loop does NOT auto-start. It must be explicitly started
+ via the start() method (triggered by the "Start Automation" button).
+ """
+
+ def __init__(
+ self,
+ runner: PersistentWorkflowRunner,
+ redis_manager: WorkflowRedisManager,
+ poll_interval: float = 0.3,
+ step_delay: float = 0.5, # Delay between steps for control checks
+ ):
+ self._runner = runner
+ self._redis = redis_manager
+ self._poll_interval = poll_interval
+ self._step_delay = step_delay
+ self._task: asyncio.Task | None = None
+ self._enabled = False
+ self._on_step_complete: Callable[[str, str, StateResult], None] | None = None
+
+ @property
+ def is_running(self) -> bool:
+ return self._task is not None and not self._task.done()
+
+ @property
+ def is_enabled(self) -> bool:
+ return self._enabled
+
+ def set_step_callback(self, callback: Callable[[str, str, StateResult], None]) -> None:
+ """Set callback for step completion: callback(item_id, step_name, result)"""
+ self._on_step_complete = callback
+
+ def start(self) -> None:
+ """Start the automation loop. Must be explicitly called."""
+ if self._task is not None and not self._task.done():
+ return # Already running
+
+ self._enabled = True
+ self._task = asyncio.create_task(self._run_loop())
+
+ self._redis.append_event(WorkflowEvent(
+ beamline=self._runner.beamline,
+ event_type="automation_started",
+ message="Automation mode enabled",
+ ))
+
+ def stop(self) -> None:
+ """Stop the automation loop gracefully."""
+ self._enabled = False
+
+ # Cancel the task if it exists
+ if self._task is not None and not self._task.done():
+ self._task.cancel()
+
+ self._redis.append_event(WorkflowEvent(
+ beamline=self._runner.beamline,
+ event_type="automation_stopped",
+ message="Automation mode disabled",
+ ))
+
+ async def _run_loop(self) -> None:
+ """Main automation loop."""
+ while self._enabled:
+ try:
+ # Check control state FIRST before doing anything
+ control = self._runner.check_control()
+ runtime = self._redis.get_runtime()
+
+ # Handle abort immediately
+ if control.abort_requested:
+ if runtime.current_item_id:
+ self._runner.complete_item(runtime.current_item_id, QueueItemStatus.ABORTED)
+ # Clear all control flags including pause
+ self._redis.clear_control()
+ self.stop()
+ await asyncio.sleep(self._poll_interval)
+ continue
+
+ # Handle pause
+ if control.pause_requested:
+ if runtime.running and not runtime.paused:
+ self._redis.patch_runtime({"paused": True})
+ self._redis.append_event(WorkflowEvent(
+ beamline=self._runner.beamline,
+ item_id=runtime.current_item_id,
+ event_type="workflow_paused",
+ message="Workflow paused by user request",
+ ))
+ await asyncio.sleep(self._poll_interval)
+ continue
+
+ # Handle resume
+ if control.resume_requested and runtime.paused:
+ self._redis.patch_runtime({"paused": False})
+ self._redis.request_control(
+ {"resume_requested": False},
+ requested_by="automation_loop",
+ )
+ self._redis.append_event(WorkflowEvent(
+ beamline=self._runner.beamline,
+ item_id=runtime.current_item_id,
+ event_type="workflow_resumed",
+ message="Workflow resumed",
+ ))
+
+ # Don't process if paused
+ if runtime.paused:
+ await asyncio.sleep(self._poll_interval)
+ continue
+
+ await self._tick()
+
+ except asyncio.CancelledError:
+ # Loop was cancelled (stop() was called)
+ break
+ except Exception as e:
+ self._redis.patch_runtime({"last_error": str(e)})
+ self._redis.append_event(WorkflowEvent(
+ beamline=self._runner.beamline,
+ event_type="automation_error",
+ message=f"Automation error: {e}",
+ ))
+ await asyncio.sleep(2.0)
+
+ await asyncio.sleep(self._poll_interval)
+
+ async def _tick(self) -> None:
+ """Single iteration of the automation loop - process one step."""
+ runtime = self._redis.get_runtime()
+ control = self._runner.check_control()
+
+ # If nothing running, try to start next item
+ if not runtime.running or not runtime.current_item_id:
+ next_item = self._redis.get_next_pending_item()
+ if next_item is None:
+ return # Queue empty, nothing to do
+
+ self._runner.start_item(next_item.item_id)
+ # Add delay after starting to allow UI to catch up
+ await asyncio.sleep(self._step_delay)
+ runtime = self._redis.get_runtime()
+
+ # Re-check control after potential start
+ control = self._runner.check_control()
+ if control.abort_requested or control.pause_requested:
+ return # Let main loop handle it
+
+ # Process current item
+ item = self._redis.get_item(runtime.current_item_id)
+ if item is None:
+ self._redis.patch_runtime({"running": False, "current_item_id": None})
+ return
+
+ # Check if item is complete
+ if item.current_step_index >= len(item.steps):
+ self._runner.complete_item(item.item_id, QueueItemStatus.COMPLETED)
+ # Clear next_sample_requested if set
+ if control.next_sample_requested:
+ self._redis.request_control(
+ {"next_sample_requested": False},
+ requested_by="automation_loop",
+ )
+ return
+
+ # Handle skip request (skip current step)
+ if control.skip_requested:
+ step_kind = item.steps[item.current_step_index].kind
+ self._redis.update_step(item.item_id, item.current_step_index, status="skipped")
+ self._redis.update_item(item.item_id, {"current_step_index": item.current_step_index + 1})
+ self._redis.request_control({"skip_requested": False}, requested_by="automation_loop")
+
+ self._redis.append_event(WorkflowEvent(
+ beamline=self._runner.beamline,
+ item_id=item.item_id,
+ step=step_kind,
+ event_type="step_skipped",
+ message=f"Step {step_kind} skipped",
+ ))
+ return
+
+ # Build context and run step
+ context = WorkflowContext(
+ mode=WorkflowMode.AUTOMATION,
+ queue_id="default",
+ item_id=item.item_id,
+ sample_id=item.sample_id,
+ current_state=WorkflowStateKind(runtime.current_state) if runtime.current_state else None,
+ current_step_index=item.current_step_index,
+ )
+
+ # Run the step (this is blocking in the async context)
+ step_name = item.steps[item.current_step_index].kind
+ result = await asyncio.to_thread(self._runner.run_current_step, context)
+
+ # Add delay after step to allow control checks
+ await asyncio.sleep(self._step_delay)
+
+ # Notify callback if set
+ if self._on_step_complete:
+ self._on_step_complete(item.item_id, step_name, result)
+
+
+if __name__ == "__main__":
+ client = redis.Redis(host="localhost", port=6379, db=0, decode_responses=True)
+ redis_mgr = WorkflowRedisManager(client, beamline="x10sa")
+
+ # Create a queue item
+ item = QueueItem(
+ item_id="",
+ beamline="x10sa",
+ sample_id=123,
+ sample_name="lysozyme_01",
+ owner_pgroup="p12345",
+ created_by="user@example.com",
+ steps=build_default_steps(),
+ )
+ item = redis_mgr.create_item(item)
+
+ # Create runner
+ runner = PersistentWorkflowRunner(
+ redis_manager=redis_mgr,
+ registry=STATE_REGISTRY,
+ handlers=HANDLER_REGISTRY,
+ )
+
+ # Start processing
+ runner.start_item(item.item_id)
+
+ # Build context
+ context = WorkflowContext(
+ mode=WorkflowMode.GUIDED_MANUAL,
+ queue_id="default",
+ item_id=item.item_id,
+ sample_id=item.sample_id,
+ )
+
+ # Run each step (in guided mode, user triggers each one)
+ while context.current_step_index < len(item.steps):
+ if runner.should_pause():
+ print("Paused by user")
+ break
+
+ if runner.should_abort():
+ runner.complete_item(item.item_id, QueueItemStatus.ABORTED)
+ break
+
+ result = runner.run_current_step(context)
+ print(f"Step {result.state.value}: {result.status.value}")
+
+ # Mark complete if all steps done
+ if context.current_step_index >= len(item.steps):
+ runner.complete_item(item.item_id, QueueItemStatus.COMPLETED)
\ No newline at end of file
diff --git a/src/aare/daq/config.py b/src/aare/daq/config.py
index 61196563..3011466f 100644
--- a/src/aare/daq/config.py
+++ b/src/aare/daq/config.py
@@ -6,7 +6,7 @@ from typing import Tuple, List
import numpy as np
import redis
import redis_lock
-from aare.common.coordinate import Coordinate
+from aare.common.coordinate import Coordinate, AerotechCoordinate
from aare.common.models import (
BeamlineSettingsModel,
BeamMarkCoeffModel,
@@ -18,13 +18,25 @@ from aare.common.models import (
FluorescenceSpectrumOutputModel, CrystalSize, SimpleStrategyInputModel, SimpleScanParameters
)
+from aare.common.auth_models import (
+ BatonStatus,
+ BatonRequest,
+ BatonHolderInfo,
+ BatonRequestStatus,
+ BatonTransferQueue,
+)
from aare.common.beamline import MXBeamline
from aare.common.logger_config import setup_logger
from aare.common.exception_handler import BeamlineBusyException
-ABR_POS_ALIGN_DEF = Coordinate(x=-18, y=-0.266, z=0)
-ABR_POS_MOUNT = Coordinate(x=0, y=0, z=0)#Coordinate(x=-18, y=0, z=0)
+#TODO WHAT SHOULD THIS BE?
+ABR_POS_ALIGN_DEF = AerotechCoordinate(at_mm=Coordinate(x=-18, y=-0.266, z=0))
+ABR_POS_MOUNT = AerotechCoordinate(
+ at_mm=Coordinate(x=0, y=0, z=0),
+ omega_deg=0
+)
+#Coordinate(x=-18, y=0, z=0)
ABR_OMEGA_MOUNT = 0.0
logger = setup_logger("aareDAQ")
@@ -74,9 +86,24 @@ class BeamlineConfig:
else:
host = f"{self.__bl}-redis.psi.ch"
self.__client = redis.Redis(host=host, port=6379, db=0, decode_responses=True)
+ self.simulated_detector = bl is MXBeamline.SIMULATED
# Session and authentication management
+ @property
+ def allow_non_staff_request_from_staff(self) -> bool:
+ raw = self.__client.get(f"{self.__bl}:allow_non_staff_request_from_staff")
+ if raw is None:
+ return False
+ return str(raw).strip().lower() in {"1", "true", "yes", "on"}
+
+ @allow_non_staff_request_from_staff.setter
+ def allow_non_staff_request_from_staff(self, enabled: bool) -> None:
+ if enabled:
+ self.__client.set(f"{self.__bl}:allow_non_staff_request_from_staff", "1")
+ else:
+ self.__client.delete(f"{self.__bl}:allow_non_staff_request_from_staff")
+
def generate_session(self) -> int:
return int(self.__client.incr(f"{self.__bl}:session"))
@@ -152,6 +179,7 @@ class BeamlineConfig:
return
if active == session:
self.__client.delete(f"{self.__bl}:active_session")
+ self.__client.delete(f"{self.__bl}:baton_holder")
def force_set_active_session(self, session: int, expiry_sec: int) -> None:
# Ensure that there is no active try-set for active session
@@ -161,6 +189,125 @@ class BeamlineConfig:
self.__client.set(f"{self.__bl}:active_session", session)
self.__client.expire(f"{self.__bl}:active_session", expiry_sec)
+ # ========== BATON SYSTEM ==========
+
+ @property
+ def baton_holder(self) -> BatonHolderInfo | None:
+ """Get information about the current baton holder."""
+ tmp = self.__client.get(f"{self.__bl}:baton_holder")
+ if tmp is None:
+ return None
+ try:
+ return BatonHolderInfo(**json.loads(tmp))
+ except Exception:
+ return None
+
+ @baton_holder.setter
+ def baton_holder(self, info: BatonHolderInfo | None) -> None:
+ if info is None:
+ self.__client.delete(f"{self.__bl}:baton_holder")
+ else:
+ self.__client.set(f"{self.__bl}:baton_holder", info.model_dump_json())
+
+ @property
+ def pending_baton_request(self) -> BatonRequest | None:
+ """Get the current pending baton request, if any."""
+ tmp = self.__client.get(f"{self.__bl}:baton_request")
+ if tmp is None:
+ return None
+ try:
+ return BatonRequest(**json.loads(tmp))
+ except Exception:
+ return None
+
+ def set_pending_baton_request(self, request: BatonRequest | None, timeout_sec: int = 30) -> None:
+ """Set a pending baton request with auto-expiry for timeout."""
+ if request is None:
+ self.__client.delete(f"{self.__bl}:baton_request")
+ else:
+ self.__client.set(f"{self.__bl}:baton_request", request.model_dump_json())
+ # Add a few seconds buffer so we can detect timeout vs expiry
+ self.__client.expire(f"{self.__bl}:baton_request", timeout_sec + 5)
+
+ def clear_pending_baton_request(self) -> None:
+ self.__client.delete(f"{self.__bl}:baton_request")
+
+ @property
+ def queued_baton_transfer(self) -> BatonTransferQueue | None:
+ """Get queued transfer waiting for beamline to be available."""
+ tmp = self.__client.get(f"{self.__bl}:baton_transfer_queue")
+ if tmp is None:
+ return None
+ try:
+ return BatonTransferQueue(**json.loads(tmp))
+ except Exception:
+ return None
+
+ @queued_baton_transfer.setter
+ def queued_baton_transfer(self, transfer: BatonTransferQueue | None) -> None:
+ if transfer is None:
+ self.__client.delete(f"{self.__bl}:baton_transfer_queue")
+ else:
+ self.__client.set(f"{self.__bl}:baton_transfer_queue", transfer.model_dump_json())
+
+ def can_transfer_baton_now(self) -> bool:
+ """Check if baton can be transferred (beamline not mid-operation)."""
+ # Can't transfer while beamline is busy
+ if self.state_busy:
+ return False
+ # Add automation queue check here when you implement it
+ # if self.automation_queue_running:
+ # return False
+ return True
+
+ def execute_baton_transfer(
+ self,
+ to_session: int,
+ to_username: str,
+ to_is_staff: bool,
+ to_pgroup: str | None,
+ expiry_sec: int
+ ) -> None:
+ """
+ Atomically transfer the baton to a new holder.
+ Use existing active_session_lock for consistency.
+ """
+ with redis_lock.Lock(
+ self.__client, f"{self.__bl}:active_session_lock", expire=10
+ ):
+ self.__client.set(f"{self.__bl}:active_session", to_session)
+ self.__client.expire(f"{self.__bl}:active_session", expiry_sec)
+ self.baton_holder = BatonHolderInfo(
+ username=to_username,
+ session=to_session,
+ is_staff=to_is_staff,
+ pgroup=to_pgroup
+ )
+ # Clear any pending request or queued transfer
+ self.clear_pending_baton_request()
+ self.queued_baton_transfer = None
+
+ def process_queued_transfer_if_ready(self, expiry_sec: int) -> bool:
+ """
+ Check if there's a queued transfer and beamline is now available.
+ Returns True if transfer was executed.
+ """
+ queued = self.queued_baton_transfer
+ if queued is None:
+ return False
+
+ if not self.can_transfer_baton_now():
+ return False
+
+ self.execute_baton_transfer(
+ to_session=queued.target_session,
+ to_username=queued.target_username,
+ to_is_staff=queued.target_is_staff,
+ to_pgroup=queued.target_pgroup,
+ expiry_sec=expiry_sec
+ )
+ return True
+
@property
def pgroup(self) -> str | None:
tmp = self.__client.get(f"{self.__bl}:pgroup")
@@ -458,16 +605,16 @@ class BeamlineConfig:
self.__client.set(f"{self.__bl}:{self.zoom_setting_string(mode)}", data.model_dump_json())
@property
- def abr_meas_pos(self) -> Coordinate:
+ def abr_meas_pos(self) -> AerotechCoordinate:
tmp = self.__client.get(f"{self.__bl}:abr_meas_pos")
if tmp is None:
return ABR_POS_ALIGN_DEF
data_dict = json.loads(tmp)
- return Coordinate(**data_dict)
+ return AerotechCoordinate(**data_dict)
@abr_meas_pos.setter
- def abr_meas_pos(self, data: Coordinate):
+ def abr_meas_pos(self, data: AerotechCoordinate):
self.__client.set(f"{self.__bl}:abr_meas_pos", data.model_dump_json())
@property
@@ -625,3 +772,8 @@ class BeamlineConfig:
def increment_failed_mount_count(self) -> int:
return int(self.__client.incr(f"{self.__bl}:failed_mount_count"))
+
+if __name__ == "__main__":
+ from aare.common.beamline import mx_beamline
+ cfg = BeamlineConfig(bl=mx_beamline())
+ cfg.allow_non_staff_request_from_staff = True
\ No newline at end of file
diff --git a/src/aare/daq/daq.py b/src/aare/daq/daq.py
index 9a72ec0e..606d812e 100644
--- a/src/aare/daq/daq.py
+++ b/src/aare/daq/daq.py
@@ -1,4 +1,4 @@
-
+import copy
import secrets
import time
import traceback
@@ -11,6 +11,7 @@ from typing import List, Tuple, Optional, Callable
import cv2
import numpy as np
+from jfjoch_client import ScanResult, ScanResultImagesInner
import aare.common.face_detection as fd
from aare.daq import workflows
@@ -21,7 +22,7 @@ from aare.daq.config import BeamlineStateEnum
from aare.daq.devices import BeamlineDevices
from aare.daq.mlbox import MlBox
from aare.common.beamline import MXBeamline
-from aare.common.coordinate import Coordinate, SmargonCoordinate
+from aare.common.coordinate import Coordinate, SmargonCoordinate, AerotechCoordinate
from aare.common.diffraction_geometry import DiffractionGeometry
from aare.common.logger_config import setup_logger
from aare.common.models import (
@@ -34,6 +35,7 @@ from aare.common.models import (
from aare.common.raster_grid import RasterGridRequest, CompletedRasterGrid, CompletedRasterGridElem
from aare.common.rotation_scan import RotationScanRequest, CompletedRotationScan
from aare.common.sample_geometry import SampleGeometryModel
+from aare.daq.spreadsheetupdater import beamline
from aare.devices.area_detector import AutoEnum
from aare.devices.jfjoch import JFJochWrapper
from aare.devices.mx_lib import clean_filename
@@ -44,7 +46,10 @@ from aare.common.exception_handler import (
MountingFailed,
WarningTellException,
CriticalTellException,
- AXCFailed, SmargonCommunicationError, TellCommunicationError,
+ AXCFailed,
+ SmargonCommunicationError,
+ TellCommunicationError,
+ JFJochCommunicationError
)
logger = setup_logger("aareDAQ")
@@ -62,7 +67,7 @@ class AareDAQ:
self.__bl = bl.value.upper()
self.__aare = AareWrapper(bl)
self.__saved_box = None
- self._smargon_trace_path = Path("logs") / "smargon_trace.csv"
+ self._smargon_trace_path = Path("/sls/mx/applications/logs") / "smargon_trace.csv"
self._face_detection_progress_cb: Callable[[dict], None] | None = None
self._last_sample_sync_ts = 0.0
self._sample_sync_min_interval_s = 2.0
@@ -358,10 +363,10 @@ class AareDAQ:
def samcam_settings(self, s: SampleCameraSettings):
self.__devs.samcam_settings = s
- def tweak_abr_meas_pos(self, c: Coordinate):
+ def tweak_abr_meas_pos(self, c: AerotechCoordinate):
self.__cfg.set_busy(BeamlineStateEnum.SampleAlignment)
try:
- new_meas_pos = self.__cfg.abr_meas_pos + c
+ new_meas_pos = AerotechCoordinate(at_mm=self.__cfg.abr_meas_pos.at_mm + c.at_mm)
self.__cfg.abr_meas_pos = new_meas_pos
self.__devs.aerotech_pos = new_meas_pos
self.__saved_box = None
@@ -405,13 +410,27 @@ class AareDAQ:
self.__cfg.state_busy = False
raise
+ def park_and_dry(self):
+ self.__cfg.try_set_busy(timeout=360)
+ try:
+ self.__devs.tell.check_enable_motion()
+ self.__devs.tell.wait_not_busy()
+ self.__devs.tell.set_in_mount_position(True)
+ self.__devs.tell.unmount(wait=True, timeout=60.0)
+ self.__devs.tell.dry(wait_cold=-1, wait=True, timeout=360.0)
+ self.__cfg.current_sample = None
+ self.__cfg.state_busy = False
+ except Exception as e:
+ self.__cfg.state_busy = False
+ logger.error(f"Failed to park and dry: {e}")
+ raise e
+
def __mount_failure_handler(self, mount_error):
pass
def __mount(self, target: SampleShortInfo | None):
self.__devs.smargon_move_home()
self.__devs.aerotech_pos = ABR_POS_MOUNT
- self.__devs.aerotech_omega = ABR_OMEGA_MOUNT
#collimator should be down!!!
self.__devs.tell.check_enable_motion()
self.__devs.tell.wait_not_busy()
@@ -458,14 +477,14 @@ class AareDAQ:
except Exception as e:
self.__cfg.state_busy = False
logger.debug(f"Failed to mount sample: {e}")
- self.__aare.sample_failed(target, f"Mount failed due to {e}")
+ # self.__aare.sample_failed(target, f"Mount failed due to {e}")
raise e
workflows.rse2sa(devs=self.__devs, cfg=self.__cfg)
self.__cfg.state_busy = False
- if target is not None:
- if target.db_id is not None:
- self.__aare.sample_mounted(target)
- self.save_screenshot_db(target.db_id, f"{target.db_id}_mounted")
+ # if target is not None:
+ # if target.db_id is not None:
+ # self.__aare.sample_mounted(target)
+ # self.save_screenshot_db(target.db_id, f"{target.db_id}_mounted")
@property
@@ -676,30 +695,120 @@ class AareDAQ:
else:
return None
+ def __setup_datacollection(self, request: RasterGridRequest | RasterGridRequest):
+
+ if request.dtz is not None:
+ logger.info(f'requesting dtz to move to {request.dtz}')
+ self.__cfg.dtz = request.dtz
+
+ if self.sample is not None and self.sample.db_id is not None:
+ self.save_screenshot_db(self.sample.db_id, f"{self.sample.db_id}_before_data_collection")
- def __raster(self, r: RasterGridRequest):# -> CompletedRasterGridElem:
- max_time = r.exp_time_s * r.n_x * r.n_y + 60
- row_time = r.exp_time_s * r.n_x
- row_width_mm = r.grid_size_mm.x * r.n_x
self.__set_state(BeamlineStateEnum.DataCollection)
- save_smargon_position = self.__devs.smargon_pos
- if r.smargon_top_left is not None:
- delta_mm = self.sample_geometry.smargon_nudge(Coordinate(x=r.grid_size_mm.x / 2, y=r.grid_size_mm.y / 2))
- self.__devs.set_smargon_pos(SmargonCoordinate(sh_mm=r.smargon_top_left.sh_mm + delta_mm,
- phi_deg=r.smargon_top_left.phi_deg,
- chi_deg=r.smargon_top_left.chi_deg))
- self.__devs.aerotech_omega = r.omega_deg
+
+ if request.transmission is not None:
+ logger.info(f'requesting transmission to move to {request.transmission}')
+ self.__devs.transmission = request.transmission
+
+ if hasattr(request, 'start') and request.start is not None:
+ self.__devs.smargon_pos = request.start
+ elif hasattr(request, 'smargon_top_left') and request.smargon_top_left is not None:
+ self.__devs.set_smargon_pos(SmargonCoordinate(sh_mm=request.smargon_top_left.sh_mm,
+ phi_deg=request.smargon_top_left.phi_deg,
+ chi_deg=request.smargon_top_left.chi_deg))
+
+ #if request.transmission is not None:
+ # self.__devs.transmission.wait()
self.__devs.smargon_wait(timeout=180)
+
+ return
+
+ def _build_fake_scan_result(self, *, file_prefix: str | None, image_count: int) -> ScanResult:
+ images = [
+ ScanResultImagesInner(
+ number=i,
+ efficiency=1.0,
+ bkg=0.0,
+ spots=0,
+ spots_low_res=0,
+ spots_indexed=0,
+ index=0,
+ b=0.0,
+ )
+ for i in range(max(1, image_count))
+ ]
+ return ScanResult(file_prefix=file_prefix, images=images)
+
+ def _build_fake_rotation_result(self, request: RotationScanRequest) -> CompletedRotationScan:
+ result = self._build_fake_scan_result(
+ file_prefix=request.file_prefix,
+ image_count=request.steps,
+ )
+ return CompletedRotationScan(
+ request=copy.deepcopy(request),
+ result=result,
+ )
+
+ def _build_fake_raster_result(self, request: RasterGridRequest) -> CompletedRasterGridElem:
+ result = self._build_fake_scan_result(
+ file_prefix=request.file_prefix,
+ image_count=request.n_x * request.n_y,
+ )
+ return CompletedRasterGridElem(
+ request=copy.deepcopy(request),
+ result=result,
+ centre_of_mass=None,
+ )
+
+ def __raster(self, request: RasterGridRequest) -> CompletedRasterGridElem:
+ self.__devs.aerotech_omega = request.omega_deg
+ self.__setup_datacollection(request=request)
+
status = self.status
- print(f"raster status {status}")
- print(f'raster grid request: {r}')
- #self.__jfjoch.measure_raster(r, status)
- self.__devs.aerotech.run_grid_scan(cell_height_mm=r.grid_size_mm.y, num_rows=r.n_x,
- row_width_mm=row_width_mm, time_per_row_s=row_time, task_id=3)
- self.__devs.aerotech_pos = self.__cfg.abr_meas_pos
- #result = self.__jfjoch.wait_till_done(60)
- return None
- #return CompletedRasterGridElem(request=copy.deepcopy(r), result=, centre_of_mass=None)
+ logger.info(f"raster status {status}")
+ logger.info(f'raster grid request: {request}')
+ total_time = request.exp_time_s*request.n_x*request.n_y
+
+ if not self.__cfg.simulated_detector:
+ try:
+ logger.info('initialise detector')
+ self.__jfjoch.measure_raster(request, status)
+ logger.info('detector initialised')
+ except JFJochCommunicationError as e:
+ logger.warning(f"Failed to communicate with jfjoch: {e}")
+
+ else:
+ logger.info("Simulated detector mode enabled; returning fake zero raster result.")
+
+ try:
+ self.__devs.aerotech.grid_scan(grid_elem_size_y_um=request.grid_size_mm.y*1000,
+ grid_elem_size_x_um=request.grid_size_mm.x*1000,
+ grid_elem_count_x=request.n_x,
+ grid_elem_count_y=request.n_y,
+ time_sec=request.exp_time_s,
+ run_async=True)
+
+ self.__devs.aerotech.wait_till_done(timeout=int(round(total_time*2,0)))
+
+ self.__devs.aerotech_pos = self.__cfg.abr_meas_pos
+
+ try:
+ result = self.__jfjoch.wait_till_done(60)
+ except JFJochCommunicationError as e:
+ logger.warning(f"Failed to communicate with jfjoch: {e}")
+ return self._build_fake_raster_result(request)
+
+ if result is None:
+ logger.warning("JFJoch returned no ScanResult; using fake result for raster scan.")
+ return self._build_fake_raster_result(request)
+
+ return CompletedRasterGridElem(request=copy.deepcopy(request), result=result, centre_of_mass=None)
+
+ except JFJochCommunicationError:
+ raise
+ except Exception as e:
+ logger.error(f"Failed during raster: {e}")
+ raise Exception(f"Failed during raster: {e}") from e
def measure_raster(self, r: RasterGridRequest, auto: bool) -> CompletedRasterGrid:
self.__cfg.try_set_busy(timeout=ceil(360))
@@ -708,25 +817,96 @@ class AareDAQ:
result = self.__auto_center(r)
else:
raster_result = self.__raster(r)
- #result = CompletedRasterGrid(r=[raster_result])
+ result = CompletedRasterGrid(r=[raster_result])
self.__set_state(BeamlineStateEnum.SampleAlignment)
self.__cfg.state_busy = False
- return CompletedRasterGrid(r=[])
- #return result
+ #return CompletedRasterGrid(r=[])
+ return result
except Exception as e:
try:
self.__aare.axc_failed(self.sample)
except Exception as axc_e:
logger.error(f"Exception while reporting AXC failure: {axc_e}")
- self.__set_state(BeamlineStateEnum.SampleAlignment)
+ if not self.busy:
+ self.__cfg.try_set_busy(timeout=ceil(360))
+ if self.status.state != BeamlineStateEnum.SampleAlignment:
+ self.__set_state(BeamlineStateEnum.SampleAlignment)
self.__cfg.state_busy = False
raise e
def __rotation(self, request: RotationScanRequest) -> CompletedRotationScan:
- return None
+ omega_start = self.omega
+ status = self.status
+
+ #self.__aare.create_rotation_run(self.sample, request, status)
+ total_time = request.exp_time_s * request.steps
+
+ try:
+
+ if self.__cfg.simulated_detector:
+ logger.info("Simulated detector mode enabled; skipping JFJoch start.")
+ else:
+ self.__jfjoch.measure_rotation(request, status, self.__cfg.xrf)
+
+ if request.screening:
+ self.__devs.aerotech.screening_scan(
+ rotation_deg=request.steps*request.incr_omega_deg,
+ wedge_deg=request.wedge_omega_deg,
+ time_sec=total_time,
+ steps=request.steps,
+ run_async=True,
+ )
+ else:
+ self.__devs.aerotech.rotation_scan(
+ rotation_deg=request.steps*request.incr_omega_deg,
+ time_sec=total_time,
+ start_pos_deg=request.start_omega_deg,
+ run_async=True,
+ )
+
+ #Is this for helical scans...? do we do smargon scans?
+ if request.start is not None and request.end is not None:
+ smargon_time_step = request.time_sec / float(request.steps)
+ pos_step = (request.end.sh_mm - request.start.sh_mm) * (1.0 / float(request.steps))
+
+ for i in range(request.steps):
+ self.__devs.smargon.target = SmargonCoordinate(
+ sh_mm=request.start.sh_mm + pos_step * i
+ )
+ time.sleep(smargon_time_step)
+
+ self.__devs.aerotech.wait_till_done(timeout=int(round(total_time + 60,0)))
+ self.__devs.aerotech_omega = omega_start
+
+ # self.__aare.sample_collected(self.sample)
+
+ if self.__cfg.simulated_detector:
+ logger.warning("Detector in simulation mode, returning fake zero rotation result.")
+ result = self._build_fake_rotation_result(request)
+
+ else:
+ # Let JFJochCommunicationError propagate
+ result = self.__jfjoch.wait_till_done(60)
+
+ # if self.sample is not None and self.sample.db_id is not None:
+ # self.save_screenshot_db(self.sample.db_id, "after_dc")
+ # try:
+ # self.__aare.ingest_scan(sample=self.sample, result=result,
+ # geom=self.sample_geometry, beam_mark_pxl=self.__cfg.get_beam_mark(self.zoom))
+ #
+ # except Exception as e:
+ # logger.error(f"Exception ingesting scan: {e}")
+ except JFJochCommunicationError:
+ raise
+ except Exception as e:
+ logger.error(f"Exception during rotation scan: {e}")
+ raise
+
+ return CompletedRotationScan(request=copy.deepcopy(request), result=result)
def measure_rotation(self, request: RotationScanRequest) -> CompletedRotationScan:
total_time = request.exp_time_s * request.steps
+ logger.info(f"received rotation scan request: {request}, total time: {total_time}s, steps: {request.steps}")
self.__cfg.try_set_busy(timeout=ceil(total_time + 360))
try:
@@ -739,31 +919,6 @@ class AareDAQ:
self.__cfg.state_busy = False
raise e
- def get_background(self):
- #if self.sample is not None:
- # raise Exception("Background cannot be measured when sample is mounted")
-
- self.__cfg.try_set_busy(timeout=360)
- self.__set_state(BeamlineStateEnum.SampleAlignment)
- try:
- self.__devs.lamp_light = 2.5
- self.__cfg.zoom_mode = ZoomModeEnum.LoopCenter
- zoom_settings = self.__cfg.zoom_settings.z
- for zoom_value in zoom_settings:
- exposure = zoom_settings[zoom_value].exposure
- gain = zoom_settings[zoom_value].gain
- self.__devs.samcam_settings = SampleCameraSettings(exposure=exposure, gain=gain)
- self.__devs.set_zoom(zoom_value, wait=True)
- time.sleep(5.0) # Wait for settings to stabilize
- image = self.__devs.samcam_get_image(gray=False)
- self.__cfg.put_alc_bkg(zoom_value, exposure, gain, image)
- bgr_array = cv2.cvtColor(image, cv2.COLOR_RGB2BGR)
- cv2.imwrite(f"bkg{zoom_value:.0f}_{exposure*1000:.0f}_{gain:.0f}.jpg", bgr_array)
- self.__cfg.state_busy = False
- except Exception as e:
- self.__cfg.state_busy = False
- raise e
-
@property
def dtz(self) -> float:
tmp = self.__cfg.dtz
@@ -823,8 +978,8 @@ class AareDAQ:
@property
def sample_geometry(self) -> SampleGeometryModel:
zoom = self.__devs.zoom
- aerotech_pos_ref = self.__cfg.abr_meas_pos
- aerotech_pos = self.__devs.aerotech_pos
+ aerotech_pos_ref = self.__cfg.abr_meas_pos.at_mm
+ aerotech_pos = self.__devs.aerotech_pos.at_mm
sample_geom = SampleGeometryModel(
beam_location_pxl=self.__cfg.beam_mark_coeff.apply(zoom),
pixel_in_mm=self.__cfg.pixel_to_mm(zoom),
@@ -1632,6 +1787,16 @@ class AareDAQ:
# Keep status flowing even if Tell code throws something unexpected
return self.__cfg.current_sample, False, f"TELL unavailable: {e}"
+ def _aerotech_status(self) -> tuple[bool, str | None]:
+ aerotech_ok = True
+ aerotech_err: str | None = None
+ try:
+ _ = self.__devs.aerotech.status()
+ except Exception as e:
+ aerotech_ok = False
+ aerotech_err = str(e)
+ return aerotech_ok, aerotech_err
+
def _safe_geom(self) -> tuple[SampleGeometryModel, bool, str | None]:
"""
Return (geom, smargon_connected, smargon_error) without raising.
@@ -1714,6 +1879,7 @@ class AareDAQ:
def status(self) -> DAQStatusModel:
safe_sample, tell_ok, tell_err = self._safe_sample()
safe_geom, smargon_ok, smargon_err = self._safe_geom()
+ aerotech_ok, aerotech_err = self._aerotech_status()
return DAQStatusModel(
state=self.state,
@@ -1731,11 +1897,13 @@ class AareDAQ:
tell_error=tell_err,
smargon_connected=smargon_ok,
smargon_error=smargon_err,
+ aerotech_connected=aerotech_ok,
+ aerotech_error=aerotech_err,
)
def cancel(self):
if self.__cfg.state == BeamlineStateEnum.DataCollection:
- self.__devs.aerotech_stop()
+ self.__devs.aerotech.cancel()
self.__jfjoch.cancel()
def anneal(self, time_s: float):
diff --git a/src/aare/daq/devices.py b/src/aare/daq/devices.py
index 0455beea..78fdd330 100644
--- a/src/aare/daq/devices.py
+++ b/src/aare/daq/devices.py
@@ -5,15 +5,11 @@
# - setter with option to do sync/async move
# - property setter, which assumes that sync move is done (excl. zoom, which is async by default)
-import time
-from enum import Enum
-
import numpy as np
from epics import PV
-from fontTools.feaLib.ast import deviceToString
from aare.common.beamline import MXBeamline
-from aare.common.coordinate import Coordinate, SmargonCoordinate
+from aare.common.coordinate import SmargonCoordinate, AerotechCoordinate
from aare.common.models import SampleCameraSettings, StagePositionEnum
from aare.devices import smargon, aerotech
from aare.devices.area_detector import epicsAD, AutoEnum
@@ -27,8 +23,7 @@ class BeamlineDevices:
def __init__(self, beamline: MXBeamline):
BEAMLINE = beamline.value.upper()
self.tell = make_tell_client(beamline)
- self.__aerotech = aerotech.AerotechControllerEpics(beamline)
- self.aerotech = aerotech.AerotechController(controller_ip="129.129.118.96")
+ self.aerotech = aerotech.AerotechController(beamline)
self.__smargon = smargon.Smargon(beamline)
self.__ring_current_pv = PV(f"ARS07-DPCT-0100:CURR")
@@ -91,19 +86,24 @@ class BeamlineDevices:
self.__cryojet_x = MyMotor(f"{BEAMLINE}-ES-CS:TRX") #currently in is 5 out is 15?
+ self.magnet_position_sensor = PV(f"{BEAMLINE}-ES-DF1:CBOX-CMP1")
+ self.magnet_position_sensor_readout = PV(f"{BEAMLINE}-ES-DF1:CBOX-USER1")
+
+
# Transmission
@property
def transmission(self) -> float:
return 1.0
- def set_transmission(self, value: float, /, wait: bool = True):
- pass
@transmission.setter
def transmission(self, value: float):
self.set_transmission(value, wait=False)
+ def set_transmission(self, value: float, /, wait: bool = True):
+ pass
+
# Lamp light
@property
def lamp_light(self) -> float:
@@ -262,8 +262,6 @@ class BeamlineDevices:
"""
return int(self.__sample_cam.uid.get())
-
-
# Detector Z
@property
def dtz(self) -> float:
@@ -284,32 +282,46 @@ class BeamlineDevices:
def dtz_high(self) -> float:
return self.__dtz.get("HLM")
-# Aertoech Automation1
@property
- def aerotech_pos(self) -> Coordinate:
- return Coordinate(x=self.__aerotech.gmx.readback,
- y=self.__aerotech.gmy.readback, z=self.__aerotech.gmz.readback)
+ def aerotech_pos(self) -> AerotechCoordinate:
+ #TODO add u/omega????
+ return self.aerotech.get_position()
@aerotech_pos.setter
- def aerotech_pos(self, pos: Coordinate):
- # self.__aerotech.gmx.move(pos.x, wait=False)
- # self.__aerotech.gmy.move(pos.y, wait=False)
- # self.__aerotech.gmz.move(pos.z, wait=False)
- self.__aerotech.gmx.move(pos.x, wait=True)
- self.__aerotech.gmy.move(pos.y, wait=True)
- self.__aerotech.gmz.move(pos.z, wait=True)
- #TODO add wait pos?
+ def aerotech_pos(
+ self,
+ coord: AerotechCoordinate,
+ /,
+ wait: bool = True,
+ incremental: bool = False,
+ ):
+ self.aerotech.position(
+ coord,
+ wait=wait,
+ incremental=incremental,
+ )
@property
def aerotech_omega(self) -> float:
- return self.__aerotech.omega.readback
+ return self.aerotech.status().u.pos
@aerotech_omega.setter
def aerotech_omega(self, val: float):
self.set_aerotech_omega(val, wait=True)
- def set_aerotech_omega(self, val: float, /, wait: bool = True):
- self.__aerotech.omega.move(val, wait=wait)
+ def set_aerotech_omega(
+ self,
+ val: float,
+ /,
+ wait: bool = True,
+ incremental: bool = False,
+ ):
+ target = AerotechCoordinate(omega_deg=val)
+ return self.aerotech.position(
+ target,
+ wait=wait,
+ incremental=incremental,
+ )
@property
def aerotech_lock(self) -> bool:
diff --git a/src/aare/daq/mlbox.py b/src/aare/daq/mlbox.py
index 094233c2..84d00fc7 100644
--- a/src/aare/daq/mlbox.py
+++ b/src/aare/daq/mlbox.py
@@ -30,7 +30,7 @@ class MlBox:
elif bl == MXBeamline.X06DA:
self.__url = "http://mx-aare-test.psi.ch:8002/predict/?model=best_v8_20102025.pt"
elif bl == MXBeamline.X10SA:
- self.__url = "http://x10sa-spark-01.psi.ch:8002/predict/?model=best_v12_22092025.engine"
+ self.__url = "http://x10sa-spark-01.psi.ch:8002/predict/?model=best_yolo26l-seg-overlap-false_2026-03-16.engine"#v12_22092025.engine"
elif bl == MXBeamline.X06SA:
self.__url = ""
raise NotImplemented(f"MLBox not implemente for {bl}")
diff --git a/src/aare/daq/server.py b/src/aare/daq/server.py
index 45f57d4f..92e27d6e 100644
--- a/src/aare/daq/server.py
+++ b/src/aare/daq/server.py
@@ -7,6 +7,9 @@ import json
import cv2
import urllib3
import uvicorn
+
+from aare.common.coordinate import SmargonCoordinate, Coordinate, AerotechCoordinate
+from aare.common.auth_models import BatonStatus, BatonRequestStatus
from aare.common.coordinate import SmargonCoordinate, Coordinate
from aare.common.error_codes import export_error_codes, export_error_codes_grouped
from aare.common.logger_config import setup_logger
@@ -34,6 +37,12 @@ from aare.common.exception_handler import (
SampleException,
UserRightsException,
)
+
+from aare.common.automation_queue_manager import WorkflowRedisManager
+from aare.common.automation_workflow import STATE_REGISTRY, HANDLER_REGISTRY, SIMULATED_HANDLER_REGISTRY
+from aare.daq.automation_runner import PersistentWorkflowRunner
+from aare.daq.automation_api_router import router as workflow_router, set_workflow_dependencies
+
logger = setup_logger("aareDAQ")
app = FastAPI()
register_exception_handlers(app)
@@ -45,6 +54,31 @@ bl = mx_beamline()
cfg = BeamlineConfig(bl)
daq = AareDAQ(cfg, bl)
+# ─────────────────────────────────────────────
+# Initialize workflow system
+# ─────────────────────────────────────────────
+
+USE_SIMULATED_WORKFLOW = os.getenv("WORKFLOW_SIMULATION", "0") == "1"
+
+workflow_redis_manager = WorkflowRedisManager(
+ client=cfg._BeamlineConfig__client, # Reuse existing Redis connection
+ beamline=bl.value,
+)
+
+workflow_runner = PersistentWorkflowRunner(
+ redis_manager=workflow_redis_manager,
+ registry=STATE_REGISTRY,
+ handlers=SIMULATED_HANDLER_REGISTRY if USE_SIMULATED_WORKFLOW else HANDLER_REGISTRY,
+)
+
+if USE_SIMULATED_WORKFLOW:
+ logger.warning("⚠️ Workflow system running in SIMULATION mode - no actual DAQ operations")
+
+set_workflow_dependencies(workflow_redis_manager, workflow_runner, cfg)
+
+# Include the workflow router
+app.include_router(workflow_router)
+
try:
daq.sync_current_sample_from_tell(force=True)
except Exception as e:
@@ -205,8 +239,7 @@ async def smargon(val: SmargonCoordinate, token: str = Depends(oauth2_scheme)):
@app.post("/beamline/tweak_abr_meas_pos")
-async def tweak_abr_meas_pos(val: Coordinate, token: str = Depends(oauth2_scheme)):
- logger.debug(f"Setting abr to {val}")
+async def tweak_abr_meas_pos(val: AerotechCoordinate, token: str = Depends(oauth2_scheme)):
auth.check_jwt_staff(cfg, auth.parse_token(token))
daq.tweak_abr_meas_pos(val)
return "OK"
@@ -316,6 +349,15 @@ async def sample(token: str = Depends(oauth2_scheme)) -> SampleShortInfo:
location=s.location
)
+async def park_and_dry(token: str = Depends(oauth2_scheme)):
+ token_data = auth.parse_token(token)
+ auth.check_jwt_ro(cfg, auth.parse_token(token))
+ daq.park_and_dry()
+ return {
+ "ok": True,
+ "message": "TELL has been dryed and parked",
+ }
+
@app.post("/sample/mount")
async def mount(dbid: int, token: str = Depends(oauth2_scheme), reference: bool = False):
@@ -602,13 +644,6 @@ async def cancel(token: str = Depends(oauth2_scheme)):
daq.cancel()
# ALC routines
-@app.post("/alc/background")
-async def alc_background(token: str = Depends(oauth2_scheme)) -> str:
- auth.check_jwt_rw(cfg, auth.parse_token(token))
- daq.get_background()
- return "OK"
-
-
@app.post("/alc/center_loop")
async def alc_center_loop(token: str = Depends(oauth2_scheme)) -> str:
logger.debug(f"ALC")
@@ -659,11 +694,19 @@ async def pgroup(token: str = Depends(oauth2_scheme)) -> str:
@app.put("/access/pgroup")
async def set_pgroup(val: str, token: str = Depends(oauth2_scheme)) -> str:
- auth.check_jwt_ro(cfg, auth.parse_token(token))
+ data = auth.parse_token(token)
+
+ holder = cfg.baton_holder
+ is_current_holder = holder is not None and holder.session == data.session
+
+ # Staff or current baton holder may change p-group even if it is not currently active.
+ # Everyone else must still belong to the active p-group.
+ if not (data.staff or is_current_holder):
+ auth.check_jwt_ro(cfg, data)
+
cfg.pgroup = val
return "OK"
-
@app.delete("/access/pgroup")
async def del_pgroup(token: str = Depends(oauth2_scheme)) -> str:
auth.check_jwt_ro(cfg, auth.parse_token(token))
@@ -696,6 +739,133 @@ async def force_current_session(token: str = Depends(oauth2_scheme)) -> str:
auth.force_current_sesion(cfg, data)
return "OK"
+# ========== BATON CONTROL ENDPOINTS ==========
+
+@app.get("/baton/status")
+async def baton_status(token: str = Depends(oauth2_scheme)) -> BatonStatus:
+ """Get the current baton status for the requesting user."""
+ data = auth.parse_token(token)
+ auth.resolve_baton_timeout_if_needed(cfg)
+ return auth.get_baton_status(cfg, data)
+
+
+@app.post("/baton/request")
+async def baton_request(token: str = Depends(oauth2_scheme)) -> dict:
+ """
+ Request control (baton) of the beamline.
+
+ - If vacant: granted immediately
+ - If staff requesting: granted immediately (or queued if busy)
+ - If same level: creates pending request with timeout
+ - Non-staff cannot request from staff
+ """
+ logger.debug(cfg.allow_non_staff_request_from_staff)
+ data = auth.parse_token(token)
+ return auth.request_baton(cfg, data)
+
+@app.post("/baton/respond")
+async def baton_respond(accept: bool, token: str = Depends(oauth2_scheme)) -> dict:
+ """
+ Current baton holder responds to a pending request.
+
+ - accept=true: transfers baton (or queues if busy)
+ - accept=false: refuses the request
+ """
+ data = auth.parse_token(token)
+ return auth.respond_to_baton_request(cfg, data, accept)
+
+
+@app.post("/baton/release")
+async def baton_release(token: str = Depends(oauth2_scheme)) -> dict:
+ """Voluntarily release the baton, making the beamline vacant."""
+ data = auth.parse_token(token)
+ return auth.release_baton(cfg, data)
+
+
+@app.post("/baton/cancel")
+async def baton_cancel(token: str = Depends(oauth2_scheme)) -> dict:
+ """Cancel your own pending baton request."""
+ data = auth.parse_token(token)
+ return auth.cancel_baton_request(cfg, data)
+
+
+@app.get("/baton/check_timeout")
+async def baton_check_timeout(token: str = Depends(oauth2_scheme)) -> dict:
+ """
+ Check if a pending request has timed out and process it.
+ Called by GUI to poll for timeout completion.
+ """
+ data = auth.parse_token(token)
+
+ auth.resolve_baton_timeout_if_needed(cfg)
+
+ pending = cfg.pending_baton_request
+ if pending is None:
+ if cfg.baton_holder and cfg.baton_holder.session == data.session:
+ return {"granted": True, "message": "Baton acquired!"}
+ queued = cfg.queued_baton_transfer
+ if queued and queued.target_session == data.session:
+ return {"queued": True, "message": "Transfer queued"}
+ return {"no_pending": True}
+
+ if pending.requester_session != data.session:
+ return {"not_your_request": True}
+
+ if pending.status == BatonRequestStatus.REFUSED:
+ cfg.clear_pending_baton_request()
+ return {"refused": True, "message": "Request refused"}
+
+ elapsed = time.time() - pending.created_at
+ if elapsed < pending.timeout_seconds:
+ return {
+ "pending": True,
+ "remaining_seconds": pending.timeout_seconds - elapsed
+ }
+
+ return auth.request_baton(cfg, data)
+
+@app.put("/access/allow_non_staff_request_from_staff")
+async def set_allow_non_staff_request_from_staff(val: bool, token: str = Depends(oauth2_scheme)) -> str:
+ data = auth.parse_token(token)
+ auth.check_jwt_staff_only(data)
+ cfg.allow_non_staff_request_from_staff = val
+ return "OK"
+
+async def baton_status_event_stream(data: TokenData) -> AsyncGenerator[str, None]:
+ """SSE stream for baton status updates."""
+ last_status = None
+ try:
+ while True:
+ auth.resolve_baton_timeout_if_needed(cfg)
+
+ new_baton_status = auth.get_baton_status(cfg, data)
+ status_json = new_baton_status.model_dump_json()
+
+ if status_json != last_status:
+ last_status = status_json
+ yield f"data: {status_json}\n\n"
+
+ cfg.process_queued_transfer_if_ready(auth.SESSION_EXPIRE_SECONDS)
+
+ await asyncio.sleep(0.5)
+ except asyncio.CancelledError:
+ return
+
+
+@app.get("/sse/baton")
+async def sse_baton(token: str = Depends(oauth2_scheme)):
+ """SSE endpoint for real-time baton status updates."""
+ data = auth.parse_token(token)
+ return StreamingResponse(
+ baton_status_event_stream(data),
+ media_type="text/event-stream",
+ headers={
+ "Cache-Control": "no-cache",
+ "Connection": "keep-alive",
+ "Access-Control-Allow-Origin": "*",
+ "Access-Control-Allow-Headers": "Cache-Control"
+ }
+ )
@app.get("/beamline/settings")
async def get_settings(token: str = Depends(oauth2_scheme)) -> BeamlineSettingsModel:
diff --git a/src/aare/daq/server_exception_handler.py b/src/aare/daq/server_exception_handler.py
index 530864db..7b4307df 100644
--- a/src/aare/daq/server_exception_handler.py
+++ b/src/aare/daq/server_exception_handler.py
@@ -17,7 +17,10 @@ from aare.common.exception_handler import (
AuthenticationException,
SampleException,
UserRightsException,
- SmargonCommunicationError, TellCommunicationError
+ SmargonCommunicationError,
+ TellCommunicationError,
+ JFJochCommunicationError,
+ AerotechCommunicationError,
)
logger = setup_logger("aareDAQ")
@@ -145,9 +148,39 @@ def register_exception_handlers(app) -> None:
),
)
+ @app.exception_handler(JFJochCommunicationError)
+ async def jfjoch_comm_handler(request: Request, exc: JFJochCommunicationError) -> JSONResponse:
+ return JSONResponse(
+ status_code=api_status.HTTP_503_SERVICE_UNAVAILABLE,
+ content=_error_payload(
+ code="JFJOCH_UNAVAILABLE",
+ message=str(exc) or "JFJoch detector is unavailable",
+ extra={
+ "operation": getattr(exc, "operation", None),
+ "endpoint": getattr(exc, "endpoint", None),
+ "base_url": getattr(exc, "base_url", None),
+ },
+ ),
+ )
+
+ @app.exception_handler(AerotechCommunicationError)
+ async def aerotech_comm_handler(request: Request, exc: AerotechCommunicationError) -> JSONResponse:
+ return JSONResponse(
+ status_code=api_status.HTTP_503_SERVICE_UNAVAILABLE,
+ content=_error_payload(
+ code="AEROTECH_UNAVAILABLE",
+ message=str(exc) or "Aerotech is unavailable",
+ extra={
+ "operation": getattr(exc, "operation", None),
+ "endpoint": getattr(exc, "endpoint", None),
+ "base_url": getattr(exc, "base_url", None),
+ },
+ ),
+ )
+
@app.exception_handler(Exception)
async def unhandled_exception_handler(request: Request, exc: Exception) -> JSONResponse:
- logger.exception("Unhandled server exception")
+ logger.exception(f"Unhandled server exception: {exc}")
return JSONResponse(
status_code=api_status.HTTP_500_INTERNAL_SERVER_ERROR,
content=_error_payload(code="INTERNAL_SERVER_ERROR", message=str(exc) or "Internal server error"),
diff --git a/src/aare/daq/spreadsheetupdater.py b/src/aare/daq/spreadsheetupdater.py
index 4f7c65f8..5cad17fb 100644
--- a/src/aare/daq/spreadsheetupdater.py
+++ b/src/aare/daq/spreadsheetupdater.py
@@ -9,7 +9,9 @@ from aare.common.beamline import MXBeamline, mx_beamline
beamline = mx_beamline()
SLOT_IDENTIFIER = beamline.value.upper()
-WS_URL = f"wss://mx-db-01.psi.ch/dispatcher/protected_router/tell_runner/ws/samples-spreadsheet/{SLOT_IDENTIFIER}"
+#WS_URL = f"wss://mx-db-01.psi.ch/dispatcher/protected_router/tell_runner/ws/samples-spreadsheet/{SLOT_IDENTIFIER}"
+WS_URL = f"wss://mx-aaredb-dmz-01.psi.ch/dispatcher/protected_router/tell_runner/ws/samples-spreadsheet/{SLOT_IDENTIFIER}"
+
# Ensure the environment variable for the shared password is set
password = os.getenv("AAREDB_SHARED_PASSWORD")
@@ -134,6 +136,25 @@ def main():
"""
while True:
try:
+ import ssl
+ import websocket
+
+ # FORCE a clean context
+ context = ssl.create_default_context(ssl.Purpose.SERVER_AUTH)
+ context.load_verify_locations(cafile="/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem")
+ context.load_cert_chain(
+ certfile="/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_dmz-01.crt",
+ keyfile="/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_dmz-01.key"
+ )
+
+ # Explicitly set the SNI hostname to match NGINX server_name
+ # This is often what's missing when NGINX says "No cert sent"
+ ssl_opt = {
+ "context": context,
+ "server_hostname": "mx-aaredb-dmz-01.psi.ch",
+ "check_hostname": True
+ }
+
ws = websocket.WebSocketApp(
WS_URL,
header=WS_HEADERS,
@@ -142,11 +163,13 @@ def main():
on_close=on_close,
on_open=on_open,
)
- ws.run_forever(sslopt={"cert_reqs": 0})
- except Exception as e:
- print(f"[MAIN][ERROR] ws.run_forever() crashed with: {e}")
- print("[WS][INFO] WebSocket connection lost. Reconnecting in 5 seconds...")
+ print(f"[WS][INFO] Connecting to {WS_URL}...")
+ ws.run_forever(sslopt=ssl_opt)
+
+ except Exception as e:
+ print(f"[MAIN][ERROR] WebSocket connection failed: {e}")
+
time.sleep(5)
if __name__ == "__main__":
diff --git a/src/aare/daq/tellupdater.py b/src/aare/daq/tellupdater.py
index e9815dec..9792741d 100644
--- a/src/aare/daq/tellupdater.py
+++ b/src/aare/daq/tellupdater.py
@@ -4,7 +4,6 @@ import threading
import websocket
import sseclient
-import requests
import time
from aareDB.models import PuckWithTellPosition
@@ -19,7 +18,8 @@ logger = setup_logger("aareDAQ")
# Configuration
beamline = mx_beamline()
SLOT_IDENTIFIER = beamline.value.upper()
-WS_URL = f"wss://mx-db-01.psi.ch/dispatcher/protected_router/wstell/ws/slot/{SLOT_IDENTIFIER}"
+#WS_URL = f"wss://mx-db-01.psi.ch/dispatcher/protected_router/wstell/ws/slot/{SLOT_IDENTIFIER}"
+WS_URL = f"wss://mx-aaredb-dmz-01.psi.ch/dispatcher/protected_router/wstell/ws/slot/{SLOT_IDENTIFIER}"
#WS_URL = f"wss://localhost:8001/protected_router/wstell/ws/slot/{SLOT_IDENTIFIER}"
WS_HEADERS = [f"X-Shared-Password: {os.getenv('AAREDB_SHARED_PASSWORD')}"]
print(WS_HEADERS)
@@ -39,19 +39,40 @@ def listen_to_sse():
print(f"[SSE][WARN] No TELL URL configured – SSE listener not started. (tell_client.url={tell_client.url})")
return
sse_url = tell_client.url + "/events"
- try:
- #response = requests.get(sse_url, stream=True)
- client = sseclient.SSEClient(sse_url)
- print("[SSE][listen_to_sse] Initial detected pucks fetch on connect")
- handle_tell_change_event()
+ while True:
+ try:
+ print(f"[SSE][INFO] Attempting to connect to {sse_url}...")
+ if sse_url.startswith("https://mx-aaredb-dmz-01"): #"https://mx-db-01"
+ # mTLS path
+ import requests
+ #cert_pair = ("/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_db-01.crt", "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_db-01.key")
+ cert_pair = ("/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_dmz-01.crt",
+ "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_dmz-01.key")
+ #ca_root = "/etc/ssl/certs/secrets/mx-db-01_Full_Chain_CA.pem"
+ ca_root = "/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem"
+ response = requests.get(sse_url, stream=True, cert=cert_pair, verify=ca_root)
+ response.raise_for_status()
+ client = sseclient.SSEClient(response)
+ else:
+ # Robot path (PC17488) - Pass URL STRING directly
+ # SSEClient will handle the simple HTTP GET itself
+ client = sseclient.SSEClient(sse_url)
- for event in client.events():
- print(f"event = {event.event} with data: {event.data}")
- if event.event == "DewarContentUpdate":
- on_sse_event(event)
- except Exception as exc:
- print(f"[SSE][listen_to_sse][ERROR] Failed to connect to {sse_url}: {exc}")
+ print("[SSE][listen_to_sse] Initial detected pucks fetch on connect")
+ handle_tell_change_event()
+
+ # Compatibility: some SSEClient versions are iterable, others expose .events().
+ events_iter = client.events() if hasattr(client, "events") else iter(client)
+ for event in events_iter:
+ # print(f"event = {event.event} with data: {event.data}")
+ if event.event == "DewarContentUpdate":
+ on_sse_event(event)
+ except Exception as exc:
+ print(f"[SSE][listen_to_sse][ERROR] Connection lost or failed: {exc}")
+
+ print("[SSE][INFO] Reconnecting to SSE in 5 seconds...")
+ time.sleep(5)
def compare_and_report_change(old, new, key_func):
"""
@@ -140,14 +161,34 @@ def on_close(ws, close_status_code, close_msg):
def on_open(ws):
print("[WS][OPEN] WebSocket opened.")
+
def main():
# Start SSE listener in a separate background thread
sse_thread = threading.Thread(target=listen_to_sse, daemon=True)
sse_thread.start()
- # Main thread runs websocket client loop
+ """
+ Main function to initiate WebSocket connection with mTLS.
+ """
while True:
try:
+ import ssl
+ # 1. Create a modern SSL Context for a TLS Client
+ context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
+
+ # 2. Load the CA to verify the NGINX server's identity
+ context.load_verify_locations(cafile="/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem")
+
+ # 3. Load the Client Certificate and Key (mTLS)
+ # Using the 'dmz-01' paths that worked in your curl
+ context.load_cert_chain(
+ certfile="/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_dmz-01.crt",
+ keyfile="/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_dmz-01.key"
+ )
+
+ # Optional: Ensure hostname matching is active (recommended)
+ context.check_hostname = True
+
ws = websocket.WebSocketApp(
WS_URL,
header=WS_HEADERS,
@@ -156,12 +197,16 @@ def main():
on_close=on_close,
on_open=on_open,
)
- ws.run_forever(sslopt={"cert_reqs": 0})
- except Exception as e:
- print(f"[MAIN][ERROR] ws.run_forever() crashed with: {e}")
- print("[WS][INFO] WebSocket connection lost. Reconnecting in 5 seconds...")
+ # 4. Pass the context directly via sslopt
+ print(f"[WS][INFO] Connecting to {WS_URL} using mTLS...")
+ ws.run_forever(sslopt={"context": context})
+
+ except Exception as e:
+ print(f"[MAIN][ERROR] WebSocket connection failed: {e}")
+
+ print("[WS][INFO] Reconnecting in 5 seconds...")
time.sleep(5)
if __name__ == "__main__":
- main()
\ No newline at end of file
+ main()
diff --git a/src/aare/daq/workflows.py b/src/aare/daq/workflows.py
index 04addece..094da900 100644
--- a/src/aare/daq/workflows.py
+++ b/src/aare/daq/workflows.py
@@ -34,8 +34,7 @@ def common_2rse(devs: BeamlineDevices, cfg: BeamlineConfig):
print('move smargon home')
devs.smargon_move_home()
print('try to move aerotech')
- devs.aerotech_pos = Coordinate(x=0.0, y=0.0, z=0.0)
- devs.aerotech_omega = 0.0
+ devs.aerotech_pos = ABR_POS_MOUNT
#print('move bl to park')
#devs.reflector_up = StagePositionEnum.PARK
print('move bs to park')
@@ -53,8 +52,7 @@ def sa2se(devs: BeamlineDevices, cfg: BeamlineConfig):
print('move smargon home')
devs.smargon_move_home()
print('try to move aerotech')
- devs.aerotech_pos = Coordinate(x=0.0, y=0.0, z=0.0)
- devs.aerotech_omega = 0.0
+ devs.aerotech_pos = ABR_POS_MOUNT
print('move bl to park')
devs.reflector_up = StagePositionEnum.PARK
print('move bs to park')
diff --git a/src/aare/devices/aerotech.py b/src/aare/devices/aerotech.py
index 5884e40b..463087f7 100644
--- a/src/aare/devices/aerotech.py
+++ b/src/aare/devices/aerotech.py
@@ -1,386 +1,246 @@
-import os
-import sys
-import threading
-import time
-from enum import Enum
-from typing import Union
+from typing import Optional, Union
-import automation1 as a1
-from epics import PV, Motor, caput, caget
+from aarescan_client.models.grid_request import GridRequest
+from aarescan_client.models.screen_request import ScreenRequest
from aare.common.beamline import MXBeamline, mx_beamline
-from aare.devices.my_motor import MyMotor
+from aarescan_client import ApiClient, Status, RotationRequest, Configuration, DefaultApi, AxisStatus, Target
-class TaskEnum(Enum):
- TASK_0 = 0
- TASK_1 = 1
- TASK_2 = 2
- TASK_3 = 3
- TASK_4 = 4
+from aare.common.coordinate import AerotechCoordinate, Coordinate
+from aare.common.exception_handler import AerotechCommunicationError
-class AxisEnum(Enum):
- X = "gmx"
- Y = "gmy"
- Z = "gmz"
- OMEGA = "Omega"
+AEROTECH_HOME = AerotechCoordinate(at_mm=Coordinate(x=0, y=0, z=0), omega_deg=0)
-class AerotechRunEnum(Enum):
- STOP = 0
- START = 1
- RUN = 2
- LOAD = 3
- PAUSE = 4
- RESET = 5
-class VariableTypeEnum(Enum):
- INT = 0
- REAL = 1
- STRING = 2
+class AerotechController(object):
-class AerotechControllerEpics:
- def __init__(self, beamline: MXBeamline):
- BEAMLINE = beamline.value.upper()
- self.aerotech_pv_prefix = f"{BEAMLINE}-ES-DF1"
- trx_prefix = f"{self.aerotech_pv_prefix}:TRX1"
- try_prefix = f"{self.aerotech_pv_prefix}:TRY1"
- trz_prefix = f"{self.aerotech_pv_prefix}:TRZ1"
- rotu_prefix = f"{self.aerotech_pv_prefix}:ROTU"
- self.gmx = MyMotor(trx_prefix)
- self.gmy = MyMotor(try_prefix)
- self.gmz = MyMotor(trz_prefix)
- self.omega = MyMotor(rotu_prefix)
- self.__enable_all = PV(f"{self.aerotech_pv_prefix}:EnableAll")
- self.__disable_all = PV(f"{self.aerotech_pv_prefix}:DisableAll")
- self.__acknowledge_all = PV(f"{self.aerotech_pv_prefix}:AckAll")
- self.__stop_all = PV(f"{self.aerotech_pv_prefix}:StopAll")
- self.__task_filename = PV(f"{self.aerotech_pv_prefix}:TASK:FILENAME")
- self.__task_id = PV(f"{self.aerotech_pv_prefix}:TASK:TASKIDX") # values can be 1 to 4 DO NOT USE 0!!!!
- self.__task_run_enum = PV(f"{self.aerotech_pv_prefix}:TASK:SWITCH") # 0 Stop, 1 Start, 2 Run,
- # 3 Load, 4 Pause, 5 Reset
- self.enable_all()
+ def __init__(self, bl: MXBeamline):
+ if bl == MXBeamline.X06DA:
+ self.__simulated = False
+ self.__base = "http://x06da-smargopolo.psi.ch:5234"
+ elif bl == MXBeamline.X10SA:
+ self.__simulated = False
+ self.__base = "http://x10sa-smargopolo.psi.ch:5234"
+ elif bl == MXBeamline.SIMULATED:
+ self.__simulated = True
+ self.__pos = AEROTECH_HOME
+ self.__vel = 0
+ else:
+ raise Exception("unknown beamline")
- def enable_all(self):
- self.__enable_all.put(1)
- self.__enable_all.put(0)
+ if not self.__simulated:
+ self.__client = ApiClient(Configuration(host=self.__base))
+ self.__api = DefaultApi(self.__client)
- def acknowledge_all(self):
- self.__acknowledge_all.put(1)
- self.__acknowledge_all.put(0)
+ def __make_aerotech_target(
+ self,
+ coord: AerotechCoordinate,
+ wait: bool = False,
+ incremental: bool = False,
+ ) -> Target:
+ at_mm = coord.at_mm
- def _disable_all(self):
- self.__disable_all.put(1)
- self.__disable_all.put(0)
+ return Target(
+ x=at_mm.x if at_mm is not None else None,
+ y=at_mm.y if at_mm is not None else None,
+ z=at_mm.z if at_mm is not None else None,
+ u=coord.omega_deg,
+ run_async=not wait,
+ incremental=incremental,
+ )
- def stop_all(self):
- self.__stop_all.put(1)
- self.__stop_all.put(0)
+ def __make_aerotech_coordinate(self, target: Target) -> AerotechCoordinate:
+ return AerotechCoordinate(at_mm=Coordinate(x=target.x,y=target.y,z=target.z), omega_deg=target.u)
- def __task_stop(self):
- self.__task_run_enum.put(0)
-
- def __task_start(self):
- self.__task_run_enum.put(1)
-
- def __task_run(self):
- self.__task_run_enum.put(2)
-
- def __task_pause(self):
- self.__task_run_enum.put(4)
-
- def __task_load(self):
- self.__task_run_enum.put(3)
-
- def __task_reset(self):
- self.__task_run_enum.put(5)
-
- def __set_task_id(self, task_id: int):
- if task_id not in range(1, 5):
- raise ValueError(f"Invalid task id {task_id}")
- self.__task_id.put(task_id)
-
- def __put(self, axis: MyMotor, attr: str, value: float):
+ def cancel(self):
try:
- axis.put(attr=attr, value=value)
+ self.__api.cancel_post()
except Exception as e:
- raise ValueError(f"Error setting {axis.name} {attr} to {value}: {e}")
+ raise AerotechCommunicationError(
+ "Aerotech cancel failed",
+ endpoint="cancel_post",
+ base_url=self.__base,
+ operation="POST",
+ ) from e
- def set_offset(self, axis: MyMotor, offset: float):
- self.__put(axis=axis, attr="OFF", value=offset)
-
- def home_all(self, task_id: int = 3):
- #self.__task_id.put(2)
- self.__task_reset()
- time.sleep(0.2)
- self.__task_id.put(task_id)
- print(self.__task_id.get())
- time.sleep(0.2)
- self.__task_filename.put("home_all.a1exe")
- print(self.__task_filename.get())
- # self.__task_load()
- #time.sleep(0.2)
- #print(self.__task_run_enum.get())
- time.sleep(1.0)
- self.__task_run()
- print(self.__task_run_enum.get())
-
- def get_tast_status(self, task_id: int = 3):
- task_status = PV(f"{self.aerotech_pv_prefix}:TASK:T{task_id}:STATUS")
- return task_status.get(as_string=True)
-
-
- def __set_global(self, index:int, value: Union[int, float, str],
- timeout:float=10.0):
- """Set a global variable in aerotech, it takes ~200 ms for the value to be set
- :param index: select a variable to change. an integer between 0..256 for int and real variable for 0..31 for strings
- :param value: the value to set can be int, float (REAL) or string
- :param timeout: timeout for vairable change, default 10.0 seconds """
- var_type = self.__get_var_type(value)
- caput(f"{self.aerotech_pv_prefix}:VAR:{var_type.upper()}-ADDR", index)
- if var_type == 'STRING':
- var_type = 'STRING-SHORT'
- caput(f"{self.aerotech_pv_prefix}:VAR:{var_type.upper()}", value)
- start = time.perf_counter()
- while time.perf_counter() - start < timeout:
- rbv = self.__read_global_from_index(index, var_type)
- if rbv == str(value):
- return
- time.sleep(0.1)
-
- raise TimeoutError(f"Timeout setting global variable {index} to {value}")
-
-
- def __read_global_feedback(self, index, var_type: str):
- caput(f"{self.aerotech_pv_prefix}:VAR:{var_type.upper()}-RBV.PROC", 1)
- time.sleep(0.5)
- return caget(f"{self.aerotech_pv_prefix}:VAR:{var_type.upper()}-RBV", as_string=True)
-
- def __read_global_from_index(self, index, var_type: str):
- return caget(f"{self.aerotech_pv_prefix}:VAR:{var_type.upper()}{index}_RBV", as_string=True)
-
- def __get_var_type(self, value):
- if type(value) is int:
- return VariableTypeEnum.INT.name
- elif type(value) is float:
- return VariableTypeEnum.REAL.name
- elif type(value) is str:
- return VariableTypeEnum.STRING.name
- else:
- raise ValueError(f"Invalid type {type(value)} for global variable")
-
- def get_global_variable(self, index: int, var_type:VariableTypeEnum):
- return self.__read_global_from_index(index, var_type.name)
-
- def set_global_variable(self, index: int, value: Union[int, float, str]):
- self.__set_global(index, value)
-
-
-class AerotechController:
- def __init__(self, controller_ip: str):
- if controller_ip is None:
- self.controller = None
- else:
- self.controller = a1.Controller.connect(controller_ip)
-
- self.status_item_configuration = a1.StatusItemConfiguration()
- self.start_controller()
-
- def start_controller(self):
- self.controller.start()
-
- def disconnect(self):
- self.controller.disconnect()
-
- def enable_motion(self, axis: str):
- self.controller.runtime.commands.motion.enable(axis.upper())
-
- def home_motor(self, axis: str):
- self.__configure_axis_status(axis, a1.AxisStatusItem.AxisEnabled)
- result = self.get_status_via_status_items(name=axis, status_item=a1.AxisStatusItem.AxisEnabled)
- if not result == a1.AxisStatusItem.AxisEnabled:
- self.enable_motion(axis)
- self.controller.runtime.commands.motion.home(axis.upper())
-
- def __configure_axis_status(self, axis_name: str, axis_status_item: a1.AxisStatusItem):
- self.status_item_configuration.axis.add(axis_status_item=axis_status_item, axis=axis_name)
-
- def __configure_task_status(self, task_id:int, task_status_item: a1.TaskStatusItem):
- self.status_item_configuration.task.add(task_status_item=task_status_item, task=f"Task {task_id}")
-
- def get_status_via_status_items(self, name: str | int, status_item: a1.TaskStatusItem | a1.AxisStatusItem):
- result = self.controller.runtime.status.get_status_items(self.status_item_configuration)
- if int(name):
- return result.task.get(status_item, f"Task {name}").value
- elif str(name):
- return result.axis.get(status_item, name).value
- else:
- raise ValueError(f"Invalid status item {status_item} for task {name}")
-
- def set_global_variable(self, index:int, value: Union[int, float, str]):
- if type(value) is int:
- self.controller.runtime.variables.global_.set_integer(index, value)
- elif type(value) is float:
- self.controller.runtime.variables.global_.set_real(index, value)
- elif type(value) is str:
- self.controller.runtime.variables.global_.set_string(index, value)
- else:
- raise ValueError(f"Invalid type {type(value)} for global variable")
-
- def get_axis_status_via_status_items(self, axis_name:str, status_item: a1.AxisStatusItem):
- result = self.controller.runtime.status.get_status_items(self.status_item_configuration)
- return result.task.get(status_item, axis_name).value
-
- def wait_for_status_to_change(self, name: str | int, status_item:a1.TaskStatusItem | a1.AxisStatusItem, enum, timeout: float = 60.0):
- start = time.perf_counter()
- while self.get_status_via_status_items(name, status_item) == enum:
- time.sleep(0.1)
- if time.perf_counter() - start > timeout:
- print(self.get_status_via_status_items(name, status_item))
- raise TimeoutError(f"Timeout waiting for task {name} to finish")
- return
-
- def wait_program_finish(self, task_id:int=3, timeout:float=60.0):
- start = time.perf_counter()
- status_item = a1.TaskStatusItem.TaskState
- while True:
- status = self.get_status_via_status_items(task_id, status_item)
- if status != a1.TaskState.ProgramRunning:
- if status == a1.TaskState.Idle:
- return
- elif status == a1.TaskState.ProgramComplete:
- print(f"Program completed on Task {task_id}")
- return
- elif status == a1.TaskState.Error:
- raise RuntimeError(f"Task {task_id} failed to start")
- elif status == a1.TaskState.ProgramPaused:
- print(f"Program paused on Task {task_id}, not sure how")
- else:
- raise RuntimeError(f"Unknown status {status} for task {task_id}")
- if time.perf_counter() - start > timeout:
- print(self.get_status_via_status_items(task_id, status_item))
- raise TimeoutError(f"Timeout waiting for task {task_id} to finish")
- time.sleep(0.05)
- print(f"Task {task_id} finished with status {status}")
-
- def __run_program(self, script_name:str, task_id:int=3, timeout:float=60.0):
- self.__configure_task_status(task_id, a1.TaskStatusItem.TaskState)
- state = self.get_status_via_status_items(task_id, a1.TaskStatusItem.TaskState)
- print(f"Task {task_id} is in state {state}: {a1.TaskState(state).name}")
- if state != a1.TaskState.ProgramComplete and state != a1.TaskState.Idle:
- if a1.TaskState.ProgramRunning == state:
- self.wait_for_status_to_change(task_id, a1.TaskStatusItem.TaskState, state, timeout)
- elif a1.TaskState.ProgramComplete == state:
- print('can continue')
- else:
- print(f"Task {task_id} is not Idle: {a1.TaskState(state).name} , aborting")
- return
+ def is_idle(self) -> bool:
try:
- print(f"Running script {script_name} on task {task_id}")
- self.controller.runtime.tasks[task_id].program.run(script_name)
- self.wait_program_finish(task_id, timeout)
+ status = self.__api.status_get()
+ return status.state == 'Idle'
except Exception as e:
- print(f"Error executing script {script_name}: {e}")
+ raise AerotechCommunicationError(
+ "Aerotech status check failed",
+ endpoint="status_get",
+ base_url=self.__base,
+ operation="GET",
+ ) from e
- def run_grid_scan(self, cell_height_mm:float, num_rows:int,
- row_width_mm:float, time_per_row_s: float, task_id:int =3):
- self.set_global_variable(0, cell_height_mm)
- self.set_global_variable(1, row_width_mm)
- self.set_global_variable(2, time_per_row_s)
- self.set_global_variable(1, num_rows)
- self.set_global_variable(0, 1)
- #self.__run_program(script_name="grid_scan.a1exe", task_id=task_id, timeout=120.0)
- def home_all(self, task_id:int = 3):
- self.__run_program(script_name="home_all.a1exe", task_id=task_id, timeout=120.0)
+ def get_position(self) -> AerotechCoordinate:
+ status = self.status()
+ return AerotechCoordinate(
+ at_mm=Coordinate(
+ x=status.x.pos,
+ y=status.y.pos,
+ z=status.z.pos,
+ ),
+ omega_deg=status.u.pos,
+ )
- def rotation_scan(self, task_id = 3, start_angle:float = 0.0, end_angle:float = 360.0, step_size:float = 1.0, num_steps:int = 10):
- self.__run_program(script_name="rotation_scan.a1exe", task_id=task_id, timeout=90.0)
+ def status(self) -> Status:
+ if self.__simulated:
+ return Status(state=Status.State.IDLE,
+ x=AxisStatus(pos=self.__pos.x, vel=self.__vel,
+ enabled=False, homed=False, moving=False,fault=False),
+ y=AxisStatus(pos=self.__pos.y, vel=self.__vel,
+ enabled=False, homed=False, moving=False,fault=False),
+ z=AxisStatus(pos=self.__pos.z, vel=self.__vel,
+ enabled=False, homed=False, moving=False,fault=False),
+ u=AxisStatus(pos=self.__pos.u, vel=self.__vel,
+ enabled=False, homed=False, moving=False,fault=False),
+ )
+ try:
+ return self.__api.status_get()
+ except Exception as e:
+ raise AerotechCommunicationError(
+ "Aerotech status request failed",
+ endpoint="status_get",
+ base_url=self.__base,
+ operation="GET",
+ ) from e
- def move_motor_absolute(self, axis:str, position:float, speed:float=1.0):
- self.controller.runtime.commands.motion.moveabsolute(axis.upper(), [position], [speed])
+ def move_home(self, wait:bool=True, incremental:bool=False):
+ if self.__simulated:
+ self.__pos = AEROTECH_HOME
+ return self.__pos
+ return self.position(AEROTECH_HOME, wait=wait, incremental=incremental)
+
+ def home_aerotech(self):
+ try:
+ return self.__api.home_post()
+ except Exception as e:
+ raise AerotechCommunicationError(
+ "Aerotech home failed",
+ endpoint="home_post",
+ base_url=self.__base,
+ operation="POST",
+ ) from e
+
+ def wait_till_done(self, timeout=60):
+ try:
+ return self.__api.wait_till_done_post(timeout=timeout)
+ except Exception as e:
+ raise AerotechCommunicationError(
+ "Aerotech wait_till_done failed",
+ endpoint="wait_till_done_post",
+ base_url=self.__base,
+ operation="POST",
+ ) from e
+
+ def position(
+ self,
+ target: AerotechCoordinate,
+ /,
+ wait: bool = True,
+ incremental: bool = False,
+ ):
+ if self.__simulated:
+ return self.__pos
+
+ payload = self.__make_aerotech_target(target, wait=wait, incremental=incremental)
+ try:
+ return self.__api.position_post(payload)
+ except Exception as e:
+ raise AerotechCommunicationError(
+ "Aerotech position move failed",
+ endpoint="position_post",
+ base_url=self.__base,
+ operation="POST",
+ ) from e
+
+ def rotation_scan(self, rotation_deg: float | int,
+ time_sec: float | int,
+ start_pos_deg: float | int,
+ run_async: bool = False):
+ payload = RotationRequest(
+ rotation_deg=rotation_deg,
+ time_sec=time_sec,
+ start_pos_deg=start_pos_deg,
+ run_async=run_async
+ )
+ if self.__simulated:
+ return payload
+
+ try:
+ return self.__api.rotation_scan_post(payload)
+ except Exception as e:
+ raise AerotechCommunicationError(
+ "Aerotech rotation scan failed",
+ endpoint="rotation_scan_post",
+ base_url=self.__base,
+ operation="POST",
+ ) from e
+
+
+ def grid_scan(self,
+ grid_elem_count_y: int,
+ grid_elem_size_y_um: int | float,
+ time_sec: int|float,
+ grid_elem_size_x_um: Optional[Union[float, int]] = None,
+ grid_elem_count_x: Optional[int] = None,
+ run_async: Optional[bool] = False
+
+ ):
+ payload = GridRequest(
+ grid_elem_count_x=grid_elem_count_x,
+ grid_elem_count_y=grid_elem_count_y,
+ grid_elem_size_x_um=grid_elem_size_x_um,
+ grid_elem_size_y_um=grid_elem_size_y_um,
+ time_sec=time_sec,
+ run_async=run_async,
+ )
+ if self.__simulated:
+ return payload
+ try:
+ return self.__api.grid_scan_post(payload)
+ except Exception as e:
+ raise AerotechCommunicationError(
+ "Aerotech grid scan failed",
+ endpoint="grid_scan_post",
+ base_url=self.__base,
+ operation="POST",
+ ) from e
+
+ def screening_scan(self,
+ rotation_deg: float | int,
+ wedge_deg: float | int,
+ time_sec: float | int,
+ steps: int,
+ run_async: bool = False
+ ):
+ payload = ScreenRequest(
+ rotation_deg=rotation_deg,
+ wedge_deg=wedge_deg,
+ time_sec=time_sec,
+ steps=steps,
+ run_async=run_async
+ )
+ if self.__simulated:
+ return payload
+ try:
+ return self.__api.screening_post(payload)
+ except Exception as e:
+ raise AerotechCommunicationError(
+ "Aerotech screening scan failed",
+ endpoint="screening_post",
+ base_url=self.__base,
+ operation="POST",
+ ) from e
- def move_motor_linear(self, axis: str, position: float, speed: float = 1.0):
- self.controller.runtime.commands.motion.movelinear(axis.upper(), [position], speed)
if __name__ == "__main__":
beamline = mx_beamline()
- print(beamline)
- ### test on 10S
- aerotech = AerotechController(controller_ip="129.129.118.96")
- #aerotech.enable_motion("X")
- #aerotech.home_all()
- rw = 0.320
- ch = 0.010
- nr = 10
- tpr_s = rw / ch * 0.02
- #aerotech.run_grid_scan(cell_height_mm=ch, num_rows=nr,
- # row_width_mm=rw, time_per_row_s=tpr_s)
- st = time.perf_counter()
- aerotech.move_motor_absolute("Z", 0, 10000)
- print(f"time to move: {time.perf_counter() - st}")
- aerotech.disconnect()
-
- #aerotech_epics = AerotechControllerEpics(beamline)
- #aerotech_epics.omega.speed = 80.0
- #print(aerotech_epics.get_global_variable(0, VariableTypeEnum.REAL))
- #aerotech_epics.set_global_variable(0, 100.0)
- #print(aerotech_epics.get_global_variable(0, VariableTypeEnum.REAL))
- # print('enable all motors')
- # aerotech_epics.enable_all()
- # print('home_all')
- # aerotech_epics.home_all()
- # print('wait for home to finish')
- # start = time.perf_counter()
- # status = aerotech_epics.get_tast_status(2)
- # print(f'current status: {status}')
- # if status == 'Idle':
- # time.sleep(0.2)
- # test_counter = 0
- # while status != 'Ready':
- # time.sleep(0.1)
- # if time.perf_counter() - start > 360.0:
- # raise TimeoutError("Timeout waiting for home all task to finish")
- # elif status == 'Idle':
- # time.sleep(0.5)
- # if test_counter == 1:
- # raise RuntimeError("Home all task failed to start")
- # time.sleep(1.0)
- # print('restarting home all')
- # aerotech_epics.home_all()
- # time.sleep(1.0)
- # test_counter = 1
- #
- # status = arotech_epics.get_tast_status(2)
- # for i in range(10):
- # print(f'moving to {(i+1)*90}')
- # aerotech_epics.omega.move(90.0, relative=True, wait=True)
- # time.sleep(1.0)
- # if i == 5:
- # print('simulating disable')
- # aerotech_epics._disable_all()
- # time.sleep(10.0)
- # print('re-enabling')
- # aerotech_epics.enable_all()
- # time.sleep(1.0)
- # print('home all')
- # aerotech_epics.home_all()
- # print('wait for home to finish')
- # start = time.perf_counter()
- # status = aerotech_epics.get_tast_status(2)
- # print(f'current status: {status}')
- # test_counter = 0
- # while status != 'Ready':
- # time.sleep(0.1)
- # if time.perf_counter() - start > 360.0:
- # raise TimeoutError("Timeout waiting for home all task to finish")
- # status = aerotech_epics.get_tast_status(2)
- # time.sleep(0.5)
- # if test_counter == 1:
- # raise RuntimeError("Home all task failed to start")
- # time.sleep(1.0)
- # print('restarting home all')
- # aerotech_epics.home_all()
- # time.sleep(1.0)
- # test_counter = 1
-
-
- #aerotech_epics.home_all()
-
+ #print(beamline)
+ controller = AerotechController(beamline)
+ #controller.print_status(colored=True, compact=True)
+ print(controller.get_position())
+ controller.cancel()
diff --git a/src/aare/devices/automation1_python_api.py b/src/aare/devices/automation1_python_api.py
new file mode 100644
index 00000000..f3ca8c8f
--- /dev/null
+++ b/src/aare/devices/automation1_python_api.py
@@ -0,0 +1,187 @@
+import os
+import sys
+import threading
+import time
+from enum import Enum
+from typing import Union
+
+import automation1 as a1
+
+from aare.common.beamline import MXBeamline, mx_beamline
+
+class AerotechController:
+ def __init__(self, controller_ip: str):
+ # if controller_ip is None:
+ # self.controller = None
+ # else:
+ self.controller = a1.Controller.connect(controller_ip)
+
+ self.status_item_configuration = a1.StatusItemConfiguration()
+ self.start_controller()
+
+ def start_controller(self):
+ self.controller.start()
+
+ def disconnect(self):
+ self.controller.disconnect()
+
+ def enable_motion(self, axis: str):
+ self.controller.runtime.commands.motion.enable(axis.upper())
+
+ def home_motor(self, axis: str):
+ self.__configure_axis_status(axis, a1.AxisStatusItem.AxisEnabled)
+ result = self.get_status_via_status_items(name=axis, status_item=a1.AxisStatusItem.AxisEnabled)
+ if not result == a1.AxisStatusItem.AxisEnabled:
+ self.enable_motion(axis)
+ self.controller.runtime.commands.motion.home(axis.upper())
+
+ def __configure_axis_status(self, axis_name: str, axis_status_item: a1.AxisStatusItem):
+ self.status_item_configuration.axis.add(axis_status_item=axis_status_item, axis=axis_name)
+
+ def __configure_task_status(self, task_id:int, task_status_item: a1.TaskStatusItem):
+ self.status_item_configuration.task.add(task_status_item=task_status_item, task=f"Task {task_id}")
+
+ def get_status_via_status_items(self, name: str | int, status_item: a1.TaskStatusItem | a1.AxisStatusItem):
+ result = self.controller.runtime.status.get_status_items(self.status_item_configuration)
+ if int(name):
+ return result.task.get(status_item, f"Task {name}").value
+ elif str(name):
+ return result.axis.get(status_item, name).value
+ else:
+ raise ValueError(f"Invalid status item {status_item} for task {name}")
+
+ def set_global_variable(self, index:int, value: Union[int, float, str]):
+ if type(value) is int:
+ self.controller.runtime.variables.global_.set_integer(index, value)
+ elif type(value) is float:
+ self.controller.runtime.variables.global_.set_real(index, value)
+ elif type(value) is str:
+ self.controller.runtime.variables.global_.set_string(index, value)
+ else:
+ raise ValueError(f"Invalid type {type(value)} for global variable")
+
+ def get_axis_status_via_status_items(self, axis_name:str, status_item: a1.AxisStatusItem):
+ result = self.controller.runtime.status.get_status_items(self.status_item_configuration)
+ return result.task.get(status_item, axis_name).value
+
+ def wait_for_status_to_change(self, name: str | int, status_item:a1.TaskStatusItem | a1.AxisStatusItem, enum, timeout: float = 60.0):
+ start = time.perf_counter()
+ while self.get_status_via_status_items(name, status_item) == enum:
+ time.sleep(0.1)
+ if time.perf_counter() - start > timeout:
+ print(self.get_status_via_status_items(name, status_item))
+ raise TimeoutError(f"Timeout waiting for task {name} to finish")
+ return
+
+ def wait_program_finish(self, task_id:int=3, timeout:float=60.0):
+ start = time.perf_counter()
+ status_item = a1.TaskStatusItem.TaskState
+ while True:
+ status = self.get_status_via_status_items(task_id, status_item)
+ if status != a1.TaskState.ProgramRunning:
+ if status == a1.TaskState.Idle:
+ return
+ elif status == a1.TaskState.ProgramComplete:
+ print(f"Program completed on Task {task_id}")
+ return
+ elif status == a1.TaskState.Error:
+ raise RuntimeError(f"Task {task_id} failed to start")
+ elif status == a1.TaskState.ProgramPaused:
+ print(f"Program paused on Task {task_id}, not sure how")
+ else:
+ raise RuntimeError(f"Unknown status {status} for task {task_id}")
+ if time.perf_counter() - start > timeout:
+ print(self.get_status_via_status_items(task_id, status_item))
+ raise TimeoutError(f"Timeout waiting for task {task_id} to finish")
+ time.sleep(0.05)
+ print(f"Task {task_id} finished with status {status}")
+
+ def __run_program(self, script_name:str, task_id:int=3, timeout:float=60.0):
+ self.__configure_task_status(task_id, a1.TaskStatusItem.TaskState)
+ state = self.get_status_via_status_items(task_id, a1.TaskStatusItem.TaskState)
+ print(f"Task {task_id} is in state {state}: {a1.TaskState(state).name}")
+ if state != a1.TaskState.ProgramComplete and state != a1.TaskState.Idle:
+ if a1.TaskState.ProgramRunning == state:
+ self.wait_for_status_to_change(task_id, a1.TaskStatusItem.TaskState, state, timeout)
+ elif a1.TaskState.ProgramComplete == state:
+ print('can continue')
+ else:
+ print(f"Task {task_id} is not Idle: {a1.TaskState(state).name} , aborting")
+ return
+ try:
+ print(f"Running script {script_name} on task {task_id}")
+ self.controller.runtime.tasks[task_id].program.run(script_name)
+ self.wait_program_finish(task_id, timeout)
+ except Exception as e:
+ print(f"Error executing script {script_name}: {e}")
+
+ def run_grid_scan(self, cell_height_mm:float, num_rows:int,
+ row_width_mm:float, time_per_row_s: float, task_id:int =3):
+ self.set_global_variable(0, cell_height_mm)
+ self.set_global_variable(1, row_width_mm)
+ self.set_global_variable(2, time_per_row_s)
+ self.set_global_variable(1, num_rows)
+ self.set_global_variable(0, 1)
+ #self.__run_program(script_name="grid_scan.a1exe", task_id=task_id, timeout=120.0)
+ def home_all(self, task_id:int = 3):
+ self.__run_program(script_name="home_all.a1exe", task_id=task_id, timeout=120.0)
+
+ def rotation_scan(self, task_id = 3, start_angle:float = 0.0, end_angle:float = 360.0, step_size:float = 1.0, num_steps:int = 10):
+ self.__run_program(script_name="rotation_scan.a1exe", task_id=task_id, timeout=90.0)
+
+ def move_motor_absolute(self, axis:str, position:float, speed:float=1.0):
+ self.controller.runtime.commands.motion.moveabsolute(axis.upper(), [position], [speed])
+
+ def move_motor_linear(self, axis: str, position: float, speed: float = 1.0):
+ self.controller.runtime.commands.motion.movelinear(axis.upper(), [position], speed)
+
+if __name__ == "__main__":
+ beamline = mx_beamline()
+ print(beamline)
+ ### test on 10S
+ aerotech = AerotechController(controller_ip="129.129.118.96")
+ #aerotech.enable_motion("X")
+ #aerotech.home_all()
+ rw = 0.320
+ ch = 0.010
+ nr = 10
+ tpr_s = rw / ch * 0.02
+ #aerotech.run_grid_scan(cell_height_mm=ch, num_rows=nr,
+ # row_width_mm=rw, time_per_row_s=tpr_s)
+ # st = time.perf_counter()
+ # aerotech.move_motor_absolute("Z", 0, 10000)
+ # print(f"time to move: {time.perf_counter() - st}")
+ # aerotech.disconnect()
+ controller = a1.Controller.connect("129.129.118.96")
+ status_item_configuration = a1.StatusItemConfiguration()
+ controller.start()
+
+ axis = "X"
+
+ pso_input = a1.PsoWindowInput.iXC4ePrimaryFeedback
+ window_number = 0
+ reverse_direction = False
+ execution_task_index = 3
+ min_x_mm = 0
+ max_x_mm = 1
+
+ counts_per_unit = controller.runtime.parameters.axes[axis].units.countsperunit.value
+ print(f"Counts per unit for {axis}: {counts_per_unit}")
+ units_to_counts = controller.runtime.commands.utility_and_conversion.unitstocounts(axis,5,execution_task_index)
+ print(f"Units to counts for {axis}: {units_to_counts}")
+ print(dir(controller.runtime.commands))
+
+ primary_emulated_quadrature_divider = controller.runtime.parameters.axes[axis].feedback.primaryemulatedquadraturedivider.value
+ pso_lower_bound = round(controller.runtime.commands.utility_and_conversion.unitstocounts(axis,min_x_mm,execution_task_index)/primary_emulated_quadrature_divider)
+ pso_upper_bound = round(controller.runtime.commands.utility_and_conversion.unitstocounts(axis,max_x_mm,execution_task_index)/primary_emulated_quadrature_divider)
+
+ controller.runtime.commands.pso.psoreset(axis)
+ controller.runtime.commands.pso.psowindowconfigureinput(axis,0, pso_input, True, execution_task_index)
+ controller.runtime.commands.pso.psowindowconfigurefixedrange(axis,window_number,pso_lower_bound,pso_upper_bound,execution_task_index)
+ controller.runtime.commands.pso.psowindowoutputon(axis,window_number, execution_task_index)
+ controller.runtime.commands.motion.movelinear(axis,[1],0.1,execution_task_index)
+ controller.runtime.commands.motion.waitformotiondone(axis,execution_task_index)
+ controller.runtime.commands.motion.movelinear(axis,[-1],0.1,execution_task_index)
+ controller.runtime.commands.motion.waitformotiondone(axis,execution_task_index)
+ controller.runtime.commands.pso.psowindowoutputoff(axis,window_number, execution_task_index)
+ controller.runtime.commands.pso.psoreset(axis,execution_task_index)
\ No newline at end of file
diff --git a/src/aare/devices/jfjoch.py b/src/aare/devices/jfjoch.py
index 964cd761..e0671e8b 100644
--- a/src/aare/devices/jfjoch.py
+++ b/src/aare/devices/jfjoch.py
@@ -3,6 +3,7 @@ import math
import jfjoch_client
from aare.common.beamline import MXBeamline
+from aare.common.exception_handler import JFJochCommunicationError
from aare.common.models import DAQStatusModel, FluorescenceSpectrumOutputModel
from aare.common.raster_grid import RasterGridRequest
from aare.common.rotation_scan import RotationScanRequest
@@ -10,6 +11,7 @@ from aare.common.rotation_scan import RotationScanRequest
class JFJochWrapper:
def __init__(self, bl: MXBeamline):
+ self.__simulated = False
match bl:
case MXBeamline.X06DA:
self.__url = "http://sls-gpu-001:8080"
@@ -17,6 +19,10 @@ class JFJochWrapper:
self.__url = "http://sls-gpu-002:8080"
case MXBeamline.SIMULATED:
self.__url = "http://localhost:8080"
+ self.__client = None
+ self.__api = None
+ self.__simulated = True
+ return
case _:
raise Exception("unknown beamline")
@@ -93,7 +99,18 @@ class JFJochWrapper:
detect_ice_rings=True,
xray_fluorescence_spectrum=xrf
)
- self.__api.start_post(dataset_settings=dataset_settings)
+ try:
+ self.__api.start_post(dataset_settings=dataset_settings)
+ except Exception as e:
+ scan_type = "rotation"
+ if r.screening:
+ scan_type = "screening"
+ raise JFJochCommunicationError(
+ f"JFJoch data collection failed to initialize for {scan_type} scan",
+ operation="POST",
+ endpoint="start_post",
+ base_url=self.__url,
+ ) from e
def measure_raster(self,
r: RasterGridRequest,
@@ -141,11 +158,29 @@ class JFJochWrapper:
max_spot_count = 1000,
detect_ice_rings = True
)
- self.__api.start_post(dataset_settings=dataset_settings)
+ try:
+ self.__api.start_post(dataset_settings=dataset_settings)
+ except Exception as e:
+ raise JFJochCommunicationError(
+ "JFJoch data collection failed to initialize for raster scan",
+ operation="POST",
+ endpoint="start_post",
+ base_url=self.__url,
+ ) from e
- def wait_till_done(self, timeout : int | float) -> jfjoch_client.models.ScanResult:
- self.__api.wait_till_done_post_with_http_info(timeout=math.ceil(timeout))
- return self.__api.result_scan_get()
+ def wait_till_done(self, timeout : int | float) -> jfjoch_client.models.ScanResult | None:
+ if self.__simulated:
+ return None
+ try:
+ self.__api.wait_till_done_post_with_http_info(timeout=math.ceil(timeout))
+ return self.__api.result_scan_get()
+ except Exception as e:
+ raise JFJochCommunicationError(
+ "JFJoch wait/result retrieval failed",
+ operation="GET",
+ endpoint="wait_till_done_post / result_scan_get",
+ base_url=self.__url,
+ ) from e
def detector(self) -> jfjoch_client.models.DetectorListElement:
l = self.__api.config_select_detector_get()
diff --git a/src/aare/devices/tell_backend.py b/src/aare/devices/tell_backend.py
new file mode 100644
index 00000000..0a6f962e
--- /dev/null
+++ b/src/aare/devices/tell_backend.py
@@ -0,0 +1,410 @@
+import json
+import re
+import time
+from typing import Any, Callable, Protocol
+from urllib.parse import urlparse
+
+import requests
+
+from aare.common.exception_handler import TellCommunicationError
+from aare.common.logger_config import setup_logger
+
+from aare.common.beamline import MXBeamline # noqa: F401
+from pshell import PShellClient
+
+
+logger = setup_logger("aareDAQ")
+
+
+VALID_DEWAR_POSITIONS = [f"{p}{n}" for n in "12345" for p in "ABCDEFX"]
+
+
+def is_valid_dewar_position(position):
+ """check if argument is a valid dewar position"""
+ return position in VALID_DEWAR_POSITIONS
+
+
+POSITION_PARK = "pPark"
+POSITION_COLD = "pCold"
+POSITION_AUX = "pAux"
+POSITION_DEWAR = "pDewar"
+POSITION_HOME = "pHome"
+POSITION_HEATER = "pHeatB"
+
+
+class ManualMountException(Exception):
+ """Custom exception for manual mounting"""
+ pass
+
+
+class SmartMagnetFaultException(Exception):
+ """Custom exception for smart magnet fault"""
+ pass
+
+
+class TellMountFailedException(Exception):
+ """Custom exception for mount failure"""
+ pass
+
+
+class TellCommandWhileBusyException(Exception):
+ """Custom exception for trying to move Tell when it is busy"""
+ pass
+
+
+class TellConnectionException(Exception):
+ """Custom exception for connection problems"""
+ pass
+
+
+class TellBackend(Protocol):
+ @property
+ def url(self) -> str | None:
+ ...
+
+ def get_state(self) -> str:
+ ...
+
+ def get_result(self, command_id: int = -1):
+ ...
+
+ def wait_state(self, state: str, timeout: float) -> None:
+ ...
+
+ def wait_state_not(self, state: str, timeout: float) -> None:
+ ...
+
+ def wait_events(self, events: dict[str, Any], timeout: float):
+ ...
+
+ def eval(self, expr: str):
+ ...
+
+ def start_eval(self, expr: str) -> int:
+ ...
+
+ def run(self, path: str, pars: list[str] | None = None, background: bool = False) -> None:
+ ...
+
+ def abort(self) -> None:
+ ...
+
+
+class PShellTellBackend:
+ def __init__(self, bl: MXBeamline):
+ self._url = self._resolve_url(bl)
+
+ print(f"Connecting TELL p-shell service at {self._url} ...", end="")
+ hostname = urlparse(self._url).hostname
+ try:
+ requests.get(f"{self._url}/history/0", timeout=1.0)
+ except requests.exceptions.RequestException as e:
+ print(f"...connection to {hostname} failed")
+ raise TellCommunicationError(
+ f"TELL connection failed ({hostname})",
+ base_url=self._url,
+ endpoint="history/0",
+ operation="GET",
+ ) from e
+ except requests.ReadTimeout as e:
+ print(f"...PShell service {hostname} is down")
+ raise TellCommunicationError(
+ f"TELL connection timedout ({hostname})",
+ base_url=self._url,
+ endpoint="history/0",
+ operation="GET",
+ ) from e
+
+ self._pshell = PShellClient(self._url)
+
+ @staticmethod
+ def _resolve_url(bl: MXBeamline) -> str:
+ beamline = bl.value.lower()
+ if bl == MXBeamline.X06DA:
+ return f"http://{beamline}-tell.psi.ch:22222"
+ if bl == MXBeamline.X10SA:
+ return "http://PC17488:22222"
+ if bl == MXBeamline.X06SA:
+ raise NotImplementedError(f"TellClient not implemented for {beamline}")
+ if bl == MXBeamline.SIMULATED:
+ raise NotImplementedError("Use SimTellBackend for MXBeamline.SIMULATED")
+ raise ValueError(f"Unknown beamline {beamline}")
+
+ @property
+ def url(self) -> str | None:
+ return self._url
+
+ def get_state(self) -> str:
+ return self._pshell.get_state()
+
+ def get_result(self, command_id: int = -1):
+ return self._pshell.get_result(command_id)
+
+ def wait_state(self, state: str, timeout: float) -> None:
+ self._pshell.wait_state(state, timeout=timeout)
+
+ def wait_state_not(self, state: str, timeout: float) -> None:
+ self._pshell.wait_state_not(state, timeout=timeout)
+
+ def wait_events(self, events: dict[str, Any], timeout: float):
+ return self._pshell.wait_events(events, timeout=timeout)
+
+ def eval(self, expr: str):
+ return self._pshell.eval(expr)
+
+ def start_eval(self, expr: str) -> int:
+ return self._pshell.start_eval(expr)
+
+ def run(self, path: str, pars: list[str] | None = None, background: bool = False) -> None:
+ self._pshell.run(path, pars=pars, background=background)
+
+ def abort(self) -> None:
+ self._pshell.abort()
+
+class SimTellBackend:
+ def __init__(self):
+ self._url: str | None = None
+ self._state = "Ready"
+ self._last_cmd_id = 1000
+ self._mounted_sample = ""
+ self._settings: dict[str, str] = {"mounted_sample_position": ""}
+ self._results: dict[int, dict[str, Any]] = {}
+ self._robot_status: dict[str, Any] = {
+ "powered": True,
+ "pos": POSITION_PARK,
+ }
+ self._current_mA = 30.0
+ self._pin_offset = 0.0
+ self._detected_pucks: list[dict[str, Any]] = []
+ self._system_check_msg = "OK"
+ self._smart_magnet_state = "Ready"
+ self._in_mount_position = False
+
+ @property
+ def url(self) -> str | None:
+ return self._url
+
+ def _next_cmd_id(self) -> int:
+ self._last_cmd_id += 1
+ return self._last_cmd_id
+
+ def _set_ready_soon(self) -> None:
+ time.sleep(0.01)
+ self._state = "Ready"
+
+ def get_state(self) -> str:
+ return self._state
+
+ def get_result(self, command_id: int = -1):
+ if command_id == -1:
+ command_id = self._last_cmd_id
+ return self._results.get(command_id, {"status": "completed"})
+
+ def wait_state(self, state: str, timeout: float) -> None:
+ if self._state != state:
+ time.sleep(min(timeout, 0.05))
+ self._state = state
+
+ def wait_state_not(self, state: str, timeout: float) -> None:
+ if self._state == state:
+ time.sleep(min(timeout, 0.05))
+ self._state = "Ready"
+
+ def wait_events(self, events: dict[str, Any], timeout: float):
+ self.wait_state_not("Busy", timeout)
+ if "Motion Sync" in events:
+ return "Motion Sync", "Robot Clear after mount"
+ if "Motion Task" in events:
+ return "Motion Task", "idle"
+ return None, self._state
+
+ def eval(self, expr: str):
+ expr = expr.strip()
+
+ if expr == "in_mount_position&":
+ return "true" if self._in_mount_position else "false"
+
+ if expr.startswith("in_mount_position = "):
+ self._in_mount_position = "True" in expr or "true" in expr
+ return None
+
+ if expr.startswith("set_setting("):
+ match = re.match(r"set_setting\('([^']+)', '([^']*)'\)&", expr)
+ if match:
+ key, value = match.groups()
+ self._settings[key] = value
+ return None
+
+ if expr.startswith("get_setting("):
+ match = re.match(r"get_setting\('([^']+)'\)&", expr)
+ if match:
+ key = match.group(1)
+ return self._settings.get(key, "")
+ return ""
+
+ if expr == "system_check_msg()&":
+ return self._system_check_msg
+
+ if expr == "robot.state&":
+ return self._state
+
+ if expr == "robot.take()&":
+ return str(self._robot_status)
+
+ if expr == "get_pucks_info()&":
+ return json.dumps(self._detected_pucks)
+
+ if expr == "get_pin_offset()&":
+ return str(self._pin_offset)
+
+ if expr == "smart_magnet.get_current_rb()&":
+ return str(self._current_mA)
+
+ if expr.startswith("smart_magnet.set_current("):
+ match = re.match(r"smart_magnet\.set_current\(([-+]?\d+(?:\.\d+)?)\)&", expr)
+ if match:
+ self._current_mA = float(match.group(1))
+ return None
+
+ if expr == "enable_motion()&":
+ self._robot_status["powered"] = True
+ return None
+
+ if expr == "smart_magnet.state&":
+ return self._smart_magnet_state
+
+ if expr == "smart_magnet.set_supress(True)&":
+ return None
+
+ if expr == "smart_magnet.set_supress(False)&":
+ return None
+
+ if expr == "smart_magnet.set_resting_current()&":
+ return None
+
+ if expr == "robot.stop_task()&":
+ self._state = "Ready"
+ return None
+
+ return None
+
+ def start_eval(self, expr: str) -> int:
+ cmd_id = self._next_cmd_id()
+ self._state = "Busy"
+
+ if expr.startswith("mount("):
+ parts = re.findall(r"'([^']*)'|([^,()]+)", expr)
+ values = [a if a else b.strip() for a, b in parts]
+ if len(values) >= 4:
+ segment = values[0]
+ puck = values[1]
+ sample = values[2]
+ mounted = f"{segment}{puck}{sample}"
+ self._mounted_sample = mounted
+ self._settings["mounted_sample_position"] = mounted
+ self._robot_status["pos"] = POSITION_DEWAR
+
+ elif expr.startswith("unmount("):
+ self._mounted_sample = ""
+ self._settings["mounted_sample_position"] = ""
+ self._robot_status["pos"] = POSITION_PARK
+
+ elif expr.startswith("move_park("):
+ self._robot_status["pos"] = POSITION_PARK
+
+ elif expr.startswith("move_cold("):
+ self._robot_status["pos"] = POSITION_COLD
+
+ elif expr.startswith("dry("):
+ self._robot_status["pos"] = POSITION_HEATER
+
+ self._results[cmd_id] = {"status": "completed", "command": expr}
+ self._set_ready_soon()
+ return cmd_id
+
+ def run(self, path: str, pars: list[str] | None = None, background: bool = False) -> None:
+ _ = background
+
+ if path == "data/set_samples_info" and pars:
+ try:
+ data = json.loads(pars[0])
+ self._detected_pucks = []
+ for item in data:
+ puck_address = item.get("puckAddress", "")
+ if puck_address:
+ self._detected_pucks.append(
+ {
+ "puckState": "Present",
+ "puckAddress": puck_address,
+ "puckBarcode": item.get("puckBarcode", ""),
+ }
+ )
+ except Exception:
+ logger.warning("Failed to load simulated samples info")
+
+ def abort(self) -> None:
+ self._state = "Ready"
+
+
+class LazyTellBackend:
+ def __init__(self, factory: Callable[[], TellBackend], *, retry_interval_s: float = 2.0):
+ self._factory = factory
+ self._backend: TellBackend | None = None
+ self._retry_interval_s = float(retry_interval_s)
+ self._last_attempt_ts = 0.0
+ self._last_error: Exception | None = None
+
+ def _get_backend(self) -> TellBackend:
+ if self._backend is not None:
+ return self._backend
+
+ now = time.monotonic()
+ if now - self._last_attempt_ts < self._retry_interval_s and self._last_error is not None:
+ raise self._last_error
+
+ self._last_attempt_ts = now
+ try:
+ self._backend = self._factory()
+ self._last_error = None
+ return self._backend
+ except TellCommunicationError as e:
+ self._last_error = e
+ raise
+ except Exception as e:
+ wrapped = TellCommunicationError(
+ "TELL connection failed",
+ operation="CONNECT",
+ )
+ self._last_error = wrapped
+ raise wrapped from e
+
+ @property
+ def url(self) -> str | None:
+ return self._get_backend().url
+
+ def get_state(self) -> str:
+ return self._get_backend().get_state()
+
+ def get_result(self, command_id: int = -1):
+ return self._get_backend().get_result(command_id)
+
+ def wait_state(self, state: str, timeout: float) -> None:
+ self._get_backend().wait_state(state, timeout)
+
+ def wait_state_not(self, state: str, timeout: float) -> None:
+ self._get_backend().wait_state_not(state, timeout)
+
+ def wait_events(self, events: dict[str, Any], timeout: float):
+ return self._get_backend().wait_events(events, timeout)
+
+ def eval(self, expr: str):
+ return self._get_backend().eval(expr)
+
+ def start_eval(self, expr: str) -> int:
+ return self._get_backend().start_eval(expr)
+
+ def run(self, path: str, pars: list[str] | None = None, background: bool = False) -> None:
+ self._get_backend().run(path, pars=pars, background=background)
+
+ def abort(self) -> None:
+ self._get_backend().abort()
\ No newline at end of file
diff --git a/src/aare/devices/tell_client.py b/src/aare/devices/tell_client.py
index 9ff31f99..f7a4d8c9 100755
--- a/src/aare/devices/tell_client.py
+++ b/src/aare/devices/tell_client.py
@@ -1,14 +1,9 @@
+import ast
import json
-import random
import re
-import time
from typing import List
-from urllib.parse import urlparse
-import requests
-
-from aare.common.exception_handler import TellCommunicationError
-from aare.common.logger_config import setup_logger
+from aare.common.beamline import MXBeamline
from aare.common.models import (
PuckLoadedInfo,
DewarAddress,
@@ -16,93 +11,28 @@ from aare.common.models import (
)
from aareDB import PuckWithTellPosition
-from aare.common.beamline import MXBeamline # noqa: F401
-from pshell import PShellClient
+from aare.devices.tell_backend import (
+ TellBackend,
+ SimTellBackend,
+ LazyTellBackend,
+ PShellTellBackend,
+ ManualMountException,
+ POSITION_COLD,
+ SmartMagnetFaultException,
+ TellConnectionException,
+ TellMountFailedException,
+ TellCommandWhileBusyException,
+)
+
+from aare.common.logger_config import setup_logger
logger = setup_logger("aareDAQ")
-class ManualMountException(Exception):
- """Custom exception for manual mounting"""
- pass
-
-
-class SmartMagnetFaultException(Exception):
- """Custom exception for smart magnet fault"""
- pass
-
-
-class TellMountFailedException(Exception):
- """Custom exception for mount failure"""
- pass
-
-
-class TellCommandWhileBusyException(Exception):
- """Custom exception for trying to move Tell when it is busy"""
- pass
-
-
-class TellConnectionException(Exception):
- """Custom exception for connection problems"""
- pass
-
-VALID_DEWAR_POSITIONS = [f"{p}{n}" for n in "12345" for p in "ABCDEFX"]
-
-def is_valid_dewar_position(position):
- """check if argument is a valid dewar position"""
- return position in VALID_DEWAR_POSITIONS
-
-POSITION_PARK = "pPark"
-POSITION_COLD = "pCold"
-POSITION_AUX = "pAux"
-POSITION_DEWAR = "pDewar"
-POSITION_HOME = "pHome"
-POSITION_HEATER = "pHeatB"
-
-#Nov 26 13:36:00 mx-x06da-queue-01.psi.ch AareDAQ[2944444]: 2025-11-26 13:36:00,388 - aareDAQ - ERROR - Error getting status: ('Connection aborted.', ConnectionResetError(104, 'Connection reset by peer'))
-
class TellClient:
- """High-level Tell robot API using PShellClient"""
- def __init__(self, bl: MXBeamline):
- self.__url = None
- beamline = bl.value.lower()
+ """High-level Tell robot API using a pluggable backend"""
+ def __init__(self, bl: MXBeamline, backend: TellBackend | None = None):
self.__beamline = bl
- if bl == MXBeamline.X06DA:
- self.__url = f"http://{beamline}-tell.psi.ch:22222"
-
- elif bl == MXBeamline.X10SA:
- self.__url = f"http://PC17488:22222"
-
- elif bl == MXBeamline.X06SA:
- self.__url = f""
- raise NotImplemented(f"TellClient not implemente for {beamline}")
- elif bl == MXBeamline.SIMULATED:
- raise NotImplemented(f"Use SimClient, generate tell client using"
- f"make_tell_client(beamline)")
- else:
- raise ValueError(f"Unknown beamline {beamline}")
-
- print(f"Connecting TELL p-shell service at {self.__url} ...", end="")
- hostname = urlparse(self.__url).hostname
- try:
- requests.get(f"{self.__url}/history/0", timeout=1.0)
- except requests.exceptions.RequestException as e:
- print(f"...connection to {hostname} failed")
- raise TellCommunicationError(
- f"TELL connection failed ({hostname})",
- base_url=self.__url,
- endpoint="history/0",
- operation="GET",
- ) from e
- except requests.ReadTimeout as e:
- print(f"...PShell service {hostname} is down")
- raise TellCommunicationError(
- f"TELL connection timedout ({hostname})",
- base_url=self.__url,
- endpoint="history/0",
- operation="GET",
- ) from e
-
- self.pshell = PShellClient(self.__url)
+ self.backend = backend or PShellTellBackend(bl)
self._aborted = False
self.state = self.get_state()
@@ -112,26 +42,26 @@ class TellClient:
@property
def url(self):
"""returns the configured base url for the Tell robot"""
- return self.__url
+ return self.backend.url
def get_state(self):
"""returns the current state of the robot"""
- self.state = self.pshell.get_state()
+ self.state = self.backend.get_state()
return self.state
def get_result(self, command_id=-1):
"""returns the result of the last command issued to the robot"""
- return self.pshell.get_result(command_id)
+ return self.backend.get_result(command_id)
def wait_ready(self, timeout: float = 360.0):
"""waits until the robot is ready to accept commands returns None if simulation
and raises an exception if the robot is not ready"""
- self.pshell.wait_state("Ready", timeout=timeout)
+ self.backend.wait_state("Ready", timeout=timeout)
def wait_not_busy(self, timeout: float = 360.0):
"""waits until the robot is not busy and returns None if simulation
and raises an exception if the robot is busy"""
- self.pshell.wait_state_not("Busy", timeout=timeout)
+ self.backend.wait_state_not("Busy", timeout=timeout)
state = self.get_state()
if state != "Ready":
if state == "Initializing":
@@ -143,11 +73,11 @@ class TellClient:
def set_in_mount_position(self, value):
"""tells the robot that the beamlien is safe and to set the in mount position flag allowing mounting
:param value """
- self.pshell.eval("in_mount_position = " + str(value) + "&")
+ self.backend.eval("in_mount_position = " + str(value) + "&")
def is_in_mount_position(self) -> bool:
"""checks to see if the robot is in the mount position and returns a boolean"""
- return self.pshell.eval("in_mount_position&").lower() == "true"
+ return self.backend.eval("in_mount_position&").lower() == "true"
def set_samples_info(self, info: List[PuckWithTellPosition]):
"""sets the samples in the robot dewar based on the given list of PuckWithTellPosition objects
@@ -160,7 +90,7 @@ class TellClient:
"userName": x.pgroup,
"dewarName": x.dewar_name or "",
"puckName": x.puck_name,
- "puckType": "Unipuck", # could use x.puck_type
+ "puckType": "Unipuck",
"puckAddress": x.tell_position or "",
"puckBarcode": x.puck_name,
"sampleBarcode": "",
@@ -171,8 +101,7 @@ class TellClient:
}
)
- self.pshell.run("data/set_samples_info", pars=[json.dumps(j)], background=True)
- # self.pshell.eval("set_samples_info(" + json.dumps(info) + ")&")
+ self.backend.run("data/set_samples_info", pars=[json.dumps(j)], background=True)
def start_cmd(self, cmd, *argv):
"""starts a command on the robot and returns the command id"""
@@ -180,7 +109,7 @@ class TellClient:
for a in argv:
cmd = cmd + (("'" + a + "'") if type(a) is str else str(a)) + ", "
cmd = cmd + ")"
- ret = self.pshell.start_eval(cmd)
+ ret = self.backend.start_eval(cmd)
self.get_state()
return ret
@@ -191,11 +120,10 @@ class TellClient:
result = self.get_result(self._last_cmd_id)
logger.debug(f"{msg} {result}")
status = result["status"]
- if "completed" != status: #FIXME this is very limiting and depends on tell reporting statuses
+ if "completed" != status:
if "removed" != status:
raise TellMountFailedException(f"{msg} {result}")
- else:
- return f"{msg} {result}"
+ return f"{msg} {result}"
def estimate_mounting_time(self, segment) -> int:
"""Adds additional time if cooling/drying is expected based on requested segment,
@@ -206,7 +134,7 @@ class TellClient:
gripper_in_cold = self.is_in_cold()
if current_mounted is None:
- unmount_needs_drying = 0 # might not have anything
+ unmount_needs_drying = 0
unmount_needs_cooling = 0
else:
segment_in_cold = current_mounted.puck.segment in "ABCDEF"
@@ -219,27 +147,18 @@ class TellClient:
needs_cooling = mount_needs_cooling + unmount_needs_cooling
needs_drying = mount_needs_drying + unmount_needs_drying
return needs_cooling * 30 + needs_drying * 120
- except:
+ except Exception:
return 0
def mount(
self,
address: SampleDewarAddress,
- force: bool = False, # kept for future
- read_dm: bool = False, # read data matrix
- auto_unmount: bool = False, # single command, if False it will raise exception
- wait: bool = False, # blocking operation
+ force: bool = False,
+ read_dm: bool = False,
+ auto_unmount: bool = False,
+ wait: bool = False,
timeout: float = 600.0,
):
- """send api request to mount sample from dewer after validating dewer address returns None or repsonse.
- If the robot is busy, mount will raise an exception.
- :param address: SampleDewarAddress
- :param force: bool
- :param read_dm: bool
- :param auto_unmount: bool
- :param wait: bool
- :param timeout: float
- """
SampleDewarAddress.model_validate(address)
segment = address.puck.segment
@@ -258,9 +177,16 @@ class TellClient:
wait_timeout = timeout + self.estimate_mounting_time(segment)
logger.info("waiting for mount to complete")
if wait and segment in "ABCDEF":
- event, value = self.pshell.wait_events({"state": None, "Motion Task": "dry",
- "Gripper detection" : None,
- "Motion Sync": "Robot Clear after mount"}, timeout=wait_timeout)
+ event, value = self.backend.wait_events(
+ {
+ "state": None,
+ "Motion Task": "dry",
+ "Gripper detection": None,
+ "Motion Sync": "Robot Clear after mount",
+ },
+ timeout=wait_timeout,
+ )
+ logger.info(f"event: {event} occurred with value: {value}")
if event is None or event == "state":
logger.info(f"event: {event} occurred with value: {value}, checking command completed okay")
self.check_command_ok(
@@ -299,11 +225,6 @@ class TellClient:
return None
def unmount(self, force=False, wait=False, timeout=360.0):
- """send api request to unmount sample from dewer returns None or repsonse.
- :param force: bool Force has a meaning, will unmount even if smart magnet is not detecting sample
- :param wait: bool If true will wait until unmount is completed
- :timeout: float"""
-
if self.is_busy():
raise TellCommandWhileBusyException("mount received while robot is busy")
@@ -315,82 +236,60 @@ class TellClient:
return self._last_cmd_id
def dry(self, heat_time=None, speed=None, wait_cold=None, wait=False):
- """send api request to dry tell gripper.
- :param: heat_time float if None Tell will use default for drying time
- :param: speed float if None Tell will use default for drying speed
- :param: wait_cold bool if -1 to go to park after dry. if None Tell will use default time to wait_cold.
- :param wait: bool If true will wait until drying is completed
- """
- self.pshell.wait_state("Ready", timeout=30.0)
+ self.backend.wait_state("Ready", timeout=30.0)
self._last_cmd_id = self.start_cmd("dry", heat_time, speed, wait_cold)
if wait:
- self.check_command_ok(timeout=360.0, msg=f"Dry failed")
+ self.check_command_ok(timeout=360.0, msg="Dry failed")
def move_park(self, wait=False):
- """send api request to move robot to park position"""
self._last_cmd_id = self.start_cmd("move_park")
-
if wait:
- self.check_command_ok(timeout=360.0, msg=f"Move to park failed")
+ self.check_command_ok(timeout=360.0, msg="Move to park failed")
def move_cold(self, reset_timestamp=False, wait=False):
- """send api request to move robot to cold position"""
self._last_cmd_id = self.start_cmd("move_cold", reset_timestamp)
-
if wait:
- self.check_command_ok(timeout=360.0, msg=f"Move to cold failed")
+ self.check_command_ok(timeout=360.0, msg="Move to cold failed")
def abort_cmd(self):
- """sends an abort pshell requesst and a robot stop task command"""
- self.pshell.abort()
- self.pshell.eval("robot.stop_task()&")
+ self.backend.abort()
+ self.backend.eval("robot.stop_task()&")
def set_setting(self, key: str, value: str):
- """wrapper for pshell eval set_setting command
- :param key str, name of a setting in tell
- :param value str, the new value of the setting as a string"""
- self.pshell.eval(f"set_setting('{key}', '{value}')&")
+ self.backend.eval(f"set_setting('{key}', '{value}')&")
def get_setting(self, key: str) -> str:
- """wrapper for pshell eval get_setting command, returns the current value for key as a string
- :param key str, name of a setting in tell"""
- return self.pshell.eval(f"get_setting('{key}')&")
+ return self.backend.eval(f"get_setting('{key}')&")
def get_mounted_sample(self) -> SampleDewarAddress | None:
- """get the current mounted sample and return a SampleDewarAddress object or None if no sample is mounted"""
- ret = self.get_setting('mounted_sample_position').strip()
- if not ret or len(ret) == 0:
+ ret = self.get_setting("mounted_sample_position").strip()
+ if not ret:
return None
match = re.match(r"([A-Z])(\d)(\d{1,2})", ret)
-
if match:
segment, puck, sample = match.groups()
dewar_location = DewarAddress(segment=segment, pos=int(puck))
return SampleDewarAddress(puck=dewar_location, pin=int(sample))
- else:
- logger.warning(f"Failed to decode mounted sample position: {ret}")
- return None
+
+ logger.warning(f"Failed to decode mounted sample position: {ret}")
+ return None
def get_system_check(self):
- """returns the current system check status"""
- return self.pshell.eval("system_check_msg()&")
+ return self.backend.eval("system_check_msg()&")
def get_robot_state(self):
- """returns the current robot state"""
- return self.pshell.eval("robot.state&")
+ return self.backend.eval("robot.state&")
def get_robot_status(self):
- """returns the current robot status"""
- status = self.pshell.eval("robot.take()&")
- return eval(status) # FIXME ALL functions must return a valid JSON object
+ status = self.backend.eval("robot.take()&")
+ #return eval(status)
+ return ast.literal_eval(status)
def get_detected_pucks(self) -> List[PuckLoadedInfo]:
- """returns a list of detected pucks as PuckLoadedInfo objects"""
- j = json.loads(self.pshell.eval("get_pucks_info()&"))
+ j = json.loads(self.backend.eval("get_pucks_info()&"))
output = []
-
for i in j:
if i["puckState"] == "Present":
puck_address = i["puckAddress"]
@@ -399,90 +298,77 @@ class TellClient:
PuckLoadedInfo(
puck_name=i["puckBarcode"],
location=DewarAddress(
- segment=puck_address[0], pos=int(puck_address[1])
+ segment=puck_address[0],
+ pos=int(puck_address[1]),
),
),
)
return output
def get_pin_offset(self):
- """get the pin offset for the smart magnet, returns offset as a float"""
try:
- offset = float(self.pshell.eval("get_pin_offset()&"))
+ offset = float(self.backend.eval("get_pin_offset()&"))
except Exception:
offset = 0.0
return offset
def get_current(self):
- """get the current drawn by the smart magnet, returns current as a float in mA"""
- current = self.pshell.eval("smart_magnet.get_current_rb()&")
+ current = self.backend.eval("smart_magnet.get_current_rb()&")
return float(current)
def set_current(self, current: float) -> float:
- """set the current drawn by the smart magnet, returns current as a float in mA"""
- self.pshell.eval("smart_magnet.set_current({:.1f})&".format(current))
- current = self.pshell.eval("smart_magnet.get_current_rb()&")
+ self.backend.eval("smart_magnet.set_current({:.1f})&".format(current))
+ current = self.backend.eval("smart_magnet.get_current_rb()&")
return float(current)
def is_powered(self):
- """returns True if the robot is powered on"""
return self.get_robot_status()["powered"]
def check_enable_motion(self):
- """check if the robot is powered on and enable motion if not"""
if not self.is_powered():
- self.pshell.eval("enable_motion()&")
+ self.backend.eval("enable_motion()&")
def is_in_cold(self):
- """Compare current robot position to the set cold position. Returns True if in cold position, False otherwise."""
return self.is_position(POSITION_COLD)
def is_position(self, position: str) -> bool:
- """Compare current robot position to a given position. Returns True if in position, False otherwise."""
return position == self.get_robot_status()["pos"]
def is_ready(self):
- """returns True if the robot is ready to receive commands"""
return "ready" == self.get_state().lower()
def is_busy(self):
- """returns True if the robot is busy"""
return "busy" == self.get_state().lower()
def check_smart_magnet_mounted(self, timeout: float = 10.0, idle_time: float = 1.0, interval: float = 0.1):
- """Reads smart_magent state and tries to infer if a sample is present
- Handles: PAUSED, Fault, Busy and Ready states.
- Raises a ManualMountException is the amgnet indicates a sample is present but get_mounted_sample is None.
- Raises a SmartMagnetFaultException if the magnet detects no sample but the robot thinks a sample is mounted"""
- #TODO tidy up
- initial_state = self.pshell.eval("smart_magnet.state&")
+ #Not sure why unused, potentially can remove them
+ _ = (timeout, idle_time, interval)
+
+ initial_state = self.backend.eval("smart_magnet.state&")
logger.debug(f"checking smart magnet_initial state: {initial_state}")
if initial_state == "Paused":
- self.pshell.eval("smart_magnet.set_supress(False)&")
- self.pshell.eval("smart_magnet.set_resting_current()&")
-
+ self.backend.eval("smart_magnet.set_supress(False)&")
+ self.backend.eval("smart_magnet.set_resting_current()&")
elif initial_state == "Fault":
logger.error(f"tell smart magnet is in unknown state {initial_state}")
raise SmartMagnetFaultException
- state = self.pshell.eval("smart_magnet.state&")
+ state = self.backend.eval("smart_magnet.state&")
try:
if state == "Busy":
- logger.debug('state busy')
- self.pshell.eval("smart_magnet.set_supress(True)&")
- self.pshell.eval("smart_magnet.state&")
- sample_present = True
+ logger.debug("state busy")
+ self.backend.eval("smart_magnet.set_supress(True)&")
+ self.backend.eval("smart_magnet.state&")
if self.get_mounted_sample() is None:
logger.warning("Check mount: A manually mounted sample is detected.")
logger.warning("Remove before mounting with the robot.")
raise ManualMountException
return True
elif state == "Ready":
- logger.debug('No sample detected, ready to mount')
- sample_present = False
+ logger.debug("No sample detected, ready to mount")
if self.get_mounted_sample():
logger.error("Check mount: No sample detected, but robot thinks is mounted")
raise SmartMagnetFaultException
@@ -491,174 +377,20 @@ class TellClient:
logger.debug("Smart magnet detection is paused")
return None
else:
- self.pshell.eval("smart_magnet.set_supress(True)&")
+ self.backend.eval("smart_magnet.set_supress(True)&")
logger.error(f"Tell smart magnet is in unknown state {state}")
raise SmartMagnetFaultException
-
except Exception as e:
logger.error(f"check_smart_magnet_mounted failed: {e}")
raise e
-class SimTellClient:
- """
- Simulation-only Tell client.
- Keeps behavior deterministic-ish and stateful without needing PShellClient.
- Implement more methods as your callers need them.
- """
- def __init__(self):
- self._state = "Ready"
- self._last_cmd_id = 1000
- self._mounted_sample: str = ""
- self._simulated_samples_info = {}
- self._simulated_detected_pucks = []
- self._simulated_current = 30.0
- self._simulated_suppress = True
- self._simulated_offset = 0.0
-
- @property
- def url(self):
- return None
-
- def _next_cmd_id(self) -> int:
- self._last_cmd_id += 1
- return self._last_cmd_id
-
- def get_state(self) -> str:
- return self._state
-
- def is_ready(self) -> bool:
- return self._state.lower() == "ready"
-
- def is_busy(self) -> bool:
- return self._state.lower() == "busy"
-
- def wait_ready(self, timeout: float = 360.0):
- # Keep it simple: flip to Ready quickly.
- time.sleep(0.05)
- self._state = "Ready"
-
- def mount(
- self,
- address: SampleDewarAddress,
- force: bool = False,
- read_dm: bool = False,
- auto_unmount: bool = False,
- wait: bool = False,
- timeout: float = 600.0,
- ):
- SampleDewarAddress.model_validate(address)
- if self.is_busy():
- raise TellCommandWhileBusyException("mount received while robot is busy")
-
- cmd_id = self._next_cmd_id()
- self._state = "Busy"
-
- segment = address.puck.segment
- puck = address.puck.pos
- sample = address.pin
- self._mounted_sample = f"{segment}{puck}{sample}"
-
- if wait:
- self.wait_ready(timeout=timeout)
- else:
- # quickly become ready anyway, but asynchronously-ish
- time.sleep(0.01)
- self._state = "Ready"
-
- return cmd_id
-
- def unmount(self, force: bool = False, wait: bool = False, timeout: float = 360.0):
- if self.is_busy():
- raise TellCommandWhileBusyException("unmount received while robot is busy")
-
- cmd_id = self._next_cmd_id()
- self._state = "Busy"
- self._mounted_sample = ""
- if wait:
- self.wait_ready(timeout=timeout)
- else:
- time.sleep(0.01)
- self._state = "Ready"
- return cmd_id
-
- def get_mounted_sample(self) -> SampleDewarAddress | None:
- ret = self._mounted_sample
- if not ret:
- return None
- match = re.match(r"([A-Z])(\d)(\d{1,2})", ret)
- if not match:
- return None
- segment, puck, sample = match.groups()
- return SampleDewarAddress(puck=DewarAddress(segment=segment, pos=int(puck)), pin=int(sample))
-
-class TellClientProxy:
- """
- Lazy-connecting Tell client proxy that retries periodically.
- - Server can start even if TELL is down.
- - First use triggers connect; failures raise TellCommunicationError.
- """
- def __init__(self, bl: MXBeamline, *, retry_interval_s: float = 2.0):
- self._bl = bl
- self._client: TellClient | None = None
- self._retry_interval_s = float(retry_interval_s)
- self._last_attempt_ts = 0.0
- self._last_error: Exception | None = None
-
- def _get_client(self) -> TellClient:
- if self._client is not None:
- return self._client
-
- now = time.monotonic()
- if now - self._last_attempt_ts < self._retry_interval_s and self._last_error is not None:
- raise self._last_error
-
- self._last_attempt_ts = now
- try:
- self._client = TellClient(self._bl)
- self._last_error = None
- return self._client
- except TellCommunicationError as e:
- self._last_error = e
- raise
- except Exception as e:
- wrapped = TellCommunicationError(
- "TELL connection failed",
- operation="CONNECT",
- )
- self._last_error = wrapped
- raise wrapped from e
-
- @property
- def url(self):
- return self._get_client().url
-
- # Delegate methods used by DAQ; add more as needed
- def get_mounted_sample(self) -> SampleDewarAddress | None:
- return self._get_client().get_mounted_sample()
-
- def get_state(self):
- return self._get_client().get_state()
-
- def wait_not_busy(self, timeout: float = 360.0):
- return self._get_client().wait_not_busy(timeout=timeout)
-
- def check_enable_motion(self):
- return self._get_client().check_enable_motion()
-
- def set_in_mount_position(self, value):
- return self._get_client().set_in_mount_position(value)
-
- def mount(self, *args, **kwargs):
- return self._get_client().mount(*args, **kwargs)
-
- def unmount(self, *args, **kwargs):
- return self._get_client().unmount(*args, **kwargs)
-
- def abort_cmd(self):
- return self._get_client().abort_cmd()
-
-def make_tell_client(bl: MXBeamline) -> TellClient | SimTellClient | TellClientProxy:
+def make_tell_client(bl: MXBeamline) -> TellClient:
if bl == MXBeamline.SIMULATED:
- return SimTellClient()
- return TellClientProxy(bl, retry_interval_s=2.0)
\ No newline at end of file
+ backend = SimTellBackend()
+ else:
+ backend = LazyTellBackend(
+ factory=lambda: PShellTellBackend(bl),
+ retry_interval_s=2.0,
+ )
+ return TellClient(bl, backend=backend)
\ No newline at end of file
diff --git a/src/aare/gui/gui.py b/src/aare/gui/gui.py
index 284ec734..712380ef 100644
--- a/src/aare/gui/gui.py
+++ b/src/aare/gui/gui.py
@@ -36,8 +36,8 @@ if __name__ == "__main__":
default_gonio_cam_addr = "axis-accc8ed2972e.psi.ch"
default_gonio_camera_id = 3
case MXBeamline.X10SA:
- default_url = "http://127.0.0.1:5210"
- default_zmq_addr = "tcp://x10sa-spark-01:9091" #"tcp://x10sa-pserv-01:9089" #
+ default_url = "http://mx-x10sa-queue-01.psi.ch:5210" #"http://127.0.0.1:5210"#
+ default_zmq_addr = "tcp://x10sa-spark-01:9091" # "tcp://x10sa-spark-01:9091" #
default_pred_zmq_addr = "tcp://x10sa-spark-01:9091" #"tcp://sls-gpu-003:9089"#""
default_beamline_cam_addr = "axis-accc8eb02488.psi.ch"
default_gonio_cam_addr = "axis-accc8ea5e463.psi.ch"
diff --git a/src/aare/gui/main_window.py b/src/aare/gui/main_window.py
index 06d0c67f..9ec75126 100644
--- a/src/aare/gui/main_window.py
+++ b/src/aare/gui/main_window.py
@@ -1,4 +1,5 @@
import time
+import requests
import jwt
from PySide6.QtCore import Qt, Slot, Signal, QTimer, QSettings
@@ -12,6 +13,7 @@ from PySide6.QtWidgets import (
QDockWidget,
QTabWidget, QFrame, QSizePolicy, QLabel)
+from aare.common.auth_models import BatonStatus, BatonRequestStatus
from aare.common.coordinate import Coordinate, SmargonCoordinate
from aare.common.diffraction_geometry import DiffractionGeometry
from aare.common.logger_config import setup_logger
@@ -47,13 +49,18 @@ from aare.gui.threads.camera_thread import SampleCameraThread
from aare.gui.threads.prediction_subscriber import PredictionSubscriber
from aare.gui.threads.daq_worker import DAQWorker
from aare.gui.threads.jfjoch_viewer import JFJochDBusClient
+from aare.gui.tutorials.tutorial_registration import register_tutorials
from aare.gui.widgets.alert_banner import AlertBanner
+from aare.gui.widgets.baton_request_dialog import BatonRequestDialog, BatonPendingDialog
from aare.gui.widgets.camera_image import SampleCameraImageLabel
from aare.gui.widgets.no_wheel_scroll_area import NoWheelScrollArea
from aare.gui.widgets.status_bar import StatusBar
from aare.gui.widgets.video_image import VideoGraphicsView
from aare.gui.panels.fluorescence_panel import FluorescencePanel
+from aare.gui.panels.automation_panel import WorkflowPanel
+from aare.gui.threads.workflow_sse_client import WorkflowSSEClient
+
logger = setup_logger("aareGUI")
class MainWindow(QMainWindow):
@@ -69,6 +76,7 @@ class MainWindow(QMainWindow):
gonio_cam_id: int | None
):
super().__init__()
+
self.__base_url = base_url
self.__token = token
self.__mounting = False
@@ -76,6 +84,11 @@ class MainWindow(QMainWindow):
self._beamline_recovery_dialog = None
self._controls_help_dialog = None
self._cleanup_done = False
+ self._default_window_state = None
+
+ self._waiting_for_baton_response: bool = False
+ self._baton_request_dialog: BatonRequestDialog | None = None
+ self._baton_pending_dialog: BatonPendingDialog | None = None
# Tutorial manager (define tutorials after widgets exist)
self._tutorial_event_bus = TutorialEventBus(self)
@@ -114,6 +127,9 @@ class MainWindow(QMainWindow):
self.alert_banner = AlertBanner(parent=root_widget)
root_layout.addWidget(self.alert_banner)
+ self.alert_banner_secondary = AlertBanner(parent=root_widget)
+ root_layout.addWidget(self.alert_banner_secondary)
+
top_widget = QWidget(parent=root_widget)
top_widget_layout = QHBoxLayout(top_widget)
top_widget.setLayout(top_widget_layout)
@@ -221,6 +237,11 @@ class MainWindow(QMainWindow):
self.ref_tools_dock.setAllowedAreas(Qt.DockWidgetArea.BottomDockWidgetArea)
self.addDockWidget(Qt.DockWidgetArea.BottomDockWidgetArea, self.ref_tools_dock)
self.tabifyDockWidget(self.ref_tools_dock, self.tell_samples_dock)
+ if self.__decoded_token.staff:
+ self.ref_tools_dock.show()
+ self.ref_tools_dock.raise_()
+ else:
+ self.ref_tools_dock.hide()
self.sample_logic = SampleMountLogic()
@@ -277,11 +298,23 @@ class MainWindow(QMainWindow):
self.addDockWidget(Qt.DockWidgetArea.RightDockWidgetArea, self.smargon_trace_dock)
self.smargon_trace_dock.hide()
+ # === Workflow Panel ===
+ self.workflow_panel = WorkflowPanel()
+ self.workflow_dock = QDockWidget("Workflow", self)
+ self.workflow_dock.setObjectName("workflow_dock")
+ self.workflow_dock.setWidget(self.workflow_panel)
+ self.workflow_dock.setAllowedAreas(
+ Qt.DockWidgetArea.RightDockWidgetArea | Qt.DockWidgetArea.LeftDockWidgetArea
+ )
+ self.addDockWidget(Qt.DockWidgetArea.RightDockWidgetArea, self.workflow_dock)
+ self.workflow_dock.hide() # Hidden by default
+
root_layout.addWidget(top_widget)
self.setCentralWidget(root_widget)
self.setWindowTitle("AareGUI")
self.create_menu_bar()
+ self._capture_default_window_state()
self._restore_window_state()
# Define tutorials now that the UI exists
@@ -300,13 +333,23 @@ class MainWindow(QMainWindow):
self.setStatusBar(self.status_bar)
self.daq = DAQWorker(base_url=self.__base_url, token=self.__token)
+
+ self.daq.baton_status_changed.connect(self.status_bar.update_baton_status)
+ self.daq.baton_status_changed.connect(self._on_baton_status_changed)
+ self.daq.baton_request_result.connect(self._on_baton_request_result)
+ self.daq.baton_response_result.connect(self._on_baton_response_result)
+ self.daq.baton_timeout_checked.connect(self._on_baton_timeout_checked)
+
+ self.status_bar.baton_request_received.connect(self._show_baton_request_dialog)
+ self.status_bar.baton_request_accepted.connect(self._accept_baton_request)
+ self.status_bar.baton_request_refused.connect(self._refuse_baton_request)
+
self.daq.spreadsheet.connect(self.tell_samples.new_sample_list)
if self.__decoded_token.staff:
self.daq.reference_tools.connect(self.ref_tools_panel.new_list)
self.beamline.samcam.changed.connect(self.daq.samcam_settings)
self.beamline.samcam.screenshot_requested.connect(self.daq.send_screenshot_db)
- self.beamline.loopctr.background.clicked.connect(self.daq.alc_background)
self.beamline.loopctr.find_tip.clicked.connect(self.daq.center_loop)
self.beamline.loopctr.bounding_box.clicked.connect(self.daq.ml_bounding_box)
self.daq.raster_generated_by_ml.connect(self.raster.update_active_grid_request)
@@ -411,9 +454,21 @@ class MainWindow(QMainWindow):
self.data_collection.simple.parameters_changed.connect(self.daq.smart_params)
self.raster.grid_scan_size_changed.connect(self.data_collection.raster.grid_scan_size_change)
+
self.status_bar.set_pgroup.connect(self.daq.set_pgroup)
self.status_bar.end_session.connect(self.daq.end_session)
self.status_bar.force_session.connect(self.daq.force_session)
+
+ self.status_bar.request_baton.connect(self.daq.request_baton)
+ self.status_bar.cancel_baton_request.connect(self.daq.cancel_baton_request)
+ self.status_bar.release_baton.connect(self.daq.release_baton)
+ self.status_bar.baton_request_accepted.connect(
+ lambda: self.daq.respond_to_baton_request(True)
+ )
+ self.status_bar.baton_request_refused.connect(
+ lambda: self.daq.respond_to_baton_request(False)
+ )
+
self.status_bar.dewar_exchange.connect(self.daq.dewar_exchange)
self.status_bar.sample_exchange.connect(self.daq.sample_exchange)
self.status_bar.sample_alignment.connect(self.daq.sample_alignment)
@@ -421,6 +476,7 @@ class MainWindow(QMainWindow):
self.status_bar.close_shutter.connect(self.daq.close_shutter)
self.status_bar.open_shutter.connect(self.daq.open_shutter)
+
self.rotation.file_ready.connect(self.viewer.load_image)
self.raster.image_selected.connect(self.viewer.load_image)
@@ -474,11 +530,59 @@ class MainWindow(QMainWindow):
self.daq.fluorimeter_spectrum_update.connect(self.fluor_panel.update_plot)
self.daq.fluorimeter_spectrum_update.connect(lambda: self.fluor_panel_dock.setVisible(True))
+ # === Alert/Status Message Routing ===
+ # Primary alert banner: Infrastructure devices (Server/Tell/Smargon/Aerotech)
+ self.daq.polled_devices_status.connect(self.alert_banner.show_message)
+
+ # Secondary alert banner: Detector errors (JFJoch)
+ self.daq.detector_error.connect(self.alert_banner_secondary.show_message)
+
+ # Status bar: General status messages (not device connection status)
self.daq.status_message.connect(self.status_bar.show_connection_message)
- self.daq.status_message.connect(self.alert_banner.show_message)
register_tutorials(self, self.tutorial_manager)
+ # Workflow SSE client
+ if self.__base_url is not None:
+ self.workflow_sse = WorkflowSSEClient(self.__base_url, self.__token, self)
+ self.workflow_sse.runtime_changed.connect(self.workflow_panel.update_runtime)
+ self.workflow_sse.control_changed.connect(self.workflow_panel.update_control)
+ self.workflow_sse.workflow_event.connect(
+ lambda e: self.workflow_panel.on_workflow_event(e.event_type, e.message)
+ )
+ self.workflow_sse.connect()
+ else:
+ self.workflow_sse = None
+
+ # Connect workflow panel signals to DAQ worker
+ self.workflow_panel.request_queue_refresh.connect(self.daq.workflow_load_queue)
+ self.workflow_panel.request_add_sample.connect(self.daq.workflow_add_sample)
+ self.workflow_panel.request_add_samples.connect(self.daq.workflow_add_samples)
+ self.workflow_panel.request_delete_item.connect(self.daq.workflow_delete_item)
+ self.workflow_panel.request_move_item.connect(self.daq.workflow_move_item)
+ self.workflow_panel.request_clear_queue.connect(self.daq.workflow_clear_queue)
+ self.workflow_panel.request_start_item.connect(self.daq.workflow_start_item)
+ self.workflow_panel.request_next_step.connect(self.daq.workflow_next_step)
+ self.workflow_panel.request_pause.connect(self.daq.workflow_pause)
+ self.workflow_panel.request_resume.connect(self.daq.workflow_resume)
+ self.workflow_panel.request_abort.connect(self.daq.workflow_abort)
+ self.workflow_panel.request_skip.connect(self.daq.workflow_skip)
+ self.workflow_panel.request_start_automation.connect(self.daq.workflow_start_automation)
+ self.workflow_panel.request_stop_automation.connect(self.daq.workflow_stop_automation)
+ self.workflow_panel.request_skip_sample.connect(self.daq.workflow_skip_sample)
+
+ # DAQ worker -> workflow panel
+ self.daq.workflow_queue_loaded.connect(self.workflow_panel.update_queue)
+ self.daq.workflow_item_updated.connect(self.workflow_panel.update_current_item)
+
+ # Forward sample list to workflow panel for "Add All" feature
+ self.daq.spreadsheet.connect(
+ lambda slist: self.workflow_panel.update_available_samples(slist.s)
+ )
+
+ # Initial load
+ QTimer.singleShot(1000, self.daq.workflow_load_queue)
+
@Slot(QPixmap)
def _on_samcam_prediction_pixmap(self, pix: QPixmap) -> None:
self._last_pred_image_ts = time.monotonic()
@@ -522,35 +626,41 @@ class MainWindow(QMainWindow):
file_menu.addAction(quit_action)
view_menu = menu_bar.addMenu("View")
+
show_samples_action = QAction("Show Sample List", self)
show_samples_action.setCheckable(True)
show_samples_action.setChecked(True)
- show_samples_action.triggered.connect(lambda: self.tell_samples_dock.setVisible(show_samples_action.isChecked()))
-
- # When dock visibility changes, update the action's checked state
+ show_samples_action.triggered.connect(lambda checked: self.tell_samples_dock.setVisible(checked))
self.tell_samples_dock.visibilityChanged.connect(show_samples_action.setChecked)
view_menu.addAction(show_samples_action)
+ if self.__decoded_token.staff:
+ show_reference_tools_action = QAction("Show Reference Tools", self)
+ show_reference_tools_action.setCheckable(True)
+ show_reference_tools_action.setChecked(True)
+ show_reference_tools_action.triggered.connect(
+ lambda checked: self.ref_tools_dock.setVisible(checked)
+ )
+ self.ref_tools_dock.visibilityChanged.connect(show_reference_tools_action.setChecked)
+ view_menu.addAction(show_reference_tools_action)
+
show_job_list_action = QAction("Show job List", self)
show_job_list_action.setCheckable(True)
show_job_list_action.setChecked(True)
- show_job_list_action.triggered.connect(self.job_list_dock.setVisible)
- # When dock visibility changes, update the action's checked state
- show_job_list_action.triggered.connect(lambda: self.job_list_dock.setVisible(show_job_list_action.isChecked()))
+ show_job_list_action.triggered.connect(lambda checked: self.job_list_dock.setVisible(checked))
+ self.job_list_dock.visibilityChanged.connect(show_job_list_action.setChecked)
view_menu.addAction(show_job_list_action)
show_manual_sample_action = QAction("Show manual sample", self)
show_manual_sample_action.setCheckable(True)
show_manual_sample_action.setChecked(True)
- show_manual_sample_action.triggered.connect(self.manual_sample_dock.setVisible)
- # When dock visibility changes, update the action's checked state
- show_manual_sample_action.triggered.connect(lambda: self.manual_sample_dock.setVisible(show_manual_sample_action.isChecked()))
+ show_manual_sample_action.triggered.connect(lambda checked: self.manual_sample_dock.setVisible(checked))
+ self.manual_sample_dock.visibilityChanged.connect(show_manual_sample_action.setChecked)
view_menu.addAction(show_manual_sample_action)
show_face_panel_action = QAction("Show face detection", self)
show_face_panel_action.setCheckable(True)
show_face_panel_action.setChecked(False)
- #show_face_panel_action.triggered.connect(self.face_panel_dock.setVisible)
show_face_panel_action.triggered.connect(lambda checked: self.face_panel_dock.setVisible(checked))
self.face_panel_dock.visibilityChanged.connect(show_face_panel_action.setChecked)
view_menu.addAction(show_face_panel_action)
@@ -566,11 +676,19 @@ class MainWindow(QMainWindow):
show_smargon_trace_action.setCheckable(True)
show_smargon_trace_action.setChecked(False)
show_smargon_trace_action.triggered.connect(lambda checked: self.smargon_trace_dock.setVisible(checked))
+ self.smargon_trace_dock.visibilityChanged.connect(show_smargon_trace_action.setChecked)
self.smargon_trace_dock.visibilityChanged.connect(
lambda visible: self.smargon_trace_panel.refresh_plot(force=True) if visible else None
)
view_menu.addAction(show_smargon_trace_action)
+ show_workflow_action = QAction("Show Workflow Panel", self)
+ show_workflow_action.setCheckable(True)
+ show_workflow_action.setChecked(False)
+ show_workflow_action.triggered.connect(lambda checked: self.workflow_dock.setVisible(checked))
+ self.workflow_dock.visibilityChanged.connect(show_workflow_action.setChecked)
+ view_menu.addAction(show_workflow_action)
+
show_log_action = QAction("Show Log", self)
show_log_action.setCheckable(True)
show_log_action.setChecked(False)
@@ -578,6 +696,12 @@ class MainWindow(QMainWindow):
self.log_dock.visibilityChanged.connect(show_log_action.setChecked)
view_menu.addAction(show_log_action)
+ view_menu.addSeparator()
+
+ restore_default_view_action = QAction("Restore Default View", self)
+ restore_default_view_action.triggered.connect(self.restore_default_view)
+ view_menu.addAction(restore_default_view_action)
+
help_menu = menu_bar.addMenu("Help")
about_action = QAction("About", self) # Action for 'About'
about_action.triggered.connect(self.show_about_dialog)
@@ -606,6 +730,32 @@ class MainWindow(QMainWindow):
start_interactive_tutorial_action.triggered.connect(self.start_interactive_tutorial)
help_menu.addAction(start_interactive_tutorial_action)
+ def _capture_default_window_state(self) -> None:
+ self._default_window_state = self.saveState()
+
+ @Slot()
+ def restore_default_view(self) -> None:
+ if self._default_window_state is not None:
+ self.restoreState(self._default_window_state)
+
+ self.tell_samples_dock.setVisible(True)
+ self.job_list_dock.setVisible(True)
+ self.manual_sample_dock.setVisible(True)
+
+ self.face_panel_dock.setVisible(False)
+ self.fluor_panel_dock.setVisible(False)
+ self.smargon_trace_dock.setVisible(False)
+ self.log_dock.setVisible(False)
+
+ if self.__decoded_token.staff:
+ self.ref_tools_dock.setVisible(True)
+ self.ref_tools_dock.raise_()
+ else:
+ self.ref_tools_dock.setVisible(False)
+ self.tell_samples_dock.raise_()
+
+ self.video_tab.setCurrentIndex(0)
+
def show_about_dialog(self):
QMessageBox.about(
self,
@@ -674,6 +824,189 @@ class MainWindow(QMainWindow):
self.__mounting = False
self.video_tab.setCurrentIndex(0)
+
+ # ========== BATON DIALOG HANDLING ==========
+
+ @Slot(dict)
+ def _show_baton_request_dialog(self, payload: dict):
+ requester = str(payload.get("requester") or "Another user")
+ timeout = int(payload.get("timeout") or 30)
+
+ if self._baton_request_dialog is not None and self._baton_request_dialog.isVisible():
+ return
+
+ self._baton_request_dialog = BatonRequestDialog(
+ requester=requester,
+ timeout_seconds=timeout,
+ parent=self,
+ )
+ self._baton_request_dialog.accepted_signal.connect(self.status_bar._on_baton_dialog_accepted)
+ self._baton_request_dialog.refused_signal.connect(self.status_bar._on_baton_dialog_refused)
+ self._baton_request_dialog.show()
+ self._baton_request_dialog.raise_()
+ self._baton_request_dialog.activateWindow()
+
+ @Slot()
+ def _accept_baton_request(self):
+ self.daq.respond_to_baton_request(True)
+
+ @Slot()
+ def _refuse_baton_request(self):
+ self.daq.respond_to_baton_request(False)
+
+ @Slot()
+ def _accept_baton_request(self):
+ self.daq.respond_to_baton_request(True)
+
+ @Slot()
+ def _refuse_baton_request(self):
+ self.daq.respond_to_baton_request(False)
+
+ @Slot(BatonStatus)
+ def _on_baton_status_changed(self, status: BatonStatus):
+ """Close the pending dialog immediately if the baton request has been resolved via SSE."""
+ if self._waiting_for_baton_response and not status.you_have_pending_request:
+ self._waiting_for_baton_response = False
+ self._close_baton_pending_dialog()
+
+ if status.you_are_holder:
+ self.alert_banner.show_message("Baton acquired!", False, auto_clear_ms=10000)
+ else:
+ self.alert_banner.show_message("Request declined or cancelled", False, auto_clear_ms=10000)
+
+ @Slot(dict)
+ def _on_baton_request_result(self, result: dict):
+ """Handle result of our baton request - show waiting banner with countdown."""
+ if result.get("granted"):
+ self._waiting_for_baton_response = False
+ self.alert_banner.show_message("Baton acquired!", False, auto_clear_ms=10000)
+ logger.info("Baton acquired")
+
+ # Close the pending dialog immediately before showing p-group prompt
+ self._close_baton_pending_dialog()
+
+ available_pgroups = [str(p).strip() for p in (self.__decoded_token.pgroups or []) if p is not None and str(p).strip()]
+ if len(available_pgroups) == 1:
+ self.status_bar.set_pgroup.emit(available_pgroups[0])
+ else:
+ self.status_bar._after_baton_granted_select_pgroup()
+
+ elif result.get("pending"):
+ self._waiting_for_baton_response = True
+ timeout = result.get("timeout_seconds", 30)
+ holder = result.get("message", "Waiting for response...")
+
+ if getattr(self, "_baton_pending_dialog", None) is None:
+ target_user = holder.replace("Request sent to ", "")
+ self._baton_pending_dialog = BatonPendingDialog(target_user=target_user, timeout_seconds=timeout,
+ parent=self)
+ self._baton_pending_dialog.cancelled_signal.connect(self.daq.cancel_baton_request)
+ self._baton_pending_dialog.show()
+ else:
+ self._baton_pending_dialog.update_remaining(timeout)
+
+ self.alert_banner.show_waiting(f"Requesting control - {holder}", timeout)
+ logger.info(f"Baton request pending - {timeout}s timeout")
+
+ elif result.get("queued"):
+ self._waiting_for_baton_response = True
+ self.alert_banner.show_waiting("Control transfer queued - waiting for beamline")
+ logger.info("Baton transfer queued")
+
+ if getattr(self, "_baton_pending_dialog", None) is None:
+ self._baton_pending_dialog = BatonPendingDialog(target_user="Current Holder", timeout_seconds=0,
+ parent=self)
+ self._baton_pending_dialog.cancelled_signal.connect(self.daq.cancel_baton_request)
+ self._baton_pending_dialog.show()
+ self._baton_pending_dialog.set_queued_state()
+
+ elif result.get("already_holder"):
+ self._waiting_for_baton_response = False
+ logger.debug("Already baton holder")
+ self._close_baton_pending_dialog()
+
+ elif result.get("error"):
+ self._waiting_for_baton_response = False
+ self.alert_banner.show_message(result.get("message", "Request failed"), True)
+ logger.warning(f"Baton request failed: {result.get('message')}")
+ self._close_baton_pending_dialog()
+
+ @Slot(dict)
+ def _on_baton_response_result(self, result: dict):
+ """Handle result after we responded to someone else's request."""
+ logger.debug(f"Baton response result: {result}")
+ if result.get("accepted"):
+ self._waiting_for_baton_response = False
+ self.alert_banner.show_message("Control transferred", False, auto_clear_ms=10000)
+ self._close_baton_dialog()
+ self.status_bar.update_baton_status(self.status_bar._baton_status) # refresh label state
+ elif result.get("refused"):
+ self._waiting_for_baton_response = False
+ self.alert_banner.show_message("Request declined", False, auto_clear_ms=10000)
+ self._close_baton_dialog()
+ self.status_bar.update_baton_status(self.status_bar._baton_status) # refresh label state
+ else:
+ logger.debug(f"replied with {result}")
+
+ @Slot(dict)
+ def _on_baton_timeout_checked(self, result: dict):
+ """Refresh waiting UI when the backend confirms timeout state."""
+ logger.debug(f"Baton timeout checked: {result}")
+ if result.get("pending"):
+ remaining = int(result.get("remaining_seconds", 0))
+ if self._waiting_for_baton_response:
+ self.alert_banner.show_waiting("Requesting control", remaining)
+ if getattr(self, "_baton_pending_dialog", None) is not None:
+ self._baton_pending_dialog.update_remaining(remaining)
+
+ elif result.get("granted"):
+ self._waiting_for_baton_response = False
+ self.alert_banner.show_message("Baton acquired!", False, auto_clear_ms=10000)
+
+ # Close the pending dialog immediately before showing p-group prompt
+ self._close_baton_pending_dialog()
+ # P-group logic will be handled automatically by the status_bar stream update
+
+ elif result.get("queued"):
+ self._waiting_for_baton_response = True
+ self.alert_banner.show_waiting("Control transfer queued - waiting for beamline")
+
+ if getattr(self, "_baton_pending_dialog", None) is not None:
+ self._baton_pending_dialog.set_queued_state()
+ else:
+ self._baton_pending_dialog = BatonPendingDialog(target_user="Current Holder", timeout_seconds=0,
+ parent=self)
+ self._baton_pending_dialog.cancelled_signal.connect(self.daq.cancel_baton_request)
+ self._baton_pending_dialog.show()
+ self._baton_pending_dialog.set_queued_state()
+
+ elif result.get("refused"):
+ self._waiting_for_baton_response = False
+ self.alert_banner.show_message("Request declined", False, auto_clear_ms=10000)
+ self._close_baton_pending_dialog()
+
+ else:
+ logger.debug(f"replied with {result}")
+ self.alert_banner.clear_message()
+ self._close_baton_pending_dialog()
+
+ def _close_baton_dialog(self) -> None:
+ if getattr(self, "_baton_request_dialog", None) is not None:
+ try:
+ self._baton_request_dialog.close()
+ finally:
+ self._baton_request_dialog = None
+
+ def _close_baton_pending_dialog(self) -> None:
+ if getattr(self, "_baton_pending_dialog", None) is not None:
+ try:
+ if hasattr(self._baton_pending_dialog, '_timer'):
+ self._baton_pending_dialog._timer.stop()
+ self._baton_pending_dialog.close()
+ finally:
+ self._baton_pending_dialog = None
+
+
def _restore_window_state(self) -> None:
settings = QSettings()
geometry = settings.value("main_window/geometry")
@@ -684,6 +1017,9 @@ class MainWindow(QMainWindow):
if state is not None:
self.restoreState(state)
+ if not self.__decoded_token.staff:
+ self.ref_tools_dock.hide()
+
def closeEvent(self, event) -> None:
try:
settings = QSettings()
@@ -692,6 +1028,12 @@ class MainWindow(QMainWindow):
except Exception as e:
logger.warning(f"Failed to save main window state: {e}")
+ # Release baton before closing
+ try:
+ self.daq.release_baton_on_close()
+ except Exception as e:
+ logger.warning(f"Failed to release baton on close: {e}")
+
try:
self.cleanup()
except Exception as e:
@@ -710,6 +1052,9 @@ class MainWindow(QMainWindow):
except Exception as e:
logger.warning(f"Failed to stop _samcam_source_timer: {e}")
+ if hasattr(self, "workflow_sse") and self.workflow_sse is not None:
+ self.workflow_sse.disconnect()
+
for attr_name in (
"camera_thread",
"prediction_thread",
diff --git a/src/aare/gui/panels/abr_tweak_panel.py b/src/aare/gui/panels/abr_tweak_panel.py
index 8741f0dd..9eb589cf 100644
--- a/src/aare/gui/panels/abr_tweak_panel.py
+++ b/src/aare/gui/panels/abr_tweak_panel.py
@@ -1,7 +1,7 @@
from PySide6.QtCore import Signal, Slot
from PySide6.QtGui import Qt
from PySide6.QtWidgets import QWidget, QGridLayout, QLabel, QPushButton
-from aare.common.coordinate import Coordinate
+from aare.common.coordinate import Coordinate, AerotechCoordinate
from aare.common.models import DAQStatusModel
from aare.gui.widgets.button_with_payload import ButtonWithPayload
@@ -11,7 +11,7 @@ from aare.gui.widgets.title_label import TitleLabel
DEFAULT_ABR_STEP_UM = 5
class AbrTweakButtons(QWidget):
- abr_tweak = Signal(Coordinate)
+ abr_tweak = Signal(AerotechCoordinate)
def __init__(self, step_mm, parent=None):
super().__init__(parent)
@@ -66,14 +66,14 @@ class AbrTweakButtons(QWidget):
@Slot(dict)
def abr_button(self, payload: dict):
- self.abr_tweak.emit(Coordinate(x=self.__step_mm * payload["x"], y=self.__step_mm * payload["y"], z=self.__step_mm * payload["z"]))
+ self.abr_tweak.emit(AerotechCoordinate(at_mm=Coordinate(x=self.__step_mm * payload["x"], y=self.__step_mm * payload["y"], z=self.__step_mm * payload["z"])))
@Slot(float)
def set_step(self, val_um: float):
self.__step_mm = val_um / 1000.0
class AbrTweakWidget(QWidget):
- abr_tweak = Signal(Coordinate)
+ abr_tweak = Signal(AerotechCoordinate)
abr_save = Signal()
abr_goto_meas = Signal()
@@ -109,8 +109,8 @@ class AbrTweakWidget(QWidget):
def goto_button_pressed(self):
self.abr_goto_meas.emit()
- @Slot(Coordinate)
- def abr_button_pressed(self, c: Coordinate):
+ @Slot(AerotechCoordinate)
+ def abr_button_pressed(self, c: AerotechCoordinate):
self.abr_tweak.emit(c)
@Slot()
diff --git a/src/aare/gui/panels/automation_panel.py b/src/aare/gui/panels/automation_panel.py
new file mode 100644
index 00000000..3fd82e5c
--- /dev/null
+++ b/src/aare/gui/panels/automation_panel.py
@@ -0,0 +1,730 @@
+"""
+Workflow automation panel for queue management and step control.
+
+Displays:
+- Queue items with status (supports drag-drop from sample list)
+- Current step progress
+- Control buttons (Start Guided/Start Automation/Pause/Resume/Abort/Skip)
+- Live status from SSE
+"""
+
+from __future__ import annotations
+
+import json
+from typing import Any
+
+from PySide6.QtCore import Qt, Signal, Slot, QTimer, QMimeData
+from PySide6.QtGui import QColor, QDragEnterEvent, QDropEvent, QKeySequence, QShortcut
+from PySide6.QtWidgets import (
+ QWidget,
+ QVBoxLayout,
+ QHBoxLayout,
+ QLabel,
+ QPushButton,
+ QListWidget,
+ QListWidgetItem,
+ QProgressBar,
+ QGroupBox,
+ QFrame,
+ QSizePolicy,
+ QAbstractItemView,
+ QMenu,
+ QMessageBox,
+)
+
+from aare.common.automation_models import (
+ QueueItem,
+ QueueItemStatus,
+ RuntimeState,
+ ControlState,
+ WorkflowStepRecord,
+)
+from aare.common.models import SampleShortInfo, SampleShortInfoList
+from aare.common.logger_config import setup_logger
+
+logger = setup_logger("aareGUI")
+
+
+class StepProgressWidget(QWidget):
+ """Shows progress through workflow steps."""
+
+ def __init__(self, parent: QWidget | None = None):
+ super().__init__(parent)
+ self._steps: list[WorkflowStepRecord] = []
+ self._current_index = 0
+ self._setup_ui()
+
+ def _setup_ui(self) -> None:
+ layout = QVBoxLayout(self)
+ layout.setContentsMargins(4, 4, 4, 4)
+ layout.setSpacing(2)
+
+ self._step_labels: list[QLabel] = []
+
+ self._container = QWidget()
+ self._container_layout = QVBoxLayout(self._container)
+ self._container_layout.setContentsMargins(0, 0, 0, 0)
+ self._container_layout.setSpacing(2)
+ layout.addWidget(self._container)
+
+ def set_steps(self, steps: list[WorkflowStepRecord], current_index: int) -> None:
+ """Update the step display."""
+ self._steps = steps
+ self._current_index = current_index
+
+ for lbl in self._step_labels:
+ lbl.deleteLater()
+ self._step_labels.clear()
+
+ for i, step in enumerate(steps):
+ lbl = QLabel(f"{i + 1}. {step.kind}")
+ lbl.setStyleSheet(self._style_for_step(i, step.status))
+ self._container_layout.addWidget(lbl)
+ self._step_labels.append(lbl)
+
+ def update_step(self, index: int, status: str) -> None:
+ """Update a single step's status."""
+ if 0 <= index < len(self._step_labels):
+ self._step_labels[index].setStyleSheet(self._style_for_step(index, status))
+
+ def clear_steps(self) -> None:
+ """Clear all step labels."""
+ for lbl in self._step_labels:
+ lbl.deleteLater()
+ self._step_labels.clear()
+ self._steps = []
+
+ def _style_for_step(self, index: int, status: str) -> str:
+ """Get stylesheet for step based on status."""
+ base = "padding: 4px; border-radius: 3px; "
+
+ if status == "success":
+ return base + "background-color: #90EE90; color: #006400;"
+ elif status == "running":
+ return base + "background-color: #87CEEB; color: #00008B; font-weight: bold;"
+ elif status == "failed":
+ return base + "background-color: #FFB6C1; color: #8B0000;"
+ elif status == "skipped":
+ return base + "background-color: #D3D3D3; color: #696969; text-decoration: line-through;"
+ elif status == "paused":
+ return base + "background-color: #FFE4B5; color: #8B4513;"
+ else: # pending
+ return base + "background-color: #F0F0F0; color: #808080;"
+
+
+class DraggableQueueListWidget(QListWidget):
+ """
+ QListWidget that accepts drops from TellSamplePanel.
+
+ Supports:
+ - Drag-drop samples from tell_sample_panel
+ - Internal reordering via drag
+ - Delete key to remove items
+ """
+
+ samples_dropped = Signal(list) # list[SampleShortInfo]
+ item_reordered = Signal(str, int) # item_id, new_index
+ delete_requested = Signal(list) # list[item_ids]
+ start_item_requested = Signal(str) # item_id
+
+ def __init__(self, parent: QWidget | None = None):
+ super().__init__(parent)
+
+ self.setAcceptDrops(True)
+ self.setDragEnabled(True)
+ self.setDragDropMode(QAbstractItemView.DragDropMode.DragDrop)
+ self.setDefaultDropAction(Qt.DropAction.MoveAction)
+ self.setSelectionMode(QAbstractItemView.SelectionMode.ExtendedSelection)
+ self.setContextMenuPolicy(Qt.ContextMenuPolicy.CustomContextMenu)
+ self.customContextMenuRequested.connect(self._show_context_menu)
+
+ # Delete shortcut
+ self._delete_shortcut = QShortcut(QKeySequence.StandardKey.Delete, self)
+ self._delete_shortcut.activated.connect(self._on_delete_pressed)
+
+ # Store item_id -> row mapping
+ self._item_ids: list[str] = []
+
+ def set_item_ids(self, item_ids: list[str]) -> None:
+ """Track item IDs for reordering."""
+ self._item_ids = item_ids
+
+ def dragEnterEvent(self, event: QDragEnterEvent) -> None:
+ """Accept drops from sample panels."""
+ mime = event.mimeData()
+ if mime.hasText():
+ try:
+ text = mime.text()
+ if text.startswith("{") or text.startswith("["):
+ event.acceptProposedAction()
+ return
+ except Exception:
+ pass
+ if event.source() == self:
+ event.acceptProposedAction()
+ return
+ event.ignore()
+
+ def dragMoveEvent(self, event) -> None:
+ """Show drop indicator."""
+ if event.mimeData().hasText() or event.source() == self:
+ event.acceptProposedAction()
+ else:
+ event.ignore()
+
+ def dropEvent(self, event: QDropEvent) -> None:
+ """Handle drop - either samples from panel or internal reorder."""
+ mime = event.mimeData()
+
+ if mime.hasText():
+ text = mime.text()
+ try:
+ sample_list = SampleShortInfoList.model_validate_json(text)
+ if sample_list.s:
+ self.samples_dropped.emit(sample_list.s)
+ event.acceptProposedAction()
+ return
+ except Exception:
+ pass
+
+ try:
+ sample = SampleShortInfo.model_validate_json(text)
+ self.samples_dropped.emit([sample])
+ event.acceptProposedAction()
+ return
+ except Exception:
+ pass
+
+ if event.source() == self:
+ drop_row = self.indexAt(event.position().toPoint()).row()
+ if drop_row < 0:
+ drop_row = self.count()
+
+ selected = self.selectedItems()
+ if selected and self._item_ids:
+ for item in selected:
+ row = self.row(item)
+ if 0 <= row < len(self._item_ids):
+ item_id = self._item_ids[row]
+ self.item_reordered.emit(item_id, drop_row)
+
+ event.acceptProposedAction()
+ return
+
+ event.ignore()
+
+ def _on_delete_pressed(self) -> None:
+ """Handle delete key press."""
+ selected = self.selectedItems()
+ if not selected:
+ return
+
+ item_ids = []
+ for item in selected:
+ row = self.row(item)
+ if 0 <= row < len(self._item_ids):
+ item_ids.append(self._item_ids[row])
+
+ if item_ids:
+ self.delete_requested.emit(item_ids)
+
+ def _show_context_menu(self, pos) -> None:
+ """Show context menu for queue items."""
+ item = self.itemAt(pos)
+ if not item:
+ return
+
+ row = self.row(item)
+ if row < 0 or row >= len(self._item_ids):
+ return
+
+ menu = QMenu(self)
+
+ delete_action = menu.addAction("🗑 Remove from queue")
+ start_action = menu.addAction("▶ Start this item")
+
+ menu.addSeparator()
+ move_top_action = menu.addAction("⬆ Move to top")
+ move_bottom_action = menu.addAction("⬇ Move to bottom")
+
+ action = menu.exec_(self.mapToGlobal(pos))
+
+ if action == delete_action:
+ self.delete_requested.emit([self._item_ids[row]])
+ elif action == start_action:
+ self.start_item_requested.emit(self._item_ids[row])
+ elif action == move_top_action:
+ self.item_reordered.emit(self._item_ids[row], 0)
+ elif action == move_bottom_action:
+ self.item_reordered.emit(self._item_ids[row], 999999)
+
+
+class WorkflowPanel(QWidget):
+ """
+ Main workflow panel combining queue display and controls.
+
+ Supports:
+ - Drag-drop samples from TELL sample panel
+ - Queue management (delete, reorder, clear, add all)
+ - Step-by-step guided mode (explicit start)
+ - Full automation mode (explicit start)
+ """
+
+ # Signals for DAQ worker
+ request_queue_refresh = Signal()
+ request_add_sample = Signal(object) # SampleShortInfo
+ request_add_samples = Signal(list) # list[SampleShortInfo]
+ request_delete_item = Signal(str) # item_id
+ request_move_item = Signal(str, int) # item_id, new_order_index
+ request_clear_queue = Signal()
+ request_start_item = Signal(str) # item_id
+ request_next_step = Signal()
+ request_pause = Signal()
+ request_resume = Signal()
+ request_abort = Signal()
+ request_skip = Signal() # Skip current step
+ request_skip_sample = Signal() # Skip entire current sample
+ request_start_automation = Signal()
+ request_stop_automation = Signal()
+ request_start_guided = Signal(str) # Start guided mode with item_id
+
+ def __init__(self, parent: QWidget | None = None):
+ super().__init__(parent)
+ self._queue_items: list[QueueItem] = []
+ self._all_samples: list[SampleShortInfo] = []
+ self._runtime: RuntimeState | None = None
+ self._control: ControlState | None = None
+ self._automation_enabled = False
+ self._was_paused_before_automation = False # Track pause state before automation
+ self._setup_ui()
+
+ def _setup_ui(self) -> None:
+ layout = QVBoxLayout(self)
+ layout.setContentsMargins(8, 8, 8, 8)
+ layout.setSpacing(8)
+
+ # === Status Section ===
+ status_group = QGroupBox("Current Status")
+ status_layout = QVBoxLayout(status_group)
+
+ self._status_label = QLabel("⏹️ Idle")
+ self._status_label.setStyleSheet("font-size: 14px; font-weight: bold;")
+ status_layout.addWidget(self._status_label)
+
+ self._current_item_label = QLabel("No item running")
+ status_layout.addWidget(self._current_item_label)
+
+ self._step_progress = StepProgressWidget()
+ status_layout.addWidget(self._step_progress)
+
+ layout.addWidget(status_group)
+
+ # === Start Buttons (explicit start required) ===
+ start_group = QGroupBox("Start Processing")
+ start_layout = QHBoxLayout(start_group)
+
+ self._start_guided_btn = QPushButton("▶ Start Guided Mode")
+ self._start_guided_btn.setStyleSheet("background-color: #4CAF50; color: white; font-weight: bold; padding: 8px;")
+ self._start_guided_btn.setToolTip("Start processing the first pending sample in guided (step-by-step) mode")
+ self._start_guided_btn.clicked.connect(self._on_start_guided)
+
+ self._start_auto_btn = QPushButton("▶▶ Start Automation")
+ self._start_auto_btn.setStyleSheet("background-color: #2196F3; color: white; font-weight: bold; padding: 8px;")
+ self._start_auto_btn.setToolTip("Start fully automated processing of all samples")
+ self._start_auto_btn.clicked.connect(self._on_start_automation)
+
+ self._stop_btn = QPushButton("⏹ Stop")
+ self._stop_btn.setStyleSheet("background-color: #9E9E9E; color: white; font-weight: bold; padding: 8px;")
+ self._stop_btn.setToolTip("Stop automation mode (current step will complete)")
+ self._stop_btn.clicked.connect(self._on_stop)
+ self._stop_btn.setVisible(False)
+
+ start_layout.addWidget(self._start_guided_btn)
+ start_layout.addWidget(self._start_auto_btn)
+ start_layout.addWidget(self._stop_btn)
+
+ layout.addWidget(start_group)
+
+ # === Control Buttons ===
+ controls_group = QGroupBox("Step Controls")
+ controls_layout = QVBoxLayout(controls_group)
+
+ # Step controls row 1
+ step_layout = QHBoxLayout()
+
+ self._next_btn = QPushButton("Next Step")
+ self._next_btn.setStyleSheet("background-color: #4CAF50; color: white;")
+ self._next_btn.setToolTip("Execute the next step (guided mode only)")
+ self._next_btn.clicked.connect(self.request_next_step.emit)
+
+ self._skip_btn = QPushButton("Skip Step")
+ self._skip_btn.setStyleSheet("background-color: #FF9800; color: white;")
+ self._skip_btn.setToolTip("Skip the current step and move to the next")
+ self._skip_btn.clicked.connect(self.request_skip.emit)
+
+ self._skip_sample_btn = QPushButton("Skip Sample")
+ self._skip_sample_btn.setStyleSheet("background-color: #FF5722; color: white;")
+ self._skip_sample_btn.setToolTip("Skip the entire current sample and move to the next")
+ self._skip_sample_btn.clicked.connect(self.request_skip_sample.emit)
+
+ step_layout.addWidget(self._next_btn)
+ step_layout.addWidget(self._skip_btn)
+ step_layout.addWidget(self._skip_sample_btn)
+
+ controls_layout.addLayout(step_layout)
+
+ # Step controls row 2
+ control_layout2 = QHBoxLayout()
+
+ self._pause_btn = QPushButton("⏸ Pause")
+ self._pause_btn.setStyleSheet("background-color: #2196F3; color: white;")
+ self._pause_btn.clicked.connect(self._on_pause_resume)
+
+ self._abort_btn = QPushButton("🛑 Abort")
+ self._abort_btn.setStyleSheet("background-color: #F44336; color: white;")
+ self._abort_btn.clicked.connect(self.request_abort.emit)
+
+ control_layout2.addWidget(self._pause_btn)
+ control_layout2.addWidget(self._abort_btn)
+
+ controls_layout.addLayout(control_layout2)
+ layout.addWidget(controls_group)
+
+ # === Queue Section ===
+ queue_group = QGroupBox("Queue (drag samples here)")
+ queue_layout = QVBoxLayout(queue_group)
+
+ self._drop_hint = QLabel("💡 Drag samples from the Sample List to add them")
+ self._drop_hint.setStyleSheet("color: #666; font-style: italic;")
+ queue_layout.addWidget(self._drop_hint)
+
+ self._queue_list = DraggableQueueListWidget()
+ self._queue_list.setMinimumHeight(150)
+ self._queue_list.samples_dropped.connect(self._on_samples_dropped)
+ self._queue_list.delete_requested.connect(self._on_delete_requested)
+ self._queue_list.item_reordered.connect(self._on_item_reordered)
+ self._queue_list.start_item_requested.connect(self._on_start_specific_item)
+ queue_layout.addWidget(self._queue_list)
+
+ # Queue management buttons
+ queue_btn_layout = QHBoxLayout()
+
+ self._add_all_btn = QPushButton("➕ Add All")
+ self._add_all_btn.clicked.connect(self._on_add_all_clicked)
+ self._add_all_btn.setToolTip("Add all samples from the Sample List")
+
+ self._remove_selected_btn = QPushButton("🗑 Remove")
+ self._remove_selected_btn.clicked.connect(self._on_remove_selected)
+
+ self._clear_btn = QPushButton("✖ Clear")
+ self._clear_btn.clicked.connect(self._on_clear_queue)
+
+ self._refresh_btn = QPushButton("🔄")
+ self._refresh_btn.setFixedWidth(40)
+ self._refresh_btn.setToolTip("Refresh queue")
+ self._refresh_btn.clicked.connect(self.request_queue_refresh.emit)
+
+ queue_btn_layout.addWidget(self._add_all_btn)
+ queue_btn_layout.addWidget(self._remove_selected_btn)
+ queue_btn_layout.addWidget(self._clear_btn)
+ queue_btn_layout.addStretch()
+ queue_btn_layout.addWidget(self._refresh_btn)
+ queue_layout.addLayout(queue_btn_layout)
+
+ layout.addWidget(queue_group)
+
+ # Initial state
+ self._update_button_states()
+
+ def _on_start_guided(self) -> None:
+ """Start guided mode with the first pending sample."""
+ pending = [i for i in self._queue_items if i.status == QueueItemStatus.PENDING]
+ if not pending:
+ QMessageBox.information(self, "No Samples", "No pending samples in the queue.")
+ return
+
+ # Start the first pending item
+ self.request_start_item.emit(pending[0].item_id)
+
+ def _on_start_specific_item(self, item_id: str) -> None:
+ """Start guided mode with a specific sample."""
+ self.request_start_item.emit(item_id)
+
+ def _on_start_automation(self) -> None:
+ """Start full automation mode."""
+ pending = [i for i in self._queue_items if i.status == QueueItemStatus.PENDING]
+ is_running = self._runtime is not None and self._runtime.running
+
+ if not pending and not is_running:
+ QMessageBox.information(self, "No Samples", "No pending samples in the queue.")
+ return
+
+ # Remember if we were paused before starting automation
+ self._was_paused_before_automation = (
+ self._control is not None and self._control.pause_requested
+ )
+
+ self._automation_enabled = True
+ self.request_start_automation.emit()
+ self._update_button_states()
+
+ def _on_stop(self) -> None:
+ """Stop automation mode."""
+ self._automation_enabled = False
+ self.request_stop_automation.emit()
+
+ # Restore pause state if it was paused before automation started
+ if self._was_paused_before_automation:
+ self.request_pause.emit()
+
+ self._update_button_states()
+
+ def _on_pause_resume(self) -> None:
+ """Toggle pause/resume."""
+ if self._runtime and self._runtime.paused:
+ self.request_resume.emit()
+ else:
+ self.request_pause.emit()
+
+ def _on_samples_dropped(self, samples: list[SampleShortInfo]) -> None:
+ """Handle samples dropped onto the queue."""
+ logger.info(f"Adding {len(samples)} samples to workflow queue")
+ for sample in samples:
+ self.request_add_sample.emit(sample)
+ QTimer.singleShot(500, self.request_queue_refresh.emit)
+
+ def _on_delete_requested(self, item_ids: list[str]) -> None:
+ """Handle delete request from queue list."""
+ for item_id in item_ids:
+ self.request_delete_item.emit(item_id)
+ QTimer.singleShot(300, self.request_queue_refresh.emit)
+
+ def _on_item_reordered(self, item_id: str, new_index: int) -> None:
+ """Handle item reorder."""
+ self.request_move_item.emit(item_id, new_index)
+ QTimer.singleShot(300, self.request_queue_refresh.emit)
+
+ def _on_remove_selected(self) -> None:
+ """Remove selected items from queue."""
+ selected = self._queue_list.selectedItems()
+ if not selected:
+ return
+
+ item_ids = self._queue_list._item_ids
+ for item in selected:
+ row = self._queue_list.row(item)
+ if 0 <= row < len(item_ids):
+ self.request_delete_item.emit(item_ids[row])
+
+ QTimer.singleShot(300, self.request_queue_refresh.emit)
+
+ def _on_clear_queue(self) -> None:
+ """Clear all items from queue."""
+ if not self._queue_items:
+ return
+
+ reply = QMessageBox.question(
+ self,
+ "Clear Queue",
+ "Remove all non-running items from the queue?",
+ QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No,
+ )
+
+ if reply == QMessageBox.StandardButton.Yes:
+ self.request_clear_queue.emit()
+
+ QTimer.singleShot(200, self._clear_local_queue)
+
+ QTimer.singleShot(500, self.request_queue_refresh.emit)
+
+ def _clear_local_queue(self) -> None:
+ """Clear local queue state (called after server clear)."""
+ # Keep only running items
+ self._queue_items = [i for i in self._queue_items if i.status == QueueItemStatus.RUNNING]
+
+ # Update GUI
+ self._queue_list.clear()
+ item_ids = []
+ for item in self._queue_items:
+ status_emoji = "▶️"
+ display_text = f"{status_emoji} {item.sample_name or item.item_id}"
+ list_item = QListWidgetItem(display_text)
+ list_item.setBackground(QColor("#E6F3FF"))
+ self._queue_list.addItem(list_item)
+ item_ids.append(item.item_id)
+
+ self._queue_list.set_item_ids(item_ids)
+ self._update_button_states()
+
+ def _on_add_all_clicked(self) -> None:
+ """Add all available samples to the queue."""
+ if not self._all_samples:
+ QMessageBox.information(
+ self,
+ "No Samples",
+ "No samples available to add. Load samples in the Sample List first.",
+ )
+ return
+
+ reply = QMessageBox.question(
+ self,
+ "Add All Samples",
+ f"Add all {len(self._all_samples)} samples to the workflow queue?",
+ QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No,
+ )
+
+ if reply == QMessageBox.StandardButton.Yes:
+ self.request_add_samples.emit(self._all_samples)
+ QTimer.singleShot(500, self.request_queue_refresh.emit)
+
+ def _update_button_states(self) -> None:
+ """Update button enabled/disabled states based on current state."""
+ is_running = self._runtime is not None and self._runtime.running
+ is_paused = self._runtime is not None and self._runtime.paused
+ has_pending = any(i.status == QueueItemStatus.PENDING for i in self._queue_items)
+
+ self._start_guided_btn.setVisible(not is_running and not self._automation_enabled)
+
+ self._start_auto_btn.setVisible(not self._automation_enabled)
+ self._stop_btn.setVisible(self._automation_enabled)
+
+ self._start_guided_btn.setEnabled(has_pending)
+ self._start_auto_btn.setEnabled(has_pending or is_running)
+
+ self._next_btn.setVisible(not self._automation_enabled)
+ self._next_btn.setEnabled(is_running and not is_paused)
+
+ # Skip buttons - enabled when running
+ self._skip_btn.setEnabled(is_running)
+ self._skip_sample_btn.setEnabled(is_running)
+ self._abort_btn.setEnabled(is_running)
+
+ # Pause/Resume button
+ if is_paused:
+ self._pause_btn.setText("▶ Resume")
+ self._pause_btn.setStyleSheet("background-color: #4CAF50; color: white;")
+ else:
+ self._pause_btn.setText("⏸ Pause")
+ self._pause_btn.setStyleSheet("background-color: #2196F3; color: white;")
+ self._pause_btn.setEnabled(is_running)
+
+ # Drop hint
+ has_items = len(self._queue_items) > 0
+ self._drop_hint.setVisible(not has_items)
+
+ # === Public Update Methods ===
+
+ @Slot(list)
+ def update_queue(self, items: list[QueueItem]) -> None:
+ """Update queue display."""
+ self._queue_items = items
+ self._queue_list.clear()
+
+ item_ids = []
+ for item in items:
+ status_emoji = {
+ QueueItemStatus.PENDING: "⏳",
+ QueueItemStatus.RUNNING: "▶️",
+ QueueItemStatus.COMPLETED: "✅",
+ QueueItemStatus.FAILED: "❌",
+ QueueItemStatus.ABORTED: "🛑",
+ QueueItemStatus.SKIPPED: "⏭️",
+ }.get(item.status, "❓")
+
+ display_text = f"{status_emoji} {item.sample_name or item.item_id}"
+ list_item = QListWidgetItem(display_text)
+
+ if item.status == QueueItemStatus.RUNNING:
+ list_item.setBackground(QColor("#E6F3FF"))
+ elif item.status == QueueItemStatus.COMPLETED:
+ list_item.setBackground(QColor("#E6FFE6"))
+ elif item.status == QueueItemStatus.FAILED:
+ list_item.setBackground(QColor("#FFE6E6"))
+
+ self._queue_list.addItem(list_item)
+ item_ids.append(item.item_id)
+
+ self._queue_list.set_item_ids(item_ids)
+ self._update_button_states()
+
+ @Slot(object)
+ def update_runtime(self, runtime: RuntimeState) -> None:
+ """Update from runtime state."""
+ self._runtime = runtime
+
+ if runtime.running:
+ if runtime.paused:
+ self._status_label.setText("⏸️ Paused")
+ self._status_label.setStyleSheet(
+ "font-size: 14px; font-weight: bold; color: #FF9800;"
+ )
+ else:
+ mode_str = "Automation" if self._automation_enabled else "Guided"
+ self._status_label.setText(f"▶️ Running ({mode_str})")
+ self._status_label.setStyleSheet(
+ "font-size: 14px; font-weight: bold; color: #4CAF50;"
+ )
+
+ self._current_item_label.setText(
+ f"Item: {runtime.current_item_id or 'Unknown'}"
+ )
+ else:
+ self._status_label.setText("⏹️ Idle")
+ self._status_label.setStyleSheet(
+ "font-size: 14px; font-weight: bold; color: #757575;"
+ )
+ self._current_item_label.setText("No item running")
+ self._step_progress.clear_steps()
+
+ # Reset automation enabled when nothing is running
+ if self._automation_enabled and not runtime.running:
+ # Check if queue has more pending items
+ pending = [i for i in self._queue_items if i.status == QueueItemStatus.PENDING]
+ if not pending:
+ self._automation_enabled = False
+
+ self._update_button_states()
+
+ @Slot(object)
+ def update_control(self, control: ControlState) -> None:
+ """Update from control state."""
+ self._control = control
+ self._update_button_states()
+
+ @Slot(object)
+ def update_current_item(self, item: QueueItem) -> None:
+ """Update step progress for current item."""
+ self._step_progress.set_steps(item.steps, item.current_step_index)
+
+ for i, step in enumerate(item.steps):
+ self._step_progress.update_step(i, step.status)
+
+ @Slot(bool)
+ def update_automation_enabled(self, enabled: bool) -> None:
+ """Update automation mode state."""
+ self._automation_enabled = enabled
+ self._update_button_states()
+
+ @Slot(list)
+ def update_available_samples(self, samples: list[SampleShortInfo]) -> None:
+ """Update the cached sample list for 'Add All' functionality."""
+ self._all_samples = samples
+
+ @Slot(str, str)
+ def on_workflow_event(self, event_type: str, message: str) -> None:
+ """Handle workflow events from SSE stream."""
+ logger.debug(f"Workflow event: {event_type} - {message}")
+
+ if event_type in ("item_started", "item_completed", "step_completed", "step_skipped", "sample_skipped"):
+ self.request_queue_refresh.emit()
+
+ # Detect automation stop
+ if event_type == "automation_stopped":
+ self._automation_enabled = False
+
+ # Restore pause state if it was paused before automation started
+ if self._was_paused_before_automation:
+ self.request_pause.emit()
+ self._was_paused_before_automation = False
+
+ self._update_button_states()
\ No newline at end of file
diff --git a/src/aare/gui/panels/face_detection_panel.py b/src/aare/gui/panels/face_detection_panel.py
index 9551551d..2335d0f6 100644
--- a/src/aare/gui/panels/face_detection_panel.py
+++ b/src/aare/gui/panels/face_detection_panel.py
@@ -30,7 +30,7 @@ class FaceDetectionPanel(QWidget):
self.fig = Figure(figsize=(5, 4))
self.fig.subplots_adjust(
- left=0.12,
+ left=0.18,
right=0.97,
bottom=0.10,
top=0.95,
diff --git a/src/aare/gui/panels/loop_centering_panel.py b/src/aare/gui/panels/loop_centering_panel.py
index 41de4548..5c34d340 100644
--- a/src/aare/gui/panels/loop_centering_panel.py
+++ b/src/aare/gui/panels/loop_centering_panel.py
@@ -14,8 +14,6 @@ class LoopCenteringPanel(QWidget):
self.find_tip = QPushButton("Center", parent=self)
grid_layout.addWidget(self.find_tip, 1, 0)
- self.background = QPushButton("Bkg", parent=self)
- grid_layout.addWidget(self.background, 1, 1)
self.bounding_box = QPushButton("Box", parent=self)
grid_layout.addWidget(self.bounding_box, 1, 2)
diff --git a/src/aare/gui/panels/smargon_trace_panel.py b/src/aare/gui/panels/smargon_trace_panel.py
index fe958702..7f031cc4 100644
--- a/src/aare/gui/panels/smargon_trace_panel.py
+++ b/src/aare/gui/panels/smargon_trace_panel.py
@@ -507,6 +507,7 @@ class SmargonTracePanel(QWidget):
project_root / self._csv_path,
project_root / "src" / "aare" / "daq" / "logs" / "smargon_trace.csv",
project_root / "src" / "aare" / "gui" / "logs" / "smargon_trace.csv",
+ Path("/sls/mx/applications/logs/smargon_trace.csv"),
]
out: list[Path] = []
diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py
index 16f77b8d..c201a938 100644
--- a/src/aare/gui/threads/daq_worker.py
+++ b/src/aare/gui/threads/daq_worker.py
@@ -8,11 +8,14 @@ from PySide6.QtCore import Signal, QUrl, Slot, QTimer, QObject, QByteArray
from PySide6.QtNetwork import QNetworkAccessManager, QNetworkRequest, QNetworkReply
from jfjoch_client import ScanResult, ScanResultImagesInner
+from aare.common.coordinate import SmargonCoordinate, Coordinate, AerotechCoordinate
+from aare.common.auth_models import BatonStatus
from aare.common.coordinate import SmargonCoordinate, Coordinate
from aare.common.error_codes import export_error_codes
+from aare.common.exception_handler import JFJochCommunicationError
from aare.common.models import DAQStatusModel, SampleShortInfoList, SampleShortInfo, SampleCameraSettings, \
AutofocusSettings, SimpleScanParameters, FluorescenceSpectrumParameterModel, FluorescenceSpectrumOutputModel
-from aare.common.raster_grid import RasterGridRequest, CompletedRasterGrid
+from aare.common.raster_grid import RasterGridRequest, CompletedRasterGrid, CompletedRasterGridElem
from aare.common.rotation_scan import RotationScanRequest, CompletedRotationScan
from aare.common.logger_config import setup_logger
@@ -28,8 +31,13 @@ class DAQWorker(QObject):
reference_tools = Signal(SampleShortInfoList)
http_error = Signal(str)
status_message = Signal(str, bool)
+
+ # New dedicated signals for polled device errors and request-time errors
+ polled_devices_status = Signal(str, bool) # (message, is_error)
+ detector_error = Signal(str, bool) # (message, is_error)
auth_error = Signal()
sample_missing = Signal(str)
+
automated_scan_done = Signal(int, bool, str) # sample ID, success
run_number_incremented = Signal()
raster_scan_completed = Signal(CompletedRasterGrid)
@@ -45,8 +53,23 @@ class DAQWorker(QObject):
last_error_payload_changed = Signal(dict)
last_error_payloads_changed = Signal(list)
+ # Workflow signals
+ workflow_queue_loaded = Signal(list) # list[QueueItem]
+ workflow_runtime_changed = Signal(object) # RuntimeState
+ workflow_control_changed = Signal(object) # ControlState
+ workflow_item_updated = Signal(object) # QueueItem
+ workflow_automation_status = Signal(bool) # enabled
+ workflow_event = Signal(object) # WorkflowEvent
+
+ baton_status_changed = Signal(BatonStatus)
+ baton_request_result = Signal(dict)
+ baton_response_result = Signal(dict)
+ baton_incoming_request = Signal(dict)
+ baton_timeout_checked = Signal(dict)
+
def __init__(self, base_url: str | None, token: str, parent=None):
super().__init__(parent)
+ self._active_status_error_key = None
self.__token = token
self.__base_url = base_url
self.__net_manager = QNetworkAccessManager()
@@ -77,14 +100,37 @@ class DAQWorker(QObject):
self._last_error_payloads = deque(maxlen=10)
self._last_tell_connected: bool | None = None
+ self._last_tell_error: str | None = None
self._last_smargon_connected: bool | None = None
+ self._last_smargon_error: str | None = None
+ self._last_aerotech_connected: bool | None = None
+ self._last_aerotech_error: str | None = None
self._server_connected: bool | None = None
self._last_server_error: str | None = None
+ self._has_seen_disconnection: bool = False
+ self._server_was_disconnected: bool = False
+
+ # Deduplication tracking for polled device messages
+ self._last_polled_msg: str | None = None
+ self._last_polled_is_error: bool | None = None
+
+ # Deduplication tracking for detector/request-time messages
+ self._last_detector_msg: str | None = None
+ self._last_detector_is_error: bool | None = None
+
+ self._baton_stream_reply: QNetworkReply | None = None
+ self._last_baton_status: BatonStatus | None = None
+
+ self._baton_timeout_timer = QTimer(self)
+ self._baton_timeout_timer.setInterval(1000)
+ self._baton_timeout_timer.timeout.connect(self.check_baton_timeout)
+
self._face_detection_stream_reply: QNetworkReply | None = None
if self.__base_url is not None:
self.start_face_detection_stream()
+ self.start_baton_stream()
def get_last_error_payload(self) -> dict:
return dict(self._last_error_payload or {})
@@ -100,6 +146,46 @@ class DAQWorker(QObject):
self.last_error_payload_changed.emit(self.get_last_error_payload())
self.last_error_payloads_changed.emit(self.get_last_error_payloads())
+ def _emit_polled_device_message(self, msg: str, is_error: bool) -> None:
+ """Emit polled device status message with deduplication."""
+ if (msg, is_error) == (self._last_polled_msg, self._last_polled_is_error):
+ return
+ self._last_polled_msg = msg
+ self._last_polled_is_error = is_error
+ self.polled_devices_status.emit(msg, is_error)
+
+ def _emit_detector_message(self, msg: str, is_error: bool) -> None:
+ """Emit detector/request-time error message with deduplication."""
+ if (msg, is_error) == (self._last_detector_msg, self._last_detector_is_error):
+ return
+ self._last_detector_msg = msg
+ self._last_detector_is_error = is_error
+ self.detector_error.emit(msg, is_error)
+
+ def _emit_status_if_changed(self, key: str | None, message: str | None, is_error: bool) -> None:
+ """
+ Emit polled device status to the primary alert banner.
+ Used for Server/Tell/Smargon/Aerotech connection status.
+ """
+ if not message:
+ self._active_status_error_key = None
+ self._emit_polled_device_message("", False)
+ return
+
+ if is_error:
+ self._has_seen_disconnection = True
+ if self._active_status_error_key != key:
+ self._active_status_error_key = key
+ self._emit_polled_device_message(message, True)
+ return
+
+ if not self._has_seen_disconnection:
+ self._active_status_error_key = None
+ return
+
+ self._active_status_error_key = None
+ self._emit_polled_device_message(message, False)
+
def _log_smargon_throttled(self, *, endpoint: str | None, message: str) -> None:
"""
Log immediately if endpoint/message changed; otherwise at most every N seconds.
@@ -175,27 +261,15 @@ class DAQWorker(QObject):
logger.error(f"Error in response: {reply.errorString()}")
raise RuntimeError(reply.errorString())
- def _emit_status_if_changed(self, key: str | None, message: str | None, is_error: bool) -> None:
- if not message:
- self._active_status_error_key = None
- return
-
- if is_error:
- if self._active_status_error_key != key:
- self._active_status_error_key = key
- self.status_message.emit(message, True)
- return
-
- self._active_status_error_key = None
- self.status_message.emit(message, False)
-
def _compose_device_status_message(
self,
*,
tell_conn: bool,
smargon_conn: bool,
+ aerotech_conn: bool,
tell_changed: bool,
smargon_changed: bool,
+ aerotech_changed: bool,
) -> tuple[str | None, str | None, bool]:
disconnected: list[str] = []
restored: list[str] = []
@@ -210,25 +284,37 @@ class DAQWorker(QObject):
elif smargon_changed:
restored.append("Smargon")
+ if not aerotech_conn:
+ disconnected.append("Aerotech")
+ elif aerotech_changed:
+ restored.append("Aerotech")
+
if disconnected:
+ if len(disconnected) == 3:
+ return (
+ "all-devices-down",
+ "TELL, Smargon, and Aerotech disconnected.",
+ True,
+ )
if len(disconnected) == 2:
return (
- "tell+smargon-down",
- "TELL and Smargon connection errors, please inform your local contact.",
+ "+".join(sorted(d.lower() for d in disconnected)) + "-down",
+ f"{disconnected[0]} and {disconnected[1]} disconnected.",
True,
)
device = disconnected[0]
return (
f"{device.lower()}-down",
- f"{device} connection error, please inform your local contact.",
+ f"{device} disconnected.",
True,
)
if restored:
+ if len(restored) == 3:
+ return (None, "TELL, Smargon, and Aerotech reconnected.", False)
if len(restored) == 2:
- return (None, "TELL and Smargon connections restored.", False,)
- device = restored[0]
- return (None, f"{device} connection restored.",False)
+ return (None, f"{restored[0]} and {restored[1]} reconnected.", False)
+ return (None, f"{restored[0]} reconnected.", False)
return (None, None, False)
@@ -239,29 +325,68 @@ class DAQWorker(QObject):
parsed_response = DAQStatusModel.model_validate_json(response_data)
self.update.emit(parsed_response)
- if self._server_connected is False:
- self._emit_status_if_changed(None, "Server connection restored.", False)
+ # Handle server reconnection - show "Server reconnected" not device messages
+ if self._server_connected is False and self._has_seen_disconnection:
+ self._emit_status_if_changed(None, "Server reconnected.", False)
+ # Reset device states so we don't also emit device reconnection messages
+ self._last_tell_connected = None
+ self._last_smargon_connected = None
+ self._last_aerotech_connected = None
+ self._server_was_disconnected = False
+
self._server_connected = True
self._last_server_error = None
- tell_conn = bool(getattr(parsed_response, "tell_connected", True))
smargon_conn = bool(getattr(parsed_response, "smargon_connected", True))
- tell_err = getattr(parsed_response, "tell_error", None)
smargon_err = getattr(parsed_response, "smargon_error", None)
+ tell_conn = bool(getattr(parsed_response, "tell_connected", True))
+ tell_err = getattr(parsed_response, "tell_error", None)
+ aerotech_conn = bool(getattr(parsed_response, "aerotech_connected", True))
+ aerotech_err = getattr(parsed_response, "aerotech_error", None)
- tell_changed = self._last_tell_connected is not None and self._last_tell_connected != tell_conn
+ tell_err_text = None if tell_err is None else str(tell_err).strip()
+ smargon_err_text = None if smargon_err is None else str(smargon_err).strip()
+ aerotech_err_text = None if aerotech_err is None else str(aerotech_err).strip()
+
+ # Skip device status processing if we just reconnected from server down
+ # (we already showed "Server reconnected")
+ if self._last_tell_connected is None and self._last_smargon_connected is None and self._last_aerotech_connected is None:
+ # First status after startup or server reconnect - just record states, don't emit
+ self._last_tell_connected = tell_conn
+ self._last_smargon_connected = smargon_conn
+ self._last_aerotech_connected = aerotech_conn
+ self._last_tell_error = tell_err_text
+ self._last_smargon_error = smargon_err_text
+ self._last_aerotech_error = aerotech_err_text
+ return
+
+ tell_changed = (
+ self._last_tell_connected != tell_conn
+ or self._last_tell_error != tell_err_text
+ )
smargon_changed = (
- self._last_smargon_connected is not None and self._last_smargon_connected != smargon_conn
+ self._last_smargon_connected != smargon_conn
+ or self._last_smargon_error != smargon_err_text
+ )
+ aerotech_changed = (
+ self._last_aerotech_connected != aerotech_conn
+ or self._last_aerotech_error != aerotech_err_text
)
self._last_tell_connected = tell_conn
self._last_smargon_connected = smargon_conn
+ self._last_aerotech_connected = aerotech_conn
+ self._last_tell_error = tell_err_text
+ self._last_smargon_error = smargon_err_text
+ self._last_aerotech_error = aerotech_err_text
status_key, status_msg, is_error = self._compose_device_status_message(
tell_conn=tell_conn,
smargon_conn=smargon_conn,
+ aerotech_conn=aerotech_conn,
tell_changed=tell_changed,
smargon_changed=smargon_changed,
+ aerotech_changed=aerotech_changed,
)
self._emit_status_if_changed(status_key, status_msg, is_error)
@@ -269,14 +394,18 @@ class DAQWorker(QObject):
self._log_device_error_throttled(device="tell", message=tell_err)
if not smargon_conn:
self._log_device_error_throttled(device="smargon", message=smargon_err)
+ if not aerotech_conn:
+ self._log_device_error_throttled(device="aerotech", message=aerotech_err)
except Exception as e:
err_msg = str(e)
if self._server_connected is not False:
+ self._has_seen_disconnection = True
+ self._server_was_disconnected = True
self._emit_status_if_changed(
"server-down",
- "Server connection lost. Trying to reconnect...",
+ "Server disconnected. Reconnecting...",
True,
)
@@ -284,6 +413,7 @@ class DAQWorker(QObject):
self._last_server_error = err_msg
self._last_tell_connected = None
self._last_smargon_connected = None
+ self._last_aerotech_connected = None
logger.error(f"Exception from status response: {e}")
@@ -449,10 +579,6 @@ class DAQWorker(QObject):
def open_shutter(self):
self.generic_post(f"beamline/shutter?val=true")
- @Slot()
- def alc_background(self):
- self.generic_post("alc/background")
-
@Slot()
def center_loop(self):
self.generic_post("alc/center_loop")
@@ -547,6 +673,7 @@ class DAQWorker(QObject):
self.generic_delete("access/pgroup")
else:
self.generic_put(f"access/pgroup?val={val}")
+ self.send_status_request()
@Slot(QNetworkReply)
def _handle_all_pgroups_response(self, reply: QNetworkReply):
@@ -589,12 +716,38 @@ class DAQWorker(QObject):
def handle_rotation_scan_response(self, reply: QNetworkReply):
try:
+ # Check for HTTP errors first
+ if reply.error() != QNetworkReply.NetworkError.NoError:
+ status = reply.attribute(QNetworkRequest.Attribute.HttpStatusCodeAttribute)
+ err_msg = reply.errorString()
+
+ try:
+ raw_body = reply.readAll().data().decode("utf-8")
+ if raw_body:
+ body_json = json.loads(raw_body)
+ if isinstance(body_json, dict):
+ err_msg = body_json.get("message", err_msg)
+ code = body_json.get("code", "")
+
+ # Check if this is a JFJoch error
+ if code == "JFJOCH_UNAVAILABLE" or status == 503:
+ self._emit_detector_message(f"JFJoch: {err_msg}", is_error=True)
+ except Exception:
+ pass
+
+ logger.error(f"Rotation scan failed: {err_msg}")
+ self.http_error.emit(err_msg)
+ reply.deleteLater()
+ return
+
response_data = self.handle_response(reply)
parsed_response = CompletedRotationScan.model_validate_json(response_data)
self.standard_scan_completed.emit(parsed_response)
except Exception as e:
logger.error(f"Exception from rotation scan response: {e}")
self.http_error.emit(str(e))
+ finally:
+ reply.deleteLater()
@Slot(RotationScanRequest)
def standard_scan(self, r: RotationScanRequest):
@@ -609,12 +762,40 @@ class DAQWorker(QObject):
def handle_raster_scan_response(self, reply: QNetworkReply):
try:
+ # Check for HTTP errors first
+ if reply.error() != QNetworkReply.NetworkError.NoError:
+ status = reply.attribute(QNetworkRequest.Attribute.HttpStatusCodeAttribute)
+ raw_body = ""
+ body_json = None
+ err_msg = reply.errorString()
+
+ try:
+ raw_body = reply.readAll().data().decode("utf-8")
+ if raw_body:
+ body_json = json.loads(raw_body)
+ if isinstance(body_json, dict):
+ err_msg = body_json.get("message", err_msg)
+ code = body_json.get("code", "")
+
+ # Check if this is a JFJoch error
+ if code == "JFJOCH_UNAVAILABLE" or status == 503:
+ self._emit_detector_message(f"JFJoch: {err_msg}", is_error=True)
+ except Exception:
+ pass
+
+ logger.error(f"Raster scan failed: {err_msg}")
+ self.http_error.emit(err_msg)
+ reply.deleteLater()
+ return
+
response_data = self.handle_response(reply)
parsed_response = CompletedRasterGrid.model_validate_json(response_data)
self.raster_scan_completed.emit(parsed_response)
except Exception as e:
logger.error(f"Exception from raster scan response: {e}")
self.http_error.emit(str(e))
+ finally:
+ reply.deleteLater()
@Slot(RasterGridRequest)
def raster_scan(self, r: RasterGridRequest):
@@ -633,13 +814,14 @@ class DAQWorker(QObject):
bkg = random.gauss(3.0, 0.1),
spots= random.randint(0, 250),
index= random.randint(0, 1),
- mos = random.uniform(0, 0.1),
b= random.uniform(15.0, 80.0)
))
- logger.debug("check that this works - raster scan - complete raster grid")
- reply = CompletedRasterGrid(request = new_copy,
- result = ScanResult(file_prefix=r.file_prefix, images=images))
- logger.debug(f"It appears to work {reply}")
+ raster_elem = CompletedRasterGridElem(
+ request=new_copy,
+ result=ScanResult(file_prefix=r.file_prefix, images=images),
+ centre_of_mass = None,
+ )
+ reply = CompletedRasterGrid(r=[raster_elem])
self.raster_scan_completed.emit(reply)
return
@@ -666,13 +848,14 @@ class DAQWorker(QObject):
bkg=random.gauss(3.0, 0.1),
spots=random.randint(0, 250),
index=random.randint(0, 1),
- mos=random.uniform(0, 0.1),
b=random.uniform(15.0, 80.0)
))
- logger.debug("check that this works - raster scan auto - complete raster grid")
- reply = CompletedRasterGrid(request=new_copy,
- result=ScanResult(file_prefix=r.file_prefix, images=images))
- logger.debug(f"It appears to work {reply}")
+ raster_elem = CompletedRasterGridElem(
+ request=new_copy,
+ result=ScanResult(file_prefix=r.file_prefix, images=images),
+ centre_of_mass = None,
+ )
+ reply = CompletedRasterGrid(r=[raster_elem])
self.raster_scan_completed.emit(reply)
return
@@ -763,8 +946,8 @@ class DAQWorker(QObject):
return
self.generic_post("scan/smart_params", p.model_dump_json())
- @Slot(Coordinate)
- def abr_tweak(self, c: Coordinate):
+ @Slot(AerotechCoordinate)
+ def abr_tweak(self, c: AerotechCoordinate):
self.generic_post("beamline/tweak_abr_meas_pos", c.model_dump_json())
@Slot()
@@ -811,6 +994,10 @@ class DAQWorker(QObject):
def mount(self, s: SampleShortInfo, reference: bool = False):
self.generic_post(f"sample/mount?dbid={s.db_id}&reference={reference}")
+ @Slot()
+ def park_and_dry(self):
+ self.generic_post("sample/park_and_dry")
+
@Slot(SampleShortInfo)
def sample_manual(self, s: SampleShortInfo):
self.generic_post(f"sample/manual", s.model_dump_json())
@@ -1022,7 +1209,6 @@ class DAQWorker(QObject):
out[str(k)] = str(v)
return out
- @Slot()
@Slot()
def get_error_codes(self) -> None:
"""
@@ -1089,4 +1275,331 @@ class DAQWorker(QObject):
query.append(f"message={quote(message)}")
suffix = f"?{'&'.join(query)}" if query else ""
- self.generic_post(f"samcam/send_screenshot_db{suffix}")
\ No newline at end of file
+ self.generic_post(f"samcam/send_screenshot_db{suffix}")
+
+ def start_baton_stream(self):
+ """Start SSE stream for baton status updates."""
+ if self.__base_url is None:
+ return
+
+ if self._baton_stream_reply is not None:
+ return
+
+ request = QNetworkRequest(QUrl(f"{self.__base_url}/sse/baton"))
+ request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8"))
+ reply = self.__net_manager.get(request)
+ reply.readyRead.connect(lambda: self._read_baton_stream(reply))
+ reply.finished.connect(self._restart_baton_stream)
+ self._baton_stream_reply = reply
+
+ def _restart_baton_stream(self):
+ self._baton_stream_reply = None
+ if self.__base_url is not None:
+ QTimer.singleShot(1000, self.start_baton_stream)
+
+ def _read_baton_stream(self, reply: QNetworkReply):
+ try:
+ chunk = reply.readAll().data().decode("utf-8")
+ for line in chunk.splitlines():
+ if line.startswith("data:"):
+ payload = line[5:].strip()
+ if payload:
+ status = BatonStatus.model_validate_json(payload)
+
+ if status.you_have_pending_request:
+ if not self._baton_timeout_timer.isActive():
+ self._baton_timeout_timer.start()
+ else:
+ if self._baton_timeout_timer.isActive():
+ self._baton_timeout_timer.stop()
+
+ if (status.incoming_request and
+ (self._last_baton_status is None or
+ not self._last_baton_status.incoming_request)):
+ self.baton_incoming_request.emit({
+ "requester": status.pending_request.requester_username if status.pending_request else "Unknown",
+ "timeout": status.pending_request.timeout_seconds if status.pending_request else 30
+ })
+
+ self._last_baton_status = status
+ self.baton_status_changed.emit(status)
+ except Exception as e:
+ logger.error(f"Baton stream parse error: {e}")
+
+ @Slot()
+ def request_baton(self):
+ """Request the baton."""
+ if self.__base_url is None:
+ logger.info("POST /baton/request")
+ return
+
+ request = QNetworkRequest(QUrl(f"{self.__base_url}/baton/request"))
+ request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8"))
+ request.setRawHeader(b"Content-Type", b"application/json")
+ reply = self.__net_manager.post(request, QByteArray(b""))
+ reply.finished.connect(lambda: self._handle_baton_request_response(reply))
+
+ def _handle_baton_request_response(self, reply: QNetworkReply):
+ try:
+ response_data = self.handle_response(reply)
+ result = json.loads(response_data) if response_data else {}
+ self.baton_request_result.emit(result)
+
+ if result.get("granted"):
+ self.status_message.emit("Baton acquired", False)
+ self.send_status_request()
+ elif result.get("pending"):
+ self.status_message.emit(
+ f"Request sent - waiting for response ({result.get('timeout_seconds', 30)}s timeout)",
+ False
+ )
+ elif result.get("error"):
+ self.status_message.emit(result.get("message", "Request failed"), True)
+ except Exception as e:
+ logger.error(f"Baton request failed: {e}")
+ self.http_error.emit(str(e))
+
+ @Slot(bool)
+ def respond_to_baton_request(self, accept: bool):
+ """Respond to an incoming baton request."""
+ if self.__base_url is None:
+ logger.info(f"POST /baton/respond?accept={accept}")
+ return
+
+ request = QNetworkRequest(QUrl(f"{self.__base_url}/baton/respond?accept={str(accept).lower()}"))
+ request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8"))
+ request.setRawHeader(b"Content-Type", b"application/json")
+ reply = self.__net_manager.post(request, QByteArray(b""))
+ reply.finished.connect(lambda: self._handle_baton_response_result(reply))
+
+ def _handle_baton_response_result(self, reply: QNetworkReply):
+ try:
+ response_data = self.handle_response(reply)
+ result = json.loads(response_data) if response_data else {}
+ self.baton_response_result.emit(result)
+
+ if result.get("accepted") or result.get("refused"):
+ self.send_status_request()
+ self.check_baton_timeout()
+ self.start_baton_stream()
+ except Exception as e:
+ logger.error(f"Baton response failed: {e}")
+ self.http_error.emit(str(e))
+
+
+ @Slot()
+ def release_baton(self):
+ """Release the baton voluntarily."""
+ self.generic_post("baton/release")
+
+ @Slot()
+ def cancel_baton_request(self):
+ """Cancel your pending baton request."""
+ self.generic_post("baton/cancel")
+
+ @Slot()
+ def check_baton_timeout(self):
+ """Poll to check if timeout has been reached."""
+ if self.__base_url is None:
+ return
+
+ request = QNetworkRequest(QUrl(f"{self.__base_url}/baton/check_timeout"))
+ request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8"))
+ reply = self.__net_manager.get(request)
+ reply.finished.connect(lambda: self._handle_baton_timeout_response(reply))
+
+ def _handle_baton_timeout_response(self, reply: QNetworkReply):
+ try:
+ response_data = self.handle_response(reply)
+ result = json.loads(response_data) if response_data else {}
+ self.baton_timeout_checked.emit(result)
+
+ if result.get("granted") or result.get("queued") or result.get("refused"):
+ self.send_status_request()
+ self.start_baton_stream()
+ except Exception as e:
+ logger.error(f"Baton timeout check failed: {e}")
+ self.http_error.emit(str(e))
+
+ def release_baton_on_close(self):
+ """Release baton when GUI is closed to free the beamline."""
+ try:
+ if hasattr(self, "_baton_timeout_timer") and self._baton_timeout_timer is not None:
+ self._baton_timeout_timer.stop()
+ self.release_baton()
+ from PySide6.QtCore import QEventLoop, QTimer
+ loop = QEventLoop()
+ QTimer.singleShot(500, loop.quit)
+ loop.exec()
+ logger.info("Baton release requested on GUI close")
+ except Exception as e:
+ logger.warning(f"Error releasing baton on close: {e}")
+
+ # ─────────────────────────────────────────────
+ # Workflow API methods
+ # ─────────────────────────────────────────────
+
+ @Slot()
+ def workflow_load_queue(self):
+ """Load workflow queue."""
+ if self.__base_url is None:
+ return
+
+ request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/queue"))
+ request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode())
+ reply = self.__net_manager.get(request)
+ reply.finished.connect(lambda: self._handle_workflow_queue_response(reply))
+
+ def _handle_workflow_queue_response(self, reply: QNetworkReply):
+ try:
+ response_data = self.handle_response(reply)
+ data = json.loads(response_data)
+ from aare.common.automation_models import QueueItem, QueueItemStatus
+ items = [QueueItem.model_validate(i) for i in data.get("items", [])]
+ self.workflow_queue_loaded.emit(items)
+
+ # Also emit the currently running item for step progress
+ for item in items:
+ if item.status == QueueItemStatus.RUNNING:
+ self.workflow_item_updated.emit(item)
+ break
+ except Exception as e:
+ logger.error(f"Failed to load workflow queue: {e}")
+
+ @Slot(object) # SampleShortInfo
+ def workflow_add_sample(self, sample: "SampleShortInfo"):
+ """Add a sample to the workflow queue."""
+ if self.__base_url is None:
+ logger.info(f"POST /workflow/queue: {sample.sample_name}")
+ return
+
+ from aare.common.automation_models import CreateQueueItemRequest
+ request_data = CreateQueueItemRequest(
+ sample_id=sample.db_id,
+ sample_name=sample.sample_name,
+ priority=int(sample.priority or 100),
+ )
+
+ request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/queue"))
+ request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode())
+ request.setRawHeader(b"Content-Type", b"application/json")
+ reply = self.__net_manager.post(request, QByteArray(request_data.model_dump_json().encode()))
+ reply.finished.connect(lambda: self.handle_req_response(reply))
+
+ @Slot(list) # list[SampleShortInfo]
+ def workflow_add_samples(self, samples: list):
+ """Add multiple samples to the workflow queue."""
+ for sample in samples:
+ self.workflow_add_sample(sample)
+
+ @Slot(str)
+ def workflow_delete_item(self, item_id: str):
+ """Delete an item from the workflow queue."""
+ if self.__base_url is None:
+ logger.info(f"DELETE /workflow/queue/{item_id}")
+ return
+
+ self.generic_delete(f"workflow/queue/{item_id}")
+
+ @Slot(str, int)
+ def workflow_move_item(self, item_id: str, new_order_index: int):
+ """Move/reorder an item in the workflow queue."""
+ if self.__base_url is None:
+ logger.info(f"POST /workflow/queue/{item_id}/move order={new_order_index}")
+ return
+
+ body = json.dumps({"new_order_index": new_order_index})
+ self.generic_post(f"workflow/queue/{item_id}/move", body)
+
+ @Slot()
+ def workflow_clear_queue(self):
+ """Clear all non-running items from the queue via server."""
+ if self.__base_url is None:
+ return
+
+ request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/queue/clear"))
+ request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode())
+ request.setRawHeader(b"Content-Type", b"application/json")
+ reply = self.__net_manager.deleteResource(request)
+ reply.finished.connect(lambda: self.handle_req_response(reply))
+
+ @Slot()
+ def workflow_skip_sample(self):
+ """Skip the entire current sample."""
+ self.generic_post("workflow/control/skip_sample")
+
+ def _handle_clear_queue_response(self, reply: QNetworkReply):
+ try:
+ response_data = self.handle_response(reply)
+ data = json.loads(response_data)
+ from aare.common.automation_models import QueueItem, QueueItemStatus
+ items = [QueueItem.model_validate(i) for i in data.get("items", [])]
+
+ # Delete all non-running items
+ for item in items:
+ if item.status != QueueItemStatus.RUNNING:
+ self.workflow_delete_item(item.item_id)
+ except Exception as e:
+ logger.error(f"Failed to clear workflow queue: {e}")
+
+ @Slot(str)
+ def workflow_start_item(self, item_id: str):
+ """Start processing a queue item."""
+ self.generic_post(f"workflow/start/{item_id}")
+ QTimer.singleShot(500, self.workflow_load_queue)
+
+ @Slot()
+ def workflow_next_step(self):
+ """Run next step in guided mode."""
+ if self.__base_url is None:
+ return
+
+ request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/next"))
+ request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode())
+ request.setRawHeader(b"Content-Type", b"application/json")
+ reply = self.__net_manager.post(request, QByteArray(b""))
+ reply.finished.connect(lambda: self._handle_workflow_step_response(reply))
+
+ def _handle_workflow_step_response(self, reply: QNetworkReply):
+ try:
+ response_data = self.handle_response(reply)
+ data = json.loads(response_data)
+ if data.get("item"):
+ from aare.common.automation_models import QueueItem
+ item = QueueItem.model_validate(data["item"])
+ self.workflow_item_updated.emit(item)
+ # Also refresh queue
+ self.workflow_load_queue()
+ except Exception as e:
+ logger.error(f"Workflow step error: {e}")
+ self.http_error.emit(str(e))
+
+ @Slot()
+ def workflow_pause(self):
+ """Request pause."""
+ self.generic_post("workflow/control/pause")
+
+ @Slot()
+ def workflow_resume(self):
+ """Request resume."""
+ self.generic_post("workflow/control/resume")
+
+ @Slot()
+ def workflow_abort(self):
+ """Request abort."""
+ self.generic_post("workflow/control/abort")
+
+ @Slot()
+ def workflow_skip(self):
+ """Request skip."""
+ self.generic_post("workflow/control/skip")
+
+ @Slot()
+ def workflow_start_automation(self):
+ """Start automation mode."""
+ self.generic_post("workflow/automation/start")
+
+ @Slot()
+ def workflow_stop_automation(self):
+ """Stop automation mode."""
+ self.generic_post("workflow/automation/stop")
\ No newline at end of file
diff --git a/src/aare/gui/threads/workflow_sse_client.py b/src/aare/gui/threads/workflow_sse_client.py
new file mode 100644
index 00000000..ec264106
--- /dev/null
+++ b/src/aare/gui/threads/workflow_sse_client.py
@@ -0,0 +1,175 @@
+"""
+SSE client for workflow events.
+
+Connects to the /workflow/sse endpoint and emits signals for:
+- Runtime state changes
+- Control state changes
+- Individual workflow events
+"""
+
+from __future__ import annotations
+
+import json
+
+from PySide6.QtCore import QObject, Signal, Slot, QTimer
+from PySide6.QtNetwork import QNetworkAccessManager, QNetworkRequest, QNetworkReply
+from PySide6.QtCore import QUrl, QByteArray
+
+from aare.common.automation_models import RuntimeState, ControlState, WorkflowEvent
+from aare.common.logger_config import setup_logger
+
+logger = setup_logger("aareGUI")
+
+
+class WorkflowSSEClient(QObject):
+ """
+ SSE client that subscribes to workflow events.
+
+ Emits signals when state changes are received.
+ """
+
+ # Signals
+ runtime_changed = Signal(object) # RuntimeState
+ control_changed = Signal(object) # ControlState
+ workflow_event = Signal(object) # WorkflowEvent
+ connected = Signal()
+ disconnected = Signal()
+ error = Signal(str)
+
+ def __init__(
+ self,
+ base_url: str,
+ token: str,
+ parent: QObject | None = None,
+ ):
+ super().__init__(parent)
+
+ self._base_url = base_url
+ self._token = token
+ self._manager = QNetworkAccessManager(self)
+ self._reply: QNetworkReply | None = None
+ self._buffer = ""
+
+ # Reconnection
+ self._reconnect_timer = QTimer(self)
+ self._reconnect_timer.setInterval(5000) # 5 seconds
+ self._reconnect_timer.timeout.connect(self.connect)
+ self._should_reconnect = False
+
+ def connect(self) -> None:
+ """Start SSE connection."""
+ if self._reply is not None:
+ return # Already connected
+
+ self._should_reconnect = True
+
+ url = QUrl(f"{self._base_url}/workflow/sse")
+
+ request = QNetworkRequest(url)
+ request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode())
+ request.setRawHeader(b"Accept", b"text/event-stream")
+ request.setRawHeader(b"Cache-Control", b"no-cache")
+
+ self._reply = self._manager.get(request)
+ self._reply.readyRead.connect(self._on_data_ready)
+ self._reply.finished.connect(self._on_finished)
+ self._reply.errorOccurred.connect(self._on_error)
+
+ self._reconnect_timer.stop()
+ logger.debug("Workflow SSE: connecting...")
+
+ def disconnect(self) -> None:
+ """Stop SSE connection."""
+ self._should_reconnect = False
+ self._reconnect_timer.stop()
+
+ if self._reply is not None:
+ try:
+ self._reply.abort()
+ except Exception:
+ pass
+ try:
+ self._reply.deleteLater()
+ except Exception:
+ pass
+ self._reply = None
+
+ self.disconnected.emit()
+
+ @Slot()
+ def _on_data_ready(self) -> None:
+ """Handle incoming SSE data."""
+ if self._reply is None:
+ return
+
+ try:
+ data = self._reply.readAll().data().decode("utf-8")
+ self._buffer += data
+
+ # Process complete events (separated by double newlines)
+ while "\n\n" in self._buffer:
+ event_data, self._buffer = self._buffer.split("\n\n", 1)
+ self._parse_event(event_data)
+ except Exception as e:
+ logger.warning(f"Workflow SSE data read error: {e}")
+
+ def _parse_event(self, event_data: str) -> None:
+ """Parse a single SSE event."""
+ event_type = "message"
+ data_lines = []
+
+ for line in event_data.split("\n"):
+ if line.startswith("event:"):
+ event_type = line[6:].strip()
+ elif line.startswith("data:"):
+ data_lines.append(line[5:].strip())
+
+ if not data_lines:
+ return
+
+ data_str = "\n".join(data_lines)
+
+ try:
+ if event_type == "runtime":
+ runtime = RuntimeState.model_validate_json(data_str)
+ self.runtime_changed.emit(runtime)
+ elif event_type == "control":
+ control = ControlState.model_validate_json(data_str)
+ self.control_changed.emit(control)
+ elif event_type == "workflow_event":
+ event = WorkflowEvent.model_validate_json(data_str)
+ self.workflow_event.emit(event)
+ except Exception as e:
+ logger.warning(f"Workflow SSE: failed to parse {event_type}: {e}")
+
+ @Slot()
+ def _on_finished(self) -> None:
+ """Handle connection finished."""
+ if self._reply is not None:
+ try:
+ self._reply.deleteLater()
+ except Exception:
+ pass
+ self._reply = None
+
+ self._buffer = ""
+ self.disconnected.emit()
+
+ # Reconnect if desired
+ if self._should_reconnect:
+ logger.debug("Workflow SSE: disconnected, will reconnect...")
+ self._reconnect_timer.start()
+
+ @Slot(QNetworkReply.NetworkError)
+ def _on_error(self, error: QNetworkReply.NetworkError) -> None:
+ """Handle connection error."""
+ error_msg = ""
+ if self._reply is not None:
+ try:
+ error_msg = self._reply.errorString()
+ except Exception:
+ error_msg = str(error)
+ else:
+ error_msg = str(error)
+ logger.warning(f"Workflow SSE error: {error_msg}")
+ self.error.emit(error_msg)
\ No newline at end of file
diff --git a/src/aare/gui/widgets/alert_banner.py b/src/aare/gui/widgets/alert_banner.py
index fe10f4c7..7f18386d 100644
--- a/src/aare/gui/widgets/alert_banner.py
+++ b/src/aare/gui/widgets/alert_banner.py
@@ -15,6 +15,13 @@ class AlertBanner(QFrame):
self._clear_timer.setSingleShot(True)
self._clear_timer.timeout.connect(self.clear_message)
+ # Countdown timer for "waiting" state
+ self._countdown_timer = QTimer(self)
+ self._countdown_timer.setInterval(1000)
+ self._countdown_timer.timeout.connect(self._tick_countdown)
+ self._countdown_remaining = 0
+ self._countdown_base_message = ""
+
self._label = QLabel("", self)
self._label.setWordWrap(True)
self._label.setAlignment(Qt.AlignmentFlag.AlignCenter)
@@ -33,7 +40,9 @@ class AlertBanner(QFrame):
self.setSizePolicy(QSizePolicy.Policy.Expanding, QSizePolicy.Policy.Fixed)
@Slot(str, bool)
- def show_message(self, msg: str, is_error: bool = True):
+ def show_message(self, msg: str, is_error: bool = True, auto_clear_ms: int | None = None):
+ """Show error (red) or success (green) message."""
+ self._stop_countdown()
self._clear_timer.stop()
if not msg:
@@ -72,13 +81,86 @@ class AlertBanner(QFrame):
" padding: 2px 6px 2px 6px;"
"}"
)
- self._clear_timer.start(5000)
+ timeout = 5000 if auto_clear_ms is None else int(auto_clear_ms)
+ self._clear_timer.start(timeout)
+ self._current_message = msg
+ self._current_is_error = is_error
self._label.setText(decorated)
self.setVisible(True)
+ @Slot(str, int)
+ def show_waiting(self, msg: str, countdown_seconds: int = 0):
+ """
+ Show a waiting/pending message (yellow) with optional countdown.
+
+ Args:
+ msg: Base message to display
+ countdown_seconds: If > 0, append countdown and auto-update
+ """
+ self._clear_timer.stop()
+ self._stop_countdown()
+
+ if not msg:
+ self.clear_message()
+ return
+
+ self._countdown_base_message = msg
+ self._countdown_remaining = countdown_seconds
+
+ self._apply_waiting_style()
+ self._update_waiting_text()
+
+ if countdown_seconds > 0:
+ self._countdown_timer.start()
+
+ self.setVisible(True)
+
+ def _apply_waiting_style(self):
+ """Apply yellow/waiting style."""
+ self.setStyleSheet(
+ "QFrame {"
+ " background-color: #fff8e1;"
+ " border: 2px solid #ffb300;"
+ " border-radius: 12px;"
+ " margin: 8px 12px 8px 12px;"
+ "}"
+ "QLabel {"
+ " color: #e65100;"
+ " font-weight: 700;"
+ " font-size: 20px;"
+ " padding: 2px 6px 2px 6px;"
+ "}"
+ )
+
+ def _update_waiting_text(self):
+ """Update the waiting message text, including countdown if active."""
+ if self._countdown_remaining > 0:
+ decorated = f"⏳ {self._countdown_base_message} ({self._countdown_remaining}s) ⏳"
+ else:
+ decorated = f"⏳ {self._countdown_base_message} ⏳"
+ self._label.setText(decorated)
+
+ def _tick_countdown(self):
+ """Called every second during countdown."""
+ self._countdown_remaining -= 1
+ if self._countdown_remaining <= 0:
+ self._stop_countdown()
+ self.clear_message()
+ return
+
+ self._update_waiting_text()
+
+ def _stop_countdown(self):
+ """Stop the countdown timer."""
+ self._countdown_timer.stop()
+ self._countdown_remaining = 0
+ self._countdown_base_message = ""
+
@Slot()
def clear_message(self):
+ self._current_message = None
+ self._current_is_error = None
self._clear_timer.stop()
self._label.clear()
self.setVisible(False)
diff --git a/src/aare/gui/widgets/baton_request_dialog.py b/src/aare/gui/widgets/baton_request_dialog.py
new file mode 100644
index 00000000..b23bca2e
--- /dev/null
+++ b/src/aare/gui/widgets/baton_request_dialog.py
@@ -0,0 +1,348 @@
+from PySide6.QtCore import Qt, Signal, QTimer
+from PySide6.QtWidgets import (
+ QDialog, QVBoxLayout, QHBoxLayout, QLabel,
+ QPushButton, QProgressBar, QFrame
+)
+from PySide6.QtGui import QFont
+
+
+class BatonRequestDialog(QDialog):
+ """
+ Dialog shown to current baton holder when someone requests control.
+
+ Based on the workflow diagram:
+ - User can Accept (transfer immediately or queue if busy)
+ - User can Refuse (deny the request)
+ - If user ignores/closes, timeout causes auto-transfer
+ """
+
+ accepted_signal = Signal()
+ refused_signal = Signal()
+
+ def __init__(self, requester: str, timeout_seconds: int = 30, parent=None):
+ super().__init__(parent)
+ self.setWindowTitle("⚡ Baton Request")
+ self.setModal(False) # Non-modal so user can see beamline status
+ self.setMinimumWidth(400)
+ self.setWindowFlags(
+ self.windowFlags() |
+ Qt.WindowType.WindowStaysOnTopHint
+ )
+
+ self._timeout = timeout_seconds
+ self._remaining = timeout_seconds
+ self._requester = requester
+
+ self._setup_ui()
+ self._start_timer()
+
+ def _setup_ui(self):
+ layout = QVBoxLayout(self)
+ layout.setSpacing(15)
+
+ # Header
+ header = QLabel("🔔 Control Request")
+ header_font = QFont()
+ header_font.setPointSize(14)
+ header_font.setBold(True)
+ header.setFont(header_font)
+ header.setAlignment(Qt.AlignmentFlag.AlignCenter)
+ layout.addWidget(header)
+
+ # Separator
+ line = QFrame()
+ line.setFrameShape(QFrame.Shape.HLine)
+ line.setFrameShadow(QFrame.Shadow.Sunken)
+ layout.addWidget(line)
+
+ # Message
+ self.message_label = QLabel(
+ f"{self._requester} is requesting control of the beamline."
+ )
+ self.message_label.setWordWrap(True)
+ self.message_label.setAlignment(Qt.AlignmentFlag.AlignCenter)
+ layout.addWidget(self.message_label)
+
+ # Timeout progress
+ progress_layout = QVBoxLayout()
+
+ self.progress = QProgressBar()
+ self.progress.setRange(0, self._timeout)
+ self.progress.setValue(self._timeout)
+ self.progress.setTextVisible(False)
+ self.progress.setFixedHeight(8)
+ self.progress.setStyleSheet("""
+ QProgressBar {
+ border: 1px solid #ccc;
+ border-radius: 4px;
+ background-color: #f0f0f0;
+ }
+ QProgressBar::chunk {
+ background-color: #4CAF50;
+ border-radius: 3px;
+ }
+ """)
+ progress_layout.addWidget(self.progress)
+
+ self.time_label = QLabel(f"{self._timeout} seconds remaining")
+ self.time_label.setAlignment(Qt.AlignmentFlag.AlignCenter)
+ self.time_label.setStyleSheet("color: #666;")
+ progress_layout.addWidget(self.time_label)
+
+ layout.addLayout(progress_layout)
+
+ # Warning about auto-transfer
+ self.warning_label = QLabel(
+ "⚠️ If you don't respond, control will transfer automatically."
+ )
+ self.warning_label.setWordWrap(True)
+ self.warning_label.setAlignment(Qt.AlignmentFlag.AlignCenter)
+ self.warning_label.setStyleSheet("color: #ff9800; font-style: italic;")
+ layout.addWidget(self.warning_label)
+
+ # Buttons
+ button_layout = QHBoxLayout()
+ button_layout.setSpacing(20)
+
+ self.accept_btn = QPushButton("✓ Accept")
+ self.accept_btn.setMinimumHeight(40)
+ self.accept_btn.setStyleSheet("""
+ QPushButton {
+ background-color: #4CAF50;
+ color: white;
+ border: none;
+ border-radius: 5px;
+ font-weight: bold;
+ font-size: 13px;
+ }
+ QPushButton:hover {
+ background-color: #45a049;
+ }
+ QPushButton:pressed {
+ background-color: #3d8b40;
+ }
+ """)
+ self.accept_btn.clicked.connect(self._on_accept)
+ button_layout.addWidget(self.accept_btn)
+
+ self.refuse_btn = QPushButton("✗ Refuse")
+ self.refuse_btn.setMinimumHeight(40)
+ self.refuse_btn.setStyleSheet("""
+ QPushButton {
+ background-color: #f44336;
+ color: white;
+ border: none;
+ border-radius: 5px;
+ font-weight: bold;
+ font-size: 13px;
+ }
+ QPushButton:hover {
+ background-color: #da190b;
+ }
+ QPushButton:pressed {
+ background-color: #c41000;
+ }
+ """)
+ self.refuse_btn.clicked.connect(self._on_refuse)
+ button_layout.addWidget(self.refuse_btn)
+
+ layout.addLayout(button_layout)
+
+ # Info text
+ info_label = QLabel(
+ "If the beamline is busy, transfer will occur after "
+ "the current operation completes."
+ )
+ info_label.setWordWrap(True)
+ info_label.setAlignment(Qt.AlignmentFlag.AlignCenter)
+ info_label.setStyleSheet("color: #999;")
+ layout.addWidget(info_label)
+
+ def _start_timer(self):
+ self._timer = QTimer(self)
+ self._timer.setInterval(1000)
+ self._timer.timeout.connect(self._tick)
+ self._timer.start()
+
+ def _tick(self):
+ self._remaining -= 1
+ self.progress.setValue(self._remaining)
+ self.time_label.setText(f"{self._remaining} seconds remaining")
+
+ # Change progress bar color as time runs out
+ if self._remaining <= 10:
+ self.progress.setStyleSheet("""
+ QProgressBar {
+ border: 1px solid #ccc;
+ border-radius: 4px;
+ background-color: #f0f0f0;
+ }
+ QProgressBar::chunk {
+ background-color: #ff9800;
+ border-radius: 3px;
+ }
+ """)
+
+ if self._remaining <= 5:
+ self.progress.setStyleSheet("""
+ QProgressBar {
+ border: 1px solid #ccc;
+ border-radius: 4px;
+ background-color: #f0f0f0;
+ }
+ QProgressBar::chunk {
+ background-color: #f44336;
+ border-radius: 3px;
+ }
+ """)
+ self.time_label.setStyleSheet("color: #f44336; font-weight: bold;")
+
+ if self._remaining <= 0:
+ self._timer.stop()
+ # Timeout = auto-accept (as per your diagram: "Ignores request" → auto transfer)
+ self._on_accept()
+
+ def _on_accept(self):
+ self._timer.stop()
+ self.accepted_signal.emit()
+ self.accept()
+
+ def _on_refuse(self):
+ self._timer.stop()
+ self.refused_signal.emit()
+ self.reject()
+
+ def closeEvent(self, event):
+ """Closing the dialog counts as ignoring = auto-accept on timeout."""
+ # Don't emit anything here - let the timeout handle it
+ # or the SSE stream will close the dialog when resolved
+ self._timer.stop()
+ super().closeEvent(event)
+
+class BatonPendingDialog(QDialog):
+ """
+ Dialog shown to the user who requested the baton while they wait for a response
+ or for the beamline queue to clear.
+ """
+ cancelled_signal = Signal()
+
+ def __init__(self, target_user: str, timeout_seconds: int = 30, parent=None):
+ super().__init__(parent)
+ self.setWindowTitle("⏳ Baton Request Pending")
+ self.setModal(False)
+ self.setMinimumWidth(400)
+ self.setWindowFlags(self.windowFlags() | Qt.WindowType.WindowStaysOnTopHint)
+
+ self._timeout = timeout_seconds
+ self._remaining = timeout_seconds
+ self._target_user = target_user
+
+ self._setup_ui()
+ self._start_timer()
+
+ def _setup_ui(self):
+ layout = QVBoxLayout(self)
+ layout.setSpacing(15)
+
+ self.header = QLabel("⏳ Requesting Control")
+ header_font = QFont()
+ header_font.setPointSize(14)
+ header_font.setBold(True)
+ self.header.setFont(header_font)
+ self.header.setAlignment(Qt.AlignmentFlag.AlignCenter)
+ layout.addWidget(self.header)
+
+ line = QFrame()
+ line.setFrameShape(QFrame.Shape.HLine)
+ line.setFrameShadow(QFrame.Shadow.Sunken)
+ layout.addWidget(line)
+
+ self.message_label = QLabel(
+ f"Waiting for {self._target_user} to respond..."
+ )
+ self.message_label.setWordWrap(True)
+ self.message_label.setAlignment(Qt.AlignmentFlag.AlignCenter)
+ layout.addWidget(self.message_label)
+
+ self.progress_layout = QVBoxLayout()
+ self.progress = QProgressBar()
+ self.progress.setRange(0, max(1, self._timeout))
+ self.progress.setValue(self._timeout)
+ self.progress.setTextVisible(False)
+ self.progress.setFixedHeight(8)
+ self.progress.setStyleSheet("""
+ QProgressBar {
+ border: 1px solid #ccc;
+ border-radius: 4px;
+ background-color: #f0f0f0;
+ }
+ QProgressBar::chunk {
+ background-color: #2196F3;
+ border-radius: 3px;
+ }
+ """)
+ self.progress_layout.addWidget(self.progress)
+
+ self.time_label = QLabel(f"{self._timeout} seconds remaining")
+ self.time_label.setAlignment(Qt.AlignmentFlag.AlignCenter)
+ self.time_label.setStyleSheet("color: #666;")
+ self.progress_layout.addWidget(self.time_label)
+
+ layout.addLayout(self.progress_layout)
+
+ button_layout = QHBoxLayout()
+ self.cancel_btn = QPushButton("✗ Cancel Request")
+ self.cancel_btn.setMinimumHeight(40)
+ self.cancel_btn.setStyleSheet("""
+ QPushButton {
+ background-color: #f44336;
+ color: white;
+ border: none;
+ border-radius: 5px;
+ font-weight: bold;
+ font-size: 13px;
+ }
+ QPushButton:hover { background-color: #da190b; }
+ QPushButton:pressed { background-color: #c41000; }
+ """)
+ self.cancel_btn.clicked.connect(self._on_cancel)
+ button_layout.addWidget(self.cancel_btn)
+ layout.addLayout(button_layout)
+
+ def _start_timer(self):
+ self._timer = QTimer(self)
+ self._timer.setInterval(1000)
+ self._timer.timeout.connect(self._tick)
+ self._timer.start()
+
+ def _tick(self):
+ self._remaining -= 1
+ if self._remaining < 0:
+ self._remaining = 0
+
+ self.progress.setValue(self._remaining)
+ self.time_label.setText(f"{self._remaining} seconds remaining")
+ if self._remaining <= 0:
+ self._timer.stop()
+
+ def update_remaining(self, remaining: int):
+ self._remaining = remaining
+ self.progress.setValue(self._remaining)
+ self.time_label.setText(f"{self._remaining} seconds remaining")
+
+ def set_queued_state(self):
+ self._timer.stop()
+ self.header.setText("⏳ Transfer Queued")
+ self.message_label.setText("Waiting for current action to finish before receiving baton...")
+ self.progress.hide()
+ self.time_label.hide()
+ # Keep cancel button so they can abort the wait if they change their mind
+
+ def _on_cancel(self):
+ self._timer.stop()
+ self.cancelled_signal.emit()
+ self.reject()
+
+ def closeEvent(self, event):
+ self._timer.stop()
+ super().closeEvent(event)
\ No newline at end of file
diff --git a/src/aare/gui/widgets/pgroup_dialog.py b/src/aare/gui/widgets/pgroup_dialog.py
index 156ddb3a..88387dcb 100644
--- a/src/aare/gui/widgets/pgroup_dialog.py
+++ b/src/aare/gui/widgets/pgroup_dialog.py
@@ -1,5 +1,8 @@
from PySide6.QtCore import Qt
-from PySide6.QtWidgets import QDialog, QLineEdit, QVBoxLayout, QPushButton, QLabel, QComboBox, QCompleter
+from PySide6.QtWidgets import (
+ QDialog, QVBoxLayout, QPushButton, QLabel,
+ QComboBox, QCompleter, QMessageBox
+)
class PGroupDialog(QDialog):
@@ -8,42 +11,107 @@ class PGroupDialog(QDialog):
self.setWindowTitle("Change current p-group")
self.setMinimumWidth(300)
- # Create a layout
+ self._pgroups = [str(p).strip() for p in (pgroups or []) if p is not None and str(p).strip()]
+
layout = QVBoxLayout(self)
- # Add a label
self.label = QLabel("Set p-group:", self)
layout.addWidget(self.label)
self.combo = QComboBox(self)
self.combo.setEditable(True)
- items = [str(p) for p in (pgroups or []) if p is not None and str(p).strip()]
- self.combo.addItems(items)
+ self.combo.setInsertPolicy(QComboBox.InsertPolicy.NoInsert)
+ self.combo.addItems(self._pgroups)
- completer = QCompleter(items, self)
+ completer = QCompleter(self._pgroups, self)
completer.setCaseSensitivity(Qt.CaseInsensitive)
- completer.setFilterMode(Qt.MatchFlag.MatchContains) # requires Qt import; fallback to default if not desired
+ completer.setFilterMode(Qt.MatchFlag.MatchContains)
+ completer.setCompletionMode(QCompleter.CompletionMode.PopupCompletion)
self.combo.setCompleter(completer)
- if curr_pgroup and curr_pgroup in items:
- self.combo.setCurrentText(curr_pgroup)
- elif curr_pgroup:
+ default_pgroup = self._latest_pgroup(self._pgroups)
+ if curr_pgroup and curr_pgroup in self._pgroups:
self.combo.setCurrentText(curr_pgroup)
+ elif default_pgroup is not None:
+ self.combo.setCurrentText(default_pgroup)
+
layout.addWidget(self.combo)
- # Create buttons
self.ok_button = QPushButton("OK", self)
self.cancel_button = QPushButton("Cancel", self)
- # Add buttons to the layout
layout.addWidget(self.ok_button)
layout.addWidget(self.cancel_button)
- # Connect button signals
- self.ok_button.clicked.connect(self.accept)
+ self.ok_button.clicked.connect(self._validate_and_accept)
self.cancel_button.clicked.connect(self.reject)
+ if self.combo.lineEdit() is not None:
+ self.combo.lineEdit().textEdited.connect(self._live_validate)
+ self._live_validate(self.combo.currentText())
+
+ @staticmethod
+ def _latest_pgroup(pgroups: list[str]) -> str | None:
+ """
+ Return the numerically largest p-group, e.g. p16371 over p01234.
+ Falls back to lexicographic max if parsing fails.
+ """
+ if not pgroups:
+ return None
+
+ def _key(pg: str):
+ s = str(pg).strip()
+ if s.startswith("p") and s[1:].isdigit():
+ return (1, int(s[1:]), s)
+ return (0, -1, s)
+
+ return max(pgroups, key=_key)
+
+ def _set_error_state(self, is_error: bool, message: str | None = None) -> None:
+ if is_error:
+ self.combo.setStyleSheet("border: 2px solid #d9534f;")
+ if message:
+ self.label.setText(f"Set p-group: {message}")
+ else:
+ self.label.setText("Set p-group:")
+ else:
+ self.combo.setStyleSheet("")
+ self.label.setText("Set p-group:")
+
+ def _live_validate(self, text: str) -> None:
+ text = (text or "").strip()
+ if not text:
+ self._set_error_state(True, "Select a p-group")
+ return
+ if self._pgroups and text not in self._pgroups:
+ self._set_error_state(True, "Not in allowed list")
+ return
+ self._set_error_state(False)
+
+ def _validate_and_accept(self) -> None:
+ entered_text = (self.combo.currentText() or "").strip()
+
+ if not entered_text:
+ QMessageBox.warning(self, "Invalid P-Group", "You must select a p-group.")
+ self.combo.setFocus()
+ return
+
+ if self._pgroups and entered_text not in self._pgroups:
+ QMessageBox.warning(
+ self,
+ "Invalid P-Group",
+ f"P-group '{entered_text}' is not in your allowed list.\n"
+ f"Please select from: {', '.join(self._pgroups)}"
+ )
+ self.combo.setFocus()
+ if self.combo.lineEdit() is not None:
+ self.combo.lineEdit().selectAll()
+ self._set_error_state(True, "Not in allowed list")
+ return
+
+ self._set_error_state(False)
+ self.accept()
+
def get_input(self):
"""Return the input text when the dialog is accepted."""
- #return self.text_entry.text()
- return self.combo.currentText()
\ No newline at end of file
+ return (self.combo.currentText() or "").strip()
\ No newline at end of file
diff --git a/src/aare/gui/widgets/status_bar.py b/src/aare/gui/widgets/status_bar.py
index b59ae947..a6b546a9 100644
--- a/src/aare/gui/widgets/status_bar.py
+++ b/src/aare/gui/widgets/status_bar.py
@@ -5,8 +5,10 @@ from PySide6.QtGui import QFont
from PySide6.QtWidgets import QStatusBar, QDialog, QMenu, QMessageBox, QLabel, QSizePolicy
from aare.common.models import TokenData, BeamlineStateEnum, DAQStatusModel, SessionsStateEnum
+from aare.gui.widgets.baton_request_dialog import BatonRequestDialog
from aare.gui.widgets.clickable_label import ClickableLabel
from aare.gui.widgets.pgroup_dialog import PGroupDialog
+from aare.common.auth_models import BatonStatus, BatonRequestStatus
from aare.gui.widgets.value_label import ValueLabel
from aare.common.logger_config import setup_logger
@@ -22,10 +24,18 @@ class StatusBar(QStatusBar):
force_session = Signal()
end_session = Signal()
- close_shutter = Signal()
- open_shutter = Signal()
+ request_baton = Signal()
+ cancel_baton_request = Signal()
+ release_baton = Signal()
+ baton_request_accepted = Signal()
+ baton_request_refused = Signal()
+
get_all_pgroups = Signal()
staff_pgroups_loaded = Signal(list)
+ baton_request_received = Signal(dict)
+
+ close_shutter = Signal()
+ open_shutter = Signal()
def __init__(self, token: TokenData, parent=None):
super().__init__(parent)
@@ -38,6 +48,11 @@ class StatusBar(QStatusBar):
self._message_clear_timer.setSingleShot(True)
self._message_clear_timer.timeout.connect(self.clear_connection_message)
+ self._baton_status: BatonStatus | None = None
+ self._has_pending_request: bool = False
+ self._pgroup_dialog_for_baton: PGroupDialog | None = None
+ self._baton_request_dialog: BatonRequestDialog | None = None
+
self.message_label = QLabel("", self)
self.message_label.setVisible(False)
self.message_label.setSizePolicy(QSizePolicy.Policy.Maximum, QSizePolicy.Policy.Preferred)
@@ -200,26 +215,147 @@ class StatusBar(QStatusBar):
html_content_session = f"""Session: {session_flag}"""
self.session_label.setText(html_content_session)
+ @Slot(BatonStatus)
+ def update_baton_status(self, status: BatonStatus):
+ """Update baton status from SSE stream."""
+ prev_incoming = bool(self._baton_status and self._baton_status.incoming_request)
+
+ # Detect if we just became the holder (e.g., from a queue resolving)
+ was_holder = bool(self._baton_status and self._baton_status.you_are_holder)
+ now_holder = bool(status and status.you_are_holder)
+
+ self._baton_status = status
+ self._has_pending_request = status.you_have_pending_request if status else False
+ self._update_session_display()
+
+ # If we just received the baton (and weren't the holder a moment ago)
+ if now_holder and not was_holder:
+ self._after_baton_granted_select_pgroup()
+
+ incoming = bool(status and status.incoming_request)
+ if incoming and not prev_incoming:
+ self._emit_incoming_baton_request(status)
+
+ if self._baton_request_dialog is not None and self._baton_request_dialog.isVisible():
+ if not incoming:
+ self._baton_request_dialog.close()
+ self._baton_request_dialog = None
+
+ def _emit_incoming_baton_request(self, status: BatonStatus) -> None:
+ requester = "Another user"
+ timeout = 30
+ if status.pending_request is not None:
+ requester = status.pending_request.requester_username or requester
+ timeout = int(status.pending_request.timeout_seconds or timeout)
+
+ self.baton_request_received.emit({
+ "requester": requester,
+ "timeout": timeout,
+ })
+
+ @Slot()
+ def _on_baton_dialog_accepted(self):
+ self.baton_request_accepted.emit()
+ self._baton_request_dialog = None
+
+ @Slot()
+ def _on_baton_dialog_refused(self):
+ self.baton_request_refused.emit()
+ self._baton_request_dialog = None
+
+ def _update_session_display(self):
+ """Update session label based on current status."""
+ if self.__status is None:
+ return
+
+ session_state = self.__status.session.session
+
+ # Base text
+ if session_state == SessionsStateEnum.OwnedByYou:
+ text = "Session: You"
+ if self._baton_status and self._baton_status.incoming_request:
+ text = "Session: You (⚡ Request)"
+ elif session_state == SessionsStateEnum.OwnedByElse:
+ holder_name = ""
+ if self._baton_status and self._baton_status.holder:
+ holder_name = self._baton_status.holder.username
+ text = f"Session: {holder_name or 'Other'}"
+ if self._has_pending_request:
+ text += " (⏳ Waiting)"
+ else:
+ text = "Session: Vacant"
+
+ self.session_label.setText(text)
+
def show_session_menu(self):
menu = QMenu(self)
is_busy = self.__status and self.__status.busy
is_vacant = self.__status and self.__status.session.session == SessionsStateEnum.Vacant
- action_1 = menu.addAction("Grab")
- action_1.setEnabled(bool(not is_busy or self.__is_staff or is_vacant))
- action_1.triggered.connect(self._on_grab_clicked)
- action_2 = menu.addAction("End")
- action_2.setEnabled(bool(not is_busy or self.__is_staff))
- action_2.triggered.connect(self.end_session_clicked)
+ is_yours = self.__status and self.__status.session.session == SessionsStateEnum.OwnedByYou
+ is_other = self.__status and self.__status.session.session == SessionsStateEnum.OwnedByElse
+
+ # Determine holder info from baton status
+ holder_is_staff = (
+ self._baton_status and
+ self._baton_status.holder and
+ self._baton_status.holder.is_staff
+ )
+
+ # --- GRAB / REQUEST ---
+ if is_vacant:
+ # Vacant - simple grab
+ action_grab = menu.addAction("Grab")
+ action_grab.setEnabled(True)
+ action_grab.triggered.connect(self._on_grab_clicked)
+ elif is_other:
+ # Someone else has it
+ if self._has_pending_request:
+ # Already have a pending request - show cancel option
+ action_cancel = menu.addAction("Cancel Request")
+ action_cancel.triggered.connect(self._on_cancel_request_clicked)
+ elif self.__is_staff:
+ # Staff can always grab (override)
+ action_grab = menu.addAction("Grab (Override)")
+ action_grab.setEnabled(not is_busy) # Still respect busy for safety
+ action_grab.triggered.connect(self._on_grab_clicked)
+ elif holder_is_staff:
+ # Non-staff cannot request from staff
+ allowed = self._baton_status and getattr(self._baton_status, "allow_non_staff_request", False)
+ action_grab = menu.addAction("Request from Staff")
+ action_grab.setEnabled(allowed)
+ if allowed:
+ action_grab.triggered.connect(self._on_grab_clicked)
+ else:
+ # Same level - request with timeout
+ action_request = menu.addAction("Request Control")
+ action_request.setEnabled(True)
+ action_request.triggered.connect(self._on_grab_clicked)
+ elif is_yours:
+ # You have it - show release option
+ action_release = menu.addAction("Release")
+ action_release.setEnabled(not is_busy)
+ action_release.triggered.connect(self._on_release_clicked)
+
+ menu.addSeparator()
+
+ # --- END SESSION (cleanup) ---
+ action_end = menu.addAction("End Session")
+ action_end.setEnabled(bool((is_yours and not is_busy) or self.__is_staff))
+ action_end.triggered.connect(self.end_session_clicked)
+
+ # --- STAFF: FORCE GRAB (emergency) ---
+ if self.__is_staff and is_other:
+ menu.addSeparator()
+ action_force = menu.addAction("⚠️ Force Take Over")
+ action_force.triggered.connect(self._on_force_session_clicked)
label_geometry = self.session_label.geometry()
menu.move(self.mapToGlobal(label_geometry.topLeft()) - QPoint(0, menu.sizeHint().height()))
-
menu.setFixedWidth(label_geometry.width())
-
menu.exec()
def show_pgroup_menu(self):
- in_curr = self.__status and self.__status.session.current_pgroup in (self.__allowed_pgroups or [])
+ in_curr = self.__status
logger.info(f"in_curr is {in_curr}")
@@ -301,40 +437,90 @@ class StatusBar(QStatusBar):
menu.exec()
+ def _latest_pgroup(self, pgroups: list[str]) -> str | None:
+ if not pgroups:
+ return None
+
+ def _key(pg: str):
+ s = str(pg).strip()
+ if s.startswith("p") and s[1:].isdigit():
+ return (1, int(s[1:]), s)
+ return (0, -1, s)
+
+ return max(pgroups, key=_key)
+
+ def _after_baton_granted_select_pgroup(self) -> None:
+ """
+ After baton grant:
+ - if exactly one allowed p-group, apply it automatically
+ - otherwise prompt user to choose from their allowed list
+ """
+ pgroups = [str(p).strip() for p in (self.__allowed_pgroups or []) if p is not None and str(p).strip()]
+ if not pgroups:
+ return
+
+ if len(pgroups) == 1:
+ self.set_pgroup.emit(pgroups[0])
+ return
+
+ default_pgroup = self._latest_pgroup(pgroups)
+ curr = None
+ if self.__status and self.__status.session:
+ curr = self.__status.session.current_pgroup or default_pgroup
+
+ self._pgroup_dialog_for_baton = PGroupDialog(
+ curr_pgroup=curr,
+ pgroups=pgroups,
+ parent=self,
+ )
+
+ if self._pgroup_dialog_for_baton.exec() == QDialog.DialogCode.Accepted:
+ selected_pgroup = self._pgroup_dialog_for_baton.get_input()
+ if selected_pgroup:
+ self.set_pgroup.emit(selected_pgroup)
+
+ self._pgroup_dialog_for_baton = None
+
+ def _show_post_grant_pgroup_dialog(self, available_pgroups: list[str]) -> None:
+ curr_pgroup = None
+ if self.__status and self.__status.session:
+ curr_pgroup = self.__status.session.current_pgroup
+
+ if self._pgroup_dialog_for_baton is not None and self._pgroup_dialog_for_baton.isVisible():
+ return
+
+ self._pgroup_dialog_for_baton = PGroupDialog(
+ curr_pgroup=curr_pgroup,
+ pgroups=available_pgroups,
+ parent=self
+ )
+
+ if self._pgroup_dialog_for_baton.exec() == QDialog.DialogCode.Accepted:
+ selected_pgroup = (self._pgroup_dialog_for_baton.get_input() or "").strip()
+ if selected_pgroup:
+ self.set_pgroup.emit(selected_pgroup)
+
+ self._pgroup_dialog_for_baton = None
+
def _on_grab_clicked(self):
- self.grab_session_clicked()
+ """Handle grab/request click - baton first, p-group after grant."""
+ self.request_baton.emit()
- def after_grab():
+ def _on_release_clicked(self):
+ """Handle release click."""
+ self.release_baton.emit()
- if not self.__status:
- logger.info("status is None when session is grabbed")
- return
+ def _on_cancel_request_clicked(self):
+ """Handle cancel request click."""
+ self.cancel_baton_request.emit()
- check_state = self.__status.state in (
- BeamlineStateEnum.SampleAlignment,
- BeamlineStateEnum.SampleExchange,
- BeamlineStateEnum.DewarTransfer,
- BeamlineStateEnum.Maintenance,
- )
-
- session_ownership = self.__status.session.session == SessionsStateEnum.OwnedByYou
- have_pgroup = (self.__status.session.current_pgroup in (self.__allowed_pgroups or []))
- allowed = (not self.__status.busy and check_state and session_ownership and have_pgroup) or self.__is_staff
-
- if allowed:
- self.show_change_dialog()
-
- else:
- logger.debug(
- "not allowed to change pgroup due to; "
- f"session ownership: {session_ownership}, allowed_pgroup: {have_pgroup}, "
- f"beamline busy: {self.__status.busy}, beamline state: {check_state}"
- )
-
- QTimer.singleShot(300, after_grab)
+ def _on_force_session_clicked(self):
+ """Staff emergency force take over (bypasses baton protocol)."""
+ self.force_session.emit()
def grab_session_clicked(self):
- self.force_session.emit()
+ """Legacy method - now routes to baton request."""
+ self.request_baton.emit()
def end_session_clicked(self):
self.end_session.emit()
@@ -359,6 +545,7 @@ class StatusBar(QStatusBar):
return
def _generate_pgroup_dialogue(self, curr: str | None = None, pgroups: list | None = None):
+ logger.info(pgroups)
dialog = PGroupDialog(curr_pgroup=curr, pgroups=pgroups)
if dialog.exec() == QDialog.DialogCode.Accepted:
entered_text = dialog.get_input()
@@ -370,6 +557,7 @@ class StatusBar(QStatusBar):
f"P-group '{entered_text}' is not in your allowed list.\n"
f"Please select from: {', '.join(pgroups)}"
)
+ self._generate_pgroup_dialogue(curr=curr, pgroups=pgroups)
return
self.set_pgroup.emit(entered_text)