Automation 2.0: WIP added endpoints to backend, frontend connected and panel made, now debugging

This commit is contained in:
2026-03-25 15:41:57 +01:00
parent c9ec2e5028
commit 3dd5f46f1e
11 changed files with 2887 additions and 88 deletions
+110 -84
View File
@@ -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)
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
+240
View File
@@ -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
+351
View File
@@ -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),
}
+5 -3
View File
@@ -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)
+694
View File
@@ -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",
},
)
+427
View File
@@ -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)
+31
View File
@@ -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:
+64
View File
@@ -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",
+640
View File
@@ -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()
+167 -1
View File
@@ -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}")
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")
+158
View File
@@ -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)