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)