From 3dd5f46f1ebcab3c7cec81931ab4bc67356bf8f8 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Wed, 25 Mar 2026 15:41:57 +0100 Subject: [PATCH] Automation 2.0: WIP added endpoints to backend, frontend connected and panel made, now debugging --- src/aare/common/automation_models.py | 194 +++--- src/aare/common/automation_queue_manager.py | 240 +++++++ src/aare/common/automation_workflow.py | 351 ++++++++++ src/aare/daq/auth.py | 8 +- src/aare/daq/automation_api_router.py | 694 ++++++++++++++++++++ src/aare/daq/automation_runner.py | 427 ++++++++++++ src/aare/daq/server.py | 31 + src/aare/gui/main_window.py | 64 ++ src/aare/gui/panels/automation_panel.py | 640 ++++++++++++++++++ src/aare/gui/threads/daq_worker.py | 168 ++++- src/aare/gui/threads/workflow_sse_client.py | 158 +++++ 11 files changed, 2887 insertions(+), 88 deletions(-) create mode 100644 src/aare/common/automation_queue_manager.py create mode 100644 src/aare/common/automation_workflow.py create mode 100644 src/aare/daq/automation_api_router.py create mode 100644 src/aare/daq/automation_runner.py create mode 100644 src/aare/gui/panels/automation_panel.py create mode 100644 src/aare/gui/threads/workflow_sse_client.py diff --git a/src/aare/common/automation_models.py b/src/aare/common/automation_models.py index 67be794b..7431e794 100644 --- a/src/aare/common/automation_models.py +++ b/src/aare/common/automation_models.py @@ -3,7 +3,10 @@ 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" @@ -64,93 +67,116 @@ class WorkflowContext: metadata: dict[str, Any] = field(default_factory=dict) -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=(), - ), -} +class QueueItemStatus(str, Enum): + PENDING = "pending" + RUNNING = "running" + COMPLETED = "completed" + FAILED = "failed" + SKIPPED = "skipped" + ABORTED = "aborted" -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 +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 -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 +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) -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) \ No newline at end of file +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..54e4d2e2 --- /dev/null +++ b/src/aare/common/automation_queue_manager.py @@ -0,0 +1,240 @@ +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, +) + + +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() + + 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..98f26ac1 --- /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(0.5) # 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(0.3) + 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(0.4) + 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(0.5) + 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/daq/auth.py b/src/aare/daq/auth.py index 3af7f247..d93edd06 100644 --- a/src/aare/daq/auth.py +++ b/src/aare/daq/auth.py @@ -5,8 +5,8 @@ from datetime import datetime, timedelta, UTC from typing import List 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.daq.config import BeamlineConfig @@ -25,6 +25,8 @@ SESSION_EXPIRE_SECONDS = 60 * 10 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 +56,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) diff --git a/src/aare/daq/automation_api_router.py b/src/aare/daq/automation_api_router.py new file mode 100644 index 00000000..1236ffc0 --- /dev/null +++ b/src/aare/daq/automation_api_router.py @@ -0,0 +1,694 @@ +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, + ) + + item = redis_mgr.create_item(item) + 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.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) + + control = redis_mgr.request_control( + {"abort_requested": True}, + 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..4626e601 --- /dev/null +++ b/src/aare/daq/automation_runner.py @@ -0,0 +1,427 @@ +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}", + )) + + 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. + """ + + def __init__( + self, + runner: PersistentWorkflowRunner, + redis_manager: WorkflowRedisManager, + poll_interval: float = 0.5, + ): + self._runner = runner + self._redis = redis_manager + self._poll_interval = poll_interval + 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.""" + 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 + + 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: + await self._tick() + 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}", + )) + # Brief pause on error before retrying + await asyncio.sleep(2.0) + + await asyncio.sleep(self._poll_interval) + + async def _tick(self) -> None: + """Single iteration of the automation loop.""" + control = self._runner.check_control() + runtime = self._redis.get_runtime() + + # Handle pause state + 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", + )) + return + + # 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: + return + + # Handle abort + if control.abort_requested and runtime.current_item_id: + self._runner.complete_item(runtime.current_item_id, QueueItemStatus.ABORTED) + self._redis.clear_control() + return + + # 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) + runtime = self._redis.get_runtime() + + # 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) + + # Check for next sample request + if control.next_sample_requested: + self._redis.request_control( + {"next_sample_requested": False}, + requested_by="automation_loop", + ) + return + + # Handle skip request + 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) + result = await asyncio.to_thread(self._runner.run_current_step, context) + + # Notify callback if set + if self._on_step_complete: + step_name = item.steps[item.current_step_index - 1].kind # -1 because index advanced + 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/server.py b/src/aare/daq/server.py index 2665f44d..cf9a02ec 100644 --- a/src/aare/daq/server.py +++ b/src/aare/daq/server.py @@ -34,6 +34,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 +51,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: diff --git a/src/aare/gui/main_window.py b/src/aare/gui/main_window.py index 9d05a250..67b75a75 100644 --- a/src/aare/gui/main_window.py +++ b/src/aare/gui/main_window.py @@ -55,6 +55,9 @@ 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): @@ -288,6 +291,17 @@ 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) @@ -497,6 +511,46 @@ class MainWindow(QMainWindow): 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) + + # 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() @@ -596,6 +650,13 @@ class MainWindow(QMainWindow): ) 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) @@ -770,6 +831,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/automation_panel.py b/src/aare/gui/panels/automation_panel.py new file mode 100644 index 00000000..e53ea753 --- /dev/null +++ b/src/aare/gui/panels/automation_panel.py @@ -0,0 +1,640 @@ +""" +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 (Next/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 _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] + + 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(): + # Check if it's sample data (JSON) + try: + text = mime.text() + if text.startswith("{") or text.startswith("["): + event.acceptProposedAction() + return + except Exception: + pass + # Accept internal moves + 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: + # Try to parse as SampleShortInfoList + 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: + # Try single sample + sample = SampleShortInfo.model_validate_json(text) + self.samples_dropped.emit([sample]) + event.acceptProposedAction() + return + except Exception: + pass + + # Internal reorder + if event.source() == self: + # Get drop position + drop_row = self.indexAt(event.position().toPoint()).row() + if drop_row < 0: + drop_row = self.count() + + # Get selected items + 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: + # Will need to emit a signal for this + pass + 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 + - Full automation mode + """ + + # 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() + request_start_automation = Signal() + request_stop_automation = Signal() + + def __init__(self, parent: QWidget | None = None): + super().__init__(parent) + self._queue_items: list[QueueItem] = [] + self._all_samples: list[SampleShortInfo] = [] # Cache for "Add All" + self._runtime: RuntimeState | None = None + self._control: ControlState | None = None + self._automation_enabled = False + 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) + + # === Control Buttons === + controls_group = QGroupBox("Controls") + controls_layout = QVBoxLayout(controls_group) + + # Mode toggle + mode_layout = QHBoxLayout() + self._mode_label = QLabel("Mode:") + self._guided_btn = QPushButton("Guided Manual") + self._guided_btn.setCheckable(True) + self._guided_btn.setChecked(True) + self._auto_btn = QPushButton("Automation") + self._auto_btn.setCheckable(True) + + self._guided_btn.clicked.connect(self._on_guided_mode) + self._auto_btn.clicked.connect(self._on_automation_mode) + + mode_layout.addWidget(self._mode_label) + mode_layout.addWidget(self._guided_btn) + mode_layout.addWidget(self._auto_btn) + mode_layout.addStretch() + controls_layout.addLayout(mode_layout) + + # Step controls + step_layout = QHBoxLayout() + + self._next_btn = QPushButton("Next Step") + self._next_btn.setStyleSheet("background-color: #4CAF50; color: white;") + self._next_btn.clicked.connect(self.request_next_step.emit) + + self._skip_btn = QPushButton("Skip") + self._skip_btn.setStyleSheet("background-color: #FF9800; color: white;") + self._skip_btn.clicked.connect(self.request_skip.emit) + + 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) + + step_layout.addWidget(self._next_btn) + step_layout.addWidget(self._skip_btn) + step_layout.addWidget(self._pause_btn) + step_layout.addWidget(self._abort_btn) + + controls_layout.addLayout(step_layout) + layout.addWidget(controls_group) + + # === Queue Section === + queue_group = QGroupBox("Queue (drag samples here)") + queue_layout = QVBoxLayout(queue_group) + + # Info label + 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) + + # Draggable queue list + self._queue_list = DraggableQueueListWidget() + self._queue_list.setMinimumHeight(150) + self._queue_list.itemDoubleClicked.connect(self._on_item_double_clicked) + 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) + queue_layout.addWidget(self._queue_list) + + # Queue management buttons + queue_btn_layout = QHBoxLayout() + + self._add_all_btn = QPushButton("➕ Add All Samples") + self._add_all_btn.clicked.connect(self._on_add_all_clicked) + self._add_all_btn.setToolTip("Add all samples from the Sample List to the queue") + + self._remove_selected_btn = QPushButton("🗑 Remove Selected") + self._remove_selected_btn.clicked.connect(self._on_remove_selected) + + self._clear_btn = QPushButton("✖ Clear Queue") + 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_guided_mode(self) -> None: + """Switch to guided manual mode.""" + self._guided_btn.setChecked(True) + self._auto_btn.setChecked(False) + if self._automation_enabled: + self.request_stop_automation.emit() + self._update_button_states() + + def _on_automation_mode(self) -> None: + """Switch to automation mode.""" + self._auto_btn.setChecked(True) + self._guided_btn.setChecked(False) + if not self._automation_enabled: + self.request_start_automation.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_item_double_clicked(self, item: QListWidgetItem) -> None: + """Start processing double-clicked item.""" + idx = self._queue_list.row(item) + if 0 <= idx < len(self._queue_items): + queue_item = self._queue_items[idx] + if queue_item.status == QueueItemStatus.PENDING: + self.request_start_item.emit(queue_item.item_id) + + 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) + # Refresh after a short delay to let server process + 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", + f"Remove all {len(self._queue_items)} items from the queue?", + QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No, + ) + + if reply == QMessageBox.StandardButton.Yes: + self.request_clear_queue.emit() + QTimer.singleShot(500, self.request_queue_refresh.emit) + + 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 + is_automation = self._auto_btn.isChecked() + + # In automation mode, hide the Next button + self._next_btn.setVisible(not is_automation) + self._next_btn.setEnabled(is_running and not is_paused) + + self._skip_btn.setEnabled(is_running) + self._abort_btn.setEnabled(is_running) + + 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) + + # Update drop hint visibility + 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: + self._status_label.setText("▶️ Running") + 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._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._auto_btn.setChecked(enabled) + self._guided_btn.setChecked(not 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"): + self.request_queue_refresh.emit() \ No newline at end of file diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py index 7984bf57..df49f5f7 100644 --- a/src/aare/gui/threads/daq_worker.py +++ b/src/aare/gui/threads/daq_worker.py @@ -51,6 +51,14 @@ 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 + def __init__(self, base_url: str | None, token: str, parent=None): super().__init__(parent) self._active_status_error_key = None @@ -1252,4 +1260,162 @@ 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}") + + # ───────────────────────────────────────────── + # 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 + items = [QueueItem.model_validate(i) for i in data.get("items", [])] + self.workflow_queue_loaded.emit(items) + 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 pending items from the queue.""" + # Load queue first, then delete all pending items + 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_clear_queue_response(reply)) + + 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..f1f4d857 --- /dev/null +++ b/src/aare/gui/threads/workflow_sse_client.py @@ -0,0 +1,158 @@ +""" +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") + # Add token as query param for SSE (can't use headers easily) + url.setQuery(f"token={self._token}") + + 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: + self._reply.abort() + self._reply.deleteLater() + self._reply = None + + self.disconnected.emit() + + @Slot() + def _on_data_ready(self) -> None: + """Handle incoming SSE data.""" + if self._reply is None: + return + + 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) + + 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: + self._reply.deleteLater() + 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 = self._reply.errorString() if self._reply else str(error) + logger.warning(f"Workflow SSE error: {error_msg}") + self.error.emit(error_msg)