DAQ/GUI: removed automation workflow manager runner, workflow, models and tidied up GUI/DAQ as it is not in use at the moment.
This commit is contained in:
@@ -2,16 +2,6 @@ from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
import time
|
||||
import uuid
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
class WorkflowMode(str, Enum):
|
||||
FLEXIBLE_MANUAL = "flexible_manual"
|
||||
GUIDED_MANUAL = "guided_manual"
|
||||
AUTOMATION = "automation"
|
||||
|
||||
|
||||
class WorkflowStateKind(str, Enum):
|
||||
@@ -36,29 +26,9 @@ class StepState:
|
||||
step: WorkflowStateKind
|
||||
status: StepStatus = StepStatus.PENDING
|
||||
message: str = ""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TransitionRule:
|
||||
to_state: WorkflowStateKind
|
||||
allowed_modes: frozenset[WorkflowMode] = frozenset()
|
||||
optional: bool = False
|
||||
condition_name: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class StateDefinition:
|
||||
kind: WorkflowStateKind
|
||||
transitions: tuple[TransitionRule, ...]
|
||||
description: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class StateResult:
|
||||
state: WorkflowStateKind
|
||||
status: StepStatus
|
||||
message: str = ""
|
||||
payload: dict[str, Any] = field(default_factory=dict)
|
||||
started_at: float | None = None
|
||||
completed_at: float | None = None
|
||||
error_code: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -67,132 +37,6 @@ class AutomationProgress:
|
||||
steps: list[StepState] = field(default_factory=list)
|
||||
finished: bool = False
|
||||
success: bool | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class WorkflowContext:
|
||||
mode: WorkflowMode
|
||||
queue_id: str
|
||||
item_id: str
|
||||
sample_id: int | None = None
|
||||
current_state: WorkflowStateKind | None = None
|
||||
current_step_index: int = 0
|
||||
paused: bool = False
|
||||
abort_requested: bool = False
|
||||
last_message: str = ""
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
class QueueItemStatus(str, Enum):
|
||||
PENDING = "pending"
|
||||
RUNNING = "running"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
SKIPPED = "skipped"
|
||||
ABORTED = "aborted"
|
||||
|
||||
|
||||
class WorkflowStepRecord(BaseModel):
|
||||
kind: str
|
||||
status: str = "pending"
|
||||
message: str = ""
|
||||
started_at: float | None = None
|
||||
completed_at: float | None = None
|
||||
error_detail: str | None = None
|
||||
|
||||
|
||||
class QueueItem(BaseModel):
|
||||
item_id: str
|
||||
beamline: str
|
||||
sample_id: int | None = None
|
||||
sample_name: str = ""
|
||||
owner_pgroup: str = ""
|
||||
created_by: str = ""
|
||||
created_at: float = Field(default_factory=time.time)
|
||||
priority: int = 100
|
||||
order_index: int = 0
|
||||
status: QueueItemStatus = QueueItemStatus.PENDING
|
||||
steps: list[WorkflowStepRecord] = Field(default_factory=list)
|
||||
current_step_index: int = 0
|
||||
recipe: dict[str, Any] = Field(default_factory=dict)
|
||||
metadata: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class RuntimeState(BaseModel):
|
||||
running: bool = False
|
||||
paused: bool = False
|
||||
current_queue_id: str = ""
|
||||
current_item_id: str | None = None
|
||||
current_state: str | None = None
|
||||
current_step_index: int = 0
|
||||
last_error: str | None = None
|
||||
last_update: float = Field(default_factory=time.time)
|
||||
|
||||
|
||||
class ControlState(BaseModel):
|
||||
pause_requested: bool = False
|
||||
resume_requested: bool = False
|
||||
abort_requested: bool = False
|
||||
skip_requested: bool = False
|
||||
next_sample_requested: bool = False
|
||||
requested_by: str | None = None
|
||||
requested_at: float | None = None
|
||||
|
||||
|
||||
class WorkflowEvent(BaseModel):
|
||||
event_id: str = Field(default_factory=lambda: str(uuid.uuid4()))
|
||||
beamline: str = ""
|
||||
item_id: str | None = None
|
||||
step: str | None = None
|
||||
event_type: str = ""
|
||||
timestamp: float = Field(default_factory=time.time)
|
||||
actor: str = ""
|
||||
message: str = ""
|
||||
payload: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class CreateQueueItemRequest(BaseModel):
|
||||
sample_id: int | None = None
|
||||
sample_name: str = ""
|
||||
priority: int = 100
|
||||
recipe: dict = Field(default_factory=dict)
|
||||
steps: list[str] | None = None # If None, use default steps
|
||||
|
||||
|
||||
class MoveItemRequest(BaseModel):
|
||||
new_order_index: int
|
||||
|
||||
|
||||
class QueueListResponse(BaseModel):
|
||||
items: list[QueueItem]
|
||||
total: int
|
||||
|
||||
|
||||
class RuntimeResponse(BaseModel):
|
||||
runtime: RuntimeState
|
||||
control: ControlState
|
||||
|
||||
|
||||
class ControlActionResponse(BaseModel):
|
||||
ok: bool
|
||||
control: ControlState
|
||||
message: str = ""
|
||||
|
||||
|
||||
class StepActionResponse(BaseModel):
|
||||
ok: bool
|
||||
item: QueueItem | None = None
|
||||
step: str | None = None
|
||||
status: str = ""
|
||||
message: str = ""
|
||||
|
||||
|
||||
class EventListResponse(BaseModel):
|
||||
events: list[WorkflowEvent]
|
||||
|
||||
|
||||
class AutomationStatusResponse(BaseModel):
|
||||
enabled: bool
|
||||
running: bool
|
||||
runtime: RuntimeState
|
||||
control: ControlState
|
||||
samples_in_queue: int = 0
|
||||
avg_time_per_sample: float = 0.0
|
||||
current_sample_name: str = ""
|
||||
@@ -1,245 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
import redis
|
||||
|
||||
from aare.common.automation_models import (
|
||||
QueueItem,
|
||||
WorkflowEvent,
|
||||
ControlState,
|
||||
RuntimeState,
|
||||
QueueItemStatus,
|
||||
WorkflowStepRecord,
|
||||
WorkflowStateKind,
|
||||
)
|
||||
|
||||
CURRENT_QUEUE_LIMIT = 576
|
||||
|
||||
def build_default_steps() -> list[WorkflowStepRecord]:
|
||||
return [
|
||||
WorkflowStepRecord(kind=WorkflowStateKind.MOUNT.value),
|
||||
WorkflowStepRecord(kind=WorkflowStateKind.LOOP_CENTRE.value),
|
||||
WorkflowStepRecord(kind=WorkflowStateKind.RASTER.value),
|
||||
WorkflowStepRecord(kind=WorkflowStateKind.DATA_COLLECTION.value),
|
||||
]
|
||||
|
||||
|
||||
class WorkflowRedisManager:
|
||||
def __init__(self, client: redis.Redis, beamline: str):
|
||||
self._client = client
|
||||
self._bl = beamline.lower()
|
||||
|
||||
def _key(self, suffix: str) -> str:
|
||||
return f"{self._bl}:workflow:{suffix}"
|
||||
|
||||
def _item_key(self, item_id: str) -> str:
|
||||
return self._key(f"item:{item_id}")
|
||||
|
||||
def _queue_key(self) -> str:
|
||||
return self._key("queue")
|
||||
|
||||
def _runtime_key(self) -> str:
|
||||
return self._key("runtime")
|
||||
|
||||
def _control_key(self) -> str:
|
||||
return self._key("control")
|
||||
|
||||
def _events_key(self) -> str:
|
||||
return self._key("events")
|
||||
|
||||
def _next_item_id(self) -> str:
|
||||
n = int(self._client.incr(self._key("item_seq")))
|
||||
return f"wf_{int(time.time())}_{n:06d}"
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Queue item operations
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
def create_item(self, item: QueueItem) -> QueueItem:
|
||||
if not item.item_id:
|
||||
item.item_id = self._next_item_id()
|
||||
|
||||
current_count = self._client.zcard(self._queue_key())
|
||||
if current_count >= CURRENT_QUEUE_LIMIT:
|
||||
raise ValueError(f"Queue size limit reached ({CURRENT_QUEUE_LIMIT} items). Please clear the queue.")
|
||||
|
||||
pipe = self._client.pipeline(transaction=True)
|
||||
pipe.set(self._item_key(item.item_id), item.model_dump_json())
|
||||
pipe.zadd(self._queue_key(), {item.item_id: float(item.order_index)})
|
||||
pipe.execute()
|
||||
|
||||
self.append_event(WorkflowEvent(
|
||||
beamline=self._bl,
|
||||
item_id=item.item_id,
|
||||
event_type="item_created",
|
||||
message="Queue item created",
|
||||
payload={"status": item.status.value},
|
||||
))
|
||||
|
||||
return item
|
||||
|
||||
def get_item(self, item_id: str) -> QueueItem | None:
|
||||
raw = self._client.get(self._item_key(item_id))
|
||||
if raw is None:
|
||||
return None
|
||||
return QueueItem.model_validate_json(raw)
|
||||
|
||||
def update_item(self, item_id: str, patch: dict[str, Any]) -> QueueItem:
|
||||
item = self.get_item(item_id)
|
||||
if item is None:
|
||||
raise KeyError(f"Queue item not found: {item_id}")
|
||||
|
||||
updated = item.model_copy(update=patch)
|
||||
self._client.set(self._item_key(item_id), updated.model_dump_json())
|
||||
return updated
|
||||
|
||||
def delete_item(self, item_id: str) -> None:
|
||||
pipe = self._client.pipeline(transaction=True)
|
||||
pipe.delete(self._item_key(item_id))
|
||||
pipe.zrem(self._queue_key(), item_id)
|
||||
pipe.execute()
|
||||
|
||||
def list_queue_order(self) -> list[str]:
|
||||
return [str(x) for x in self._client.zrange(self._queue_key(), 0, -1)]
|
||||
|
||||
def list_items(self, *, include_finished: bool = True) -> list[QueueItem]:
|
||||
out: list[QueueItem] = []
|
||||
for item_id in self.list_queue_order():
|
||||
item = self.get_item(item_id)
|
||||
if item is None:
|
||||
continue
|
||||
if not include_finished and item.status in {
|
||||
QueueItemStatus.COMPLETED,
|
||||
QueueItemStatus.FAILED,
|
||||
QueueItemStatus.SKIPPED,
|
||||
QueueItemStatus.ABORTED,
|
||||
}:
|
||||
continue
|
||||
out.append(item)
|
||||
return out
|
||||
|
||||
def get_next_pending_item(self) -> QueueItem | None:
|
||||
for item_id in self.list_queue_order():
|
||||
item = self.get_item(item_id)
|
||||
if item is not None and item.status == QueueItemStatus.PENDING:
|
||||
return item
|
||||
return None
|
||||
|
||||
def move_item(self, item_id: str, new_order_index: int) -> None:
|
||||
if self.get_item(item_id) is None:
|
||||
raise KeyError(f"Queue item not found: {item_id}")
|
||||
self._client.zadd(self._queue_key(), {item_id: float(new_order_index)})
|
||||
self.update_item(item_id, {"order_index": new_order_index})
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Runtime state
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
def get_runtime(self) -> RuntimeState:
|
||||
raw = self._client.get(self._runtime_key())
|
||||
if raw is None:
|
||||
return RuntimeState()
|
||||
return RuntimeState.model_validate_json(raw)
|
||||
|
||||
def set_runtime(self, runtime: RuntimeState) -> RuntimeState:
|
||||
runtime.last_update = time.time()
|
||||
self._client.set(self._runtime_key(), runtime.model_dump_json())
|
||||
return runtime
|
||||
|
||||
def patch_runtime(self, patch: dict[str, Any]) -> RuntimeState:
|
||||
runtime = self.get_runtime()
|
||||
updated = runtime.model_copy(update=patch)
|
||||
return self.set_runtime(updated)
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Control state
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
def get_control(self) -> ControlState:
|
||||
raw = self._client.get(self._control_key())
|
||||
if raw is None:
|
||||
return ControlState()
|
||||
return ControlState.model_validate_json(raw)
|
||||
|
||||
def request_control(self, patch: dict[str, Any], *, requested_by: str) -> ControlState:
|
||||
control = self.get_control()
|
||||
updated = control.model_copy(update={
|
||||
**patch,
|
||||
"requested_by": requested_by,
|
||||
"requested_at": time.time(),
|
||||
})
|
||||
self._client.set(self._control_key(), updated.model_dump_json())
|
||||
|
||||
self.append_event(WorkflowEvent(
|
||||
beamline=self._bl,
|
||||
item_id=self.get_runtime().current_item_id,
|
||||
event_type="control_requested",
|
||||
actor=requested_by,
|
||||
payload=patch,
|
||||
))
|
||||
|
||||
return updated
|
||||
|
||||
def clear_control(self) -> ControlState:
|
||||
control = ControlState()
|
||||
self._client.set(self._control_key(), control.model_dump_json())
|
||||
return control
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Event log
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
def append_event(self, event: WorkflowEvent) -> None:
|
||||
payload = event.model_dump_json()
|
||||
self._client.xadd(self._events_key(), {"json": payload}, maxlen=5000, approximate=True)
|
||||
|
||||
def read_events(self, *, limit: int = 200) -> list[WorkflowEvent]:
|
||||
rows = self._client.xrevrange(self._events_key(), count=limit)
|
||||
out: list[WorkflowEvent] = []
|
||||
for _, fields in reversed(rows):
|
||||
raw = fields.get("json")
|
||||
if raw:
|
||||
out.append(WorkflowEvent.model_validate_json(raw))
|
||||
return out
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Step status updates
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
def update_step(
|
||||
self,
|
||||
item_id: str,
|
||||
step_index: int,
|
||||
*,
|
||||
status: str | None = None,
|
||||
message: str | None = None,
|
||||
error_detail: str | None = None,
|
||||
) -> QueueItem:
|
||||
item = self.get_item(item_id)
|
||||
if item is None:
|
||||
raise KeyError(f"Queue item not found: {item_id}")
|
||||
|
||||
if not (0 <= step_index < len(item.steps)):
|
||||
raise IndexError(f"Step index out of range: {step_index}")
|
||||
|
||||
step = item.steps[step_index]
|
||||
|
||||
if status is not None:
|
||||
step.status = status
|
||||
if status == "running" and step.started_at is None:
|
||||
step.started_at = time.time()
|
||||
elif status in ("completed", "failed", "skipped"):
|
||||
step.completed_at = time.time()
|
||||
|
||||
if message is not None:
|
||||
step.message = message
|
||||
|
||||
if error_detail is not None:
|
||||
step.error_detail = error_detail
|
||||
|
||||
item.steps[step_index] = step
|
||||
self._client.set(self._item_key(item_id), item.model_dump_json())
|
||||
|
||||
return item
|
||||
@@ -1,351 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from aare.common.automation_models import (
|
||||
WorkflowStateKind,
|
||||
WorkflowMode,
|
||||
StepStatus,
|
||||
StateResult,
|
||||
WorkflowContext,
|
||||
TransitionRule,
|
||||
StateDefinition,
|
||||
)
|
||||
|
||||
|
||||
STATE_REGISTRY: dict[WorkflowStateKind, StateDefinition] = {
|
||||
WorkflowStateKind.MOUNT: StateDefinition(
|
||||
kind=WorkflowStateKind.MOUNT,
|
||||
description="Mount the sample",
|
||||
transitions=(
|
||||
TransitionRule(
|
||||
to_state=WorkflowStateKind.LOOP_CENTRE,
|
||||
allowed_modes=frozenset({
|
||||
WorkflowMode.FLEXIBLE_MANUAL,
|
||||
WorkflowMode.GUIDED_MANUAL,
|
||||
WorkflowMode.AUTOMATION,
|
||||
}),
|
||||
),
|
||||
),
|
||||
),
|
||||
WorkflowStateKind.LOOP_CENTRE: StateDefinition(
|
||||
kind=WorkflowStateKind.LOOP_CENTRE,
|
||||
description="Centre the loop",
|
||||
transitions=(
|
||||
TransitionRule(
|
||||
to_state=WorkflowStateKind.RASTER,
|
||||
allowed_modes=frozenset({
|
||||
WorkflowMode.FLEXIBLE_MANUAL,
|
||||
WorkflowMode.GUIDED_MANUAL,
|
||||
WorkflowMode.AUTOMATION,
|
||||
}),
|
||||
),
|
||||
TransitionRule(
|
||||
to_state=WorkflowStateKind.DATA_COLLECTION,
|
||||
allowed_modes=frozenset({
|
||||
WorkflowMode.FLEXIBLE_MANUAL,
|
||||
WorkflowMode.GUIDED_MANUAL,
|
||||
}),
|
||||
optional=True,
|
||||
),
|
||||
),
|
||||
),
|
||||
WorkflowStateKind.RASTER: StateDefinition(
|
||||
kind=WorkflowStateKind.RASTER,
|
||||
description="Run raster scan",
|
||||
transitions=(
|
||||
TransitionRule(
|
||||
to_state=WorkflowStateKind.DATA_COLLECTION,
|
||||
allowed_modes=frozenset({
|
||||
WorkflowMode.FLEXIBLE_MANUAL,
|
||||
WorkflowMode.GUIDED_MANUAL,
|
||||
WorkflowMode.AUTOMATION,
|
||||
}),
|
||||
),
|
||||
),
|
||||
),
|
||||
WorkflowStateKind.DATA_COLLECTION: StateDefinition(
|
||||
kind=WorkflowStateKind.DATA_COLLECTION,
|
||||
description="Collect diffraction data",
|
||||
transitions=(),
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def get_state_definition(kind: WorkflowStateKind) -> StateDefinition:
|
||||
try:
|
||||
return STATE_REGISTRY[kind]
|
||||
except KeyError as exc:
|
||||
raise KeyError(f"Unknown workflow state: {kind}") from exc
|
||||
|
||||
|
||||
def get_allowed_next_states(
|
||||
kind: WorkflowStateKind,
|
||||
mode: WorkflowMode | None = None,
|
||||
) -> list[WorkflowStateKind]:
|
||||
definition = get_state_definition(kind)
|
||||
out: list[WorkflowStateKind] = []
|
||||
|
||||
for transition in definition.transitions:
|
||||
if mode is None:
|
||||
out.append(transition.to_state)
|
||||
continue
|
||||
|
||||
if not transition.allowed_modes or mode in transition.allowed_modes:
|
||||
out.append(transition.to_state)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
def can_transition(
|
||||
from_state: WorkflowStateKind,
|
||||
to_state: WorkflowStateKind,
|
||||
mode: WorkflowMode | None = None,
|
||||
) -> bool:
|
||||
return to_state in get_allowed_next_states(from_state, mode=mode)
|
||||
|
||||
|
||||
class StateHandler(ABC):
|
||||
state_kind: WorkflowStateKind
|
||||
|
||||
def __init__(self, registry: dict[WorkflowStateKind, StateDefinition] | None = None):
|
||||
self._registry = registry or STATE_REGISTRY
|
||||
|
||||
def definition(self) -> StateDefinition:
|
||||
return get_state_definition(self.state_kind)
|
||||
|
||||
def can_run(self, context: WorkflowContext) -> bool:
|
||||
return context.current_state in (None, self.state_kind)
|
||||
|
||||
@abstractmethod
|
||||
def validate(self, context: WorkflowContext) -> None:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def execute(self, context: WorkflowContext) -> StateResult:
|
||||
pass
|
||||
|
||||
|
||||
class MountHandler(StateHandler):
|
||||
state_kind = WorkflowStateKind.MOUNT
|
||||
|
||||
def validate(self, context: WorkflowContext) -> None:
|
||||
if context.abort_requested:
|
||||
raise RuntimeError("Abort requested; cannot mount.")
|
||||
|
||||
def execute(self, context: WorkflowContext) -> StateResult:
|
||||
self.validate(context)
|
||||
context.current_state = WorkflowStateKind.MOUNT
|
||||
context.last_message = "Sample mounted"
|
||||
return StateResult(
|
||||
state=self.state_kind,
|
||||
status=StepStatus.SUCCESS,
|
||||
message="Sample mounted successfully.",
|
||||
payload={"mounted": True},
|
||||
)
|
||||
|
||||
|
||||
class LoopCentreHandler(StateHandler):
|
||||
state_kind = WorkflowStateKind.LOOP_CENTRE
|
||||
|
||||
def validate(self, context: WorkflowContext) -> None:
|
||||
if context.abort_requested:
|
||||
raise RuntimeError("Abort requested; cannot loop-centre.")
|
||||
|
||||
def execute(self, context: WorkflowContext) -> StateResult:
|
||||
self.validate(context)
|
||||
context.current_state = WorkflowStateKind.LOOP_CENTRE
|
||||
context.last_message = "Loop centred"
|
||||
return StateResult(
|
||||
state=self.state_kind,
|
||||
status=StepStatus.SUCCESS,
|
||||
message="Loop centring completed.",
|
||||
payload={"centred": True},
|
||||
)
|
||||
|
||||
|
||||
class RasterHandler(StateHandler):
|
||||
state_kind = WorkflowStateKind.RASTER
|
||||
|
||||
def validate(self, context: WorkflowContext) -> None:
|
||||
if context.abort_requested:
|
||||
raise RuntimeError("Abort requested; cannot raster.")
|
||||
|
||||
def execute(self, context: WorkflowContext) -> StateResult:
|
||||
self.validate(context)
|
||||
context.current_state = WorkflowStateKind.RASTER
|
||||
context.last_message = "Raster completed"
|
||||
return StateResult(
|
||||
state=self.state_kind,
|
||||
status=StepStatus.SUCCESS,
|
||||
message="Raster scan completed.",
|
||||
payload={"best_spot_found": True},
|
||||
)
|
||||
|
||||
|
||||
class DataCollectionHandler(StateHandler):
|
||||
state_kind = WorkflowStateKind.DATA_COLLECTION
|
||||
|
||||
def validate(self, context: WorkflowContext) -> None:
|
||||
if context.abort_requested:
|
||||
raise RuntimeError("Abort requested; cannot collect data.")
|
||||
|
||||
def execute(self, context: WorkflowContext) -> StateResult:
|
||||
self.validate(context)
|
||||
context.current_state = WorkflowStateKind.DATA_COLLECTION
|
||||
context.last_message = "Data collected"
|
||||
return StateResult(
|
||||
state=self.state_kind,
|
||||
status=StepStatus.SUCCESS,
|
||||
message="Data collection completed.",
|
||||
payload={"frames_collected": 1},
|
||||
)
|
||||
|
||||
|
||||
class WorkflowRunner:
|
||||
"""Simple in-memory runner (no persistence)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
registry: dict[WorkflowStateKind, StateDefinition] | None = None,
|
||||
handlers: dict[WorkflowStateKind, StateHandler] | None = None,
|
||||
):
|
||||
self._registry = registry or STATE_REGISTRY
|
||||
self._handlers = handlers or HANDLER_REGISTRY
|
||||
|
||||
def get_handler(self, state: WorkflowStateKind) -> StateHandler:
|
||||
try:
|
||||
return self._handlers[state]
|
||||
except KeyError as exc:
|
||||
raise KeyError(f"No handler registered for state: {state}") from exc
|
||||
|
||||
def can_move_to(
|
||||
self,
|
||||
current: WorkflowStateKind,
|
||||
next_state: WorkflowStateKind,
|
||||
mode: WorkflowMode,
|
||||
) -> bool:
|
||||
return can_transition(current, next_state, mode)
|
||||
|
||||
def run_state(
|
||||
self,
|
||||
context: WorkflowContext,
|
||||
state: WorkflowStateKind,
|
||||
) -> StateResult:
|
||||
if context.current_state is not None:
|
||||
if not self.can_move_to(context.current_state, state, context.mode):
|
||||
raise RuntimeError(
|
||||
f"Transition not allowed: {context.current_state} -> {state}"
|
||||
)
|
||||
|
||||
handler = self.get_handler(state)
|
||||
result = handler.execute(context)
|
||||
|
||||
context.current_state = state
|
||||
context.current_step_index += 1
|
||||
context.last_message = result.message
|
||||
|
||||
return result
|
||||
|
||||
|
||||
class SimulatedMountHandler(StateHandler):
|
||||
"""Simulated mount handler for testing - doesn't actually mount."""
|
||||
state_kind = WorkflowStateKind.MOUNT
|
||||
|
||||
def validate(self, context: WorkflowContext) -> None:
|
||||
if context.abort_requested:
|
||||
raise RuntimeError("Abort requested; cannot mount.")
|
||||
|
||||
def execute(self, context: WorkflowContext) -> StateResult:
|
||||
self.validate(context)
|
||||
import time
|
||||
time.sleep(10) # Simulate some work
|
||||
context.current_state = WorkflowStateKind.MOUNT
|
||||
context.last_message = "[SIMULATION] Sample would be mounted"
|
||||
return StateResult(
|
||||
state=self.state_kind,
|
||||
status=StepStatus.SUCCESS,
|
||||
message=f"[SIMULATION] Would mount sample_id={context.sample_id}",
|
||||
payload={"mounted": True, "simulated": True},
|
||||
)
|
||||
|
||||
|
||||
class SimulatedLoopCentreHandler(StateHandler):
|
||||
"""Simulated loop centre handler for testing."""
|
||||
state_kind = WorkflowStateKind.LOOP_CENTRE
|
||||
|
||||
def validate(self, context: WorkflowContext) -> None:
|
||||
if context.abort_requested:
|
||||
raise RuntimeError("Abort requested; cannot loop-centre.")
|
||||
|
||||
def execute(self, context: WorkflowContext) -> StateResult:
|
||||
self.validate(context)
|
||||
import time
|
||||
time.sleep(10)
|
||||
context.current_state = WorkflowStateKind.LOOP_CENTRE
|
||||
context.last_message = "[SIMULATION] Loop would be centred"
|
||||
return StateResult(
|
||||
state=self.state_kind,
|
||||
status=StepStatus.SUCCESS,
|
||||
message="[SIMULATION] Would run loop centering algorithm",
|
||||
payload={"centred": True, "simulated": True},
|
||||
)
|
||||
|
||||
|
||||
class SimulatedRasterHandler(StateHandler):
|
||||
"""Simulated raster handler for testing."""
|
||||
state_kind = WorkflowStateKind.RASTER
|
||||
|
||||
def validate(self, context: WorkflowContext) -> None:
|
||||
if context.abort_requested:
|
||||
raise RuntimeError("Abort requested; cannot raster.")
|
||||
|
||||
def execute(self, context: WorkflowContext) -> StateResult:
|
||||
self.validate(context)
|
||||
import time
|
||||
time.sleep(10)
|
||||
context.current_state = WorkflowStateKind.RASTER
|
||||
context.last_message = "[SIMULATION] Raster scan would be performed"
|
||||
return StateResult(
|
||||
state=self.state_kind,
|
||||
status=StepStatus.SUCCESS,
|
||||
message="[SIMULATION] Would run raster scan, find best diffraction spot",
|
||||
payload={"best_spot_found": True, "simulated": True},
|
||||
)
|
||||
|
||||
|
||||
class SimulatedDataCollectionHandler(StateHandler):
|
||||
"""Simulated data collection handler for testing."""
|
||||
state_kind = WorkflowStateKind.DATA_COLLECTION
|
||||
|
||||
def validate(self, context: WorkflowContext) -> None:
|
||||
if context.abort_requested:
|
||||
raise RuntimeError("Abort requested; cannot collect data.")
|
||||
|
||||
def execute(self, context: WorkflowContext) -> StateResult:
|
||||
self.validate(context)
|
||||
import time
|
||||
time.sleep(10)
|
||||
context.current_state = WorkflowStateKind.DATA_COLLECTION
|
||||
context.last_message = "[SIMULATION] Data collection would be performed"
|
||||
return StateResult(
|
||||
state=self.state_kind,
|
||||
status=StepStatus.SUCCESS,
|
||||
message="[SIMULATION] Would collect 1800 frames at 0.2° oscillation",
|
||||
payload={"frames_collected": 1800, "simulated": True},
|
||||
)
|
||||
|
||||
|
||||
# Simulated handler registry for testing
|
||||
SIMULATED_HANDLER_REGISTRY: dict[WorkflowStateKind, StateHandler] = {
|
||||
WorkflowStateKind.MOUNT: SimulatedMountHandler(STATE_REGISTRY),
|
||||
WorkflowStateKind.LOOP_CENTRE: SimulatedLoopCentreHandler(STATE_REGISTRY),
|
||||
WorkflowStateKind.RASTER: SimulatedRasterHandler(STATE_REGISTRY),
|
||||
WorkflowStateKind.DATA_COLLECTION: SimulatedDataCollectionHandler(STATE_REGISTRY),
|
||||
}
|
||||
|
||||
HANDLER_REGISTRY: dict[WorkflowStateKind, StateHandler] = {
|
||||
WorkflowStateKind.MOUNT: MountHandler(STATE_REGISTRY),
|
||||
WorkflowStateKind.LOOP_CENTRE: LoopCentreHandler(STATE_REGISTRY),
|
||||
WorkflowStateKind.RASTER: RasterHandler(STATE_REGISTRY),
|
||||
WorkflowStateKind.DATA_COLLECTION: DataCollectionHandler(STATE_REGISTRY),
|
||||
}
|
||||
@@ -1,758 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from typing import AsyncGenerator
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
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.daq.automation_runner import PersistentWorkflowRunner
|
||||
from aare.daq.auth import parse_token, check_jwt_rw, check_jwt_ro, oauth2_scheme
|
||||
from aare.daq.config import BeamlineConfig
|
||||
|
||||
from aare.daq.automation_runner import AutomationLoop
|
||||
|
||||
router = APIRouter(prefix="/workflow", tags=["workflow"])
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Dependency: get managers
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
# These will be set up when the router is included
|
||||
_redis_manager: WorkflowRedisManager | None = None
|
||||
_runner: PersistentWorkflowRunner | None = None
|
||||
_cfg: BeamlineConfig | None = None
|
||||
_automation_loop: AutomationLoop | None = None
|
||||
|
||||
|
||||
def set_workflow_dependencies(
|
||||
redis_manager: WorkflowRedisManager,
|
||||
runner: PersistentWorkflowRunner,
|
||||
cfg: BeamlineConfig,
|
||||
) -> None:
|
||||
global _redis_manager, _runner, _cfg, _automation_loop
|
||||
_redis_manager = redis_manager
|
||||
_runner = runner
|
||||
_cfg = cfg
|
||||
_automation_loop = AutomationLoop(runner, redis_manager)
|
||||
|
||||
|
||||
def get_automation_loop() -> AutomationLoop:
|
||||
if _automation_loop is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="Automation loop not initialized",
|
||||
)
|
||||
return _automation_loop
|
||||
|
||||
|
||||
def get_cfg() -> BeamlineConfig:
|
||||
if _cfg is None:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail="Workflow system not initialized",
|
||||
)
|
||||
return _cfg
|
||||
|
||||
|
||||
def get_redis_manager() -> WorkflowRedisManager:
|
||||
if _redis_manager is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="Workflow system not initialized",
|
||||
)
|
||||
return _redis_manager
|
||||
|
||||
|
||||
def get_runner() -> PersistentWorkflowRunner:
|
||||
if _runner is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="Workflow runner not initialized",
|
||||
)
|
||||
return _runner
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Queue management endpoints
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.get("/queue", response_model=QueueListResponse)
|
||||
async def list_queue(
|
||||
include_finished: bool = False,
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""List all items in the workflow queue."""
|
||||
data = parse_token(token)
|
||||
check_jwt_ro(get_cfg(), data)
|
||||
items = redis_mgr.list_items(include_finished=include_finished)
|
||||
|
||||
# Filter by pgroup if not staff
|
||||
if not data.staff:
|
||||
items = [i for i in items if i.owner_pgroup in data.pgroups]
|
||||
|
||||
return QueueListResponse(items=items, total=len(items))
|
||||
|
||||
|
||||
@router.post("/queue", response_model=QueueItem)
|
||||
async def create_queue_item(
|
||||
request: CreateQueueItemRequest,
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Add a new item to the workflow queue."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
# Build steps
|
||||
if request.steps:
|
||||
steps = [
|
||||
WorkflowStepRecord(kind=s)
|
||||
for s in request.steps
|
||||
if s in [sk.value for sk in WorkflowStateKind]
|
||||
]
|
||||
else:
|
||||
steps = build_default_steps()
|
||||
|
||||
item = QueueItem(
|
||||
item_id="",
|
||||
beamline=redis_mgr._bl,
|
||||
sample_id=request.sample_id,
|
||||
sample_name=request.sample_name,
|
||||
owner_pgroup=data.pgroups[0] if data.pgroups else "",
|
||||
created_by=data.sub,
|
||||
priority=request.priority,
|
||||
order_index=int(asyncio.get_event_loop().time() * 1000),
|
||||
steps=steps,
|
||||
recipe=request.recipe,
|
||||
)
|
||||
|
||||
try:
|
||||
item = redis_mgr.create_item(item)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
|
||||
return item
|
||||
|
||||
|
||||
@router.get("/queue/{item_id}", response_model=QueueItem)
|
||||
async def get_queue_item(
|
||||
item_id: str,
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Get a single queue item by ID."""
|
||||
data = parse_token(token)
|
||||
check_jwt_ro(get_cfg(), data)
|
||||
|
||||
item = redis_mgr.get_item(item_id)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=404, detail="Item not found")
|
||||
|
||||
# Check access
|
||||
if not data.staff and item.owner_pgroup not in data.pgroups:
|
||||
raise HTTPException(status_code=403, detail="Access denied")
|
||||
|
||||
return item
|
||||
|
||||
|
||||
@router.delete("/queue/{item_id}")
|
||||
async def delete_queue_item(
|
||||
item_id: str,
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Remove an item from the queue."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
item = redis_mgr.get_item(item_id)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=404, detail="Item not found")
|
||||
|
||||
# Check access
|
||||
if not data.staff and item.owner_pgroup not in data.pgroups:
|
||||
raise HTTPException(status_code=403, detail="Access denied")
|
||||
|
||||
# Don't allow deleting running items
|
||||
if item.status == QueueItemStatus.RUNNING:
|
||||
raise HTTPException(status_code=409, detail="Cannot delete running item")
|
||||
|
||||
redis_mgr.delete_item(item_id)
|
||||
return {"ok": True, "message": f"Item {item_id} deleted"}
|
||||
|
||||
@router.delete("/queue/clear")
|
||||
async def clear_queue(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Clear all non-running items from the queue."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
items = redis_mgr.list_items(include_finished=True)
|
||||
deleted_count = 0
|
||||
|
||||
for item in items:
|
||||
# Skip running items
|
||||
if item.status == QueueItemStatus.RUNNING:
|
||||
continue
|
||||
|
||||
# Check access
|
||||
if not data.staff and item.owner_pgroup not in data.pgroups:
|
||||
continue
|
||||
|
||||
redis_mgr.delete_item(item.item_id)
|
||||
deleted_count += 1
|
||||
|
||||
return {"ok": True, "message": f"Cleared {deleted_count} items from queue"}
|
||||
|
||||
|
||||
@router.post("/control/skip_sample", response_model=ControlActionResponse)
|
||||
async def request_skip_sample(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
runner: PersistentWorkflowRunner = Depends(get_runner),
|
||||
):
|
||||
"""Skip the current sample entirely (abort current sample and move to next)."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
runtime = redis_mgr.get_runtime()
|
||||
|
||||
if not runtime.running or not runtime.current_item_id:
|
||||
return ControlActionResponse(
|
||||
ok=False,
|
||||
control=redis_mgr.get_control(),
|
||||
message="No sample currently running",
|
||||
)
|
||||
|
||||
# Mark current item as skipped and complete it
|
||||
runner.complete_item(runtime.current_item_id, QueueItemStatus.SKIPPED)
|
||||
|
||||
redis_mgr.append_event(WorkflowEvent(
|
||||
beamline=runner.beamline,
|
||||
item_id=runtime.current_item_id,
|
||||
event_type="sample_skipped",
|
||||
actor=data.sub,
|
||||
message="Sample skipped by user",
|
||||
))
|
||||
|
||||
return ControlActionResponse(
|
||||
ok=True,
|
||||
control=redis_mgr.get_control(),
|
||||
message="Sample skipped",
|
||||
)
|
||||
|
||||
@router.post("/queue/{item_id}/move")
|
||||
async def move_queue_item(
|
||||
item_id: str,
|
||||
request: MoveItemRequest,
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Reorder an item in the queue."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
item = redis_mgr.get_item(item_id)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=404, detail="Item not found")
|
||||
|
||||
if not data.staff and item.owner_pgroup not in data.pgroups:
|
||||
raise HTTPException(status_code=403, detail="Access denied")
|
||||
|
||||
redis_mgr.move_item(item_id, request.new_order_index)
|
||||
return {"ok": True, "message": f"Item {item_id} moved"}
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Runtime and control endpoints
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.get("/runtime", response_model=RuntimeResponse)
|
||||
async def get_runtime(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Get current runtime and control state."""
|
||||
data = parse_token(token)
|
||||
check_jwt_ro(get_cfg(), data)
|
||||
|
||||
return RuntimeResponse(
|
||||
runtime=redis_mgr.get_runtime(),
|
||||
control=redis_mgr.get_control(),
|
||||
)
|
||||
|
||||
|
||||
@router.get("/control", response_model=ControlState)
|
||||
async def get_control(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Get current control state."""
|
||||
data = parse_token(token)
|
||||
check_jwt_ro(get_cfg(), data)
|
||||
return redis_mgr.get_control()
|
||||
|
||||
|
||||
@router.post("/control/pause", response_model=ControlActionResponse)
|
||||
async def request_pause(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Request workflow pause after current step."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
control = redis_mgr.request_control(
|
||||
{"pause_requested": True, "resume_requested": False},
|
||||
requested_by=data.sub,
|
||||
)
|
||||
|
||||
return ControlActionResponse(
|
||||
ok=True,
|
||||
control=control,
|
||||
message="Pause requested",
|
||||
)
|
||||
|
||||
|
||||
@router.post("/control/resume", response_model=ControlActionResponse)
|
||||
async def request_resume(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Resume paused workflow."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
control = redis_mgr.request_control(
|
||||
{"pause_requested": False, "resume_requested": True},
|
||||
requested_by=data.sub,
|
||||
)
|
||||
|
||||
return ControlActionResponse(
|
||||
ok=True,
|
||||
control=control,
|
||||
message="Resume requested",
|
||||
)
|
||||
|
||||
|
||||
@router.post("/control/abort", response_model=ControlActionResponse)
|
||||
async def request_abort(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Abort current workflow execution."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
# Clear pause_requested when aborting
|
||||
control = redis_mgr.request_control(
|
||||
{"abort_requested": True, "pause_requested": False, "resume_requested": False},
|
||||
requested_by=data.sub,
|
||||
)
|
||||
|
||||
return ControlActionResponse(
|
||||
ok=True,
|
||||
control=control,
|
||||
message="Abort requested",
|
||||
)
|
||||
|
||||
|
||||
@router.post("/control/skip", response_model=ControlActionResponse)
|
||||
async def request_skip(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Skip current step."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
control = redis_mgr.request_control(
|
||||
{"skip_requested": True},
|
||||
requested_by=data.sub,
|
||||
)
|
||||
|
||||
return ControlActionResponse(
|
||||
ok=True,
|
||||
control=control,
|
||||
message="Skip requested",
|
||||
)
|
||||
|
||||
|
||||
@router.post("/control/clear", response_model=ControlActionResponse)
|
||||
async def clear_control(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Clear all control requests."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
control = redis_mgr.clear_control()
|
||||
|
||||
return ControlActionResponse(
|
||||
ok=True,
|
||||
control=control,
|
||||
message="Control state cleared",
|
||||
)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Execution endpoints
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.post("/start/{item_id}", response_model=StepActionResponse)
|
||||
async def start_item(
|
||||
item_id: str,
|
||||
mode: WorkflowMode = WorkflowMode.GUIDED_MANUAL,
|
||||
token: str = Depends(oauth2_scheme),
|
||||
runner: PersistentWorkflowRunner = Depends(get_runner),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Start processing a queue item."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
item = redis_mgr.get_item(item_id)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=404, detail="Item not found")
|
||||
|
||||
if not data.staff and item.owner_pgroup not in data.pgroups:
|
||||
raise HTTPException(status_code=403, detail="Access denied")
|
||||
|
||||
if item.status == QueueItemStatus.RUNNING:
|
||||
raise HTTPException(status_code=409, detail="Item already running")
|
||||
|
||||
# Check if another item is running
|
||||
runtime = redis_mgr.get_runtime()
|
||||
if runtime.running:
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
detail=f"Another item is running: {runtime.current_item_id}",
|
||||
)
|
||||
|
||||
# Clear any stale control requests
|
||||
redis_mgr.clear_control()
|
||||
|
||||
# Start the item
|
||||
item = runner.start_item(item_id)
|
||||
|
||||
return StepActionResponse(
|
||||
ok=True,
|
||||
item=item,
|
||||
step=item.steps[0].kind if item.steps else None,
|
||||
status="started",
|
||||
message=f"Started processing {item.sample_name}",
|
||||
)
|
||||
|
||||
|
||||
@router.post("/next", response_model=StepActionResponse)
|
||||
async def run_next_step(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
runner: PersistentWorkflowRunner = Depends(get_runner),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""
|
||||
Run the next step in guided mode.
|
||||
|
||||
This is the main endpoint for guided manual operation.
|
||||
User clicks "Next" and this executes one step.
|
||||
"""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
runtime = redis_mgr.get_runtime()
|
||||
|
||||
if not runtime.running or not runtime.current_item_id:
|
||||
raise HTTPException(status_code=409, detail="No item is currently running")
|
||||
|
||||
item = redis_mgr.get_item(runtime.current_item_id)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=404, detail="Running item not found")
|
||||
|
||||
if not data.staff and item.owner_pgroup not in data.pgroups:
|
||||
raise HTTPException(status_code=403, detail="Access denied")
|
||||
|
||||
# Check if all steps completed
|
||||
if item.current_step_index >= len(item.steps):
|
||||
item = runner.complete_item(item.item_id, QueueItemStatus.COMPLETED)
|
||||
return StepActionResponse(
|
||||
ok=True,
|
||||
item=item,
|
||||
status="completed",
|
||||
message="All steps completed",
|
||||
)
|
||||
|
||||
# Check for abort
|
||||
if runner.should_abort():
|
||||
item = runner.complete_item(item.item_id, QueueItemStatus.ABORTED)
|
||||
redis_mgr.clear_control()
|
||||
return StepActionResponse(
|
||||
ok=True,
|
||||
item=item,
|
||||
status="aborted",
|
||||
message="Workflow aborted by user",
|
||||
)
|
||||
|
||||
# Check for skip
|
||||
control = runner.check_control()
|
||||
if control.skip_requested:
|
||||
step_index = item.current_step_index
|
||||
step_kind = item.steps[step_index].kind
|
||||
|
||||
redis_mgr.update_step(item.item_id, step_index, status="skipped")
|
||||
redis_mgr.update_item(item.item_id, {"current_step_index": step_index + 1})
|
||||
redis_mgr.request_control({"skip_requested": False}, requested_by=data.sub)
|
||||
|
||||
redis_mgr.append_event(WorkflowEvent(
|
||||
beamline=runner.beamline,
|
||||
item_id=item.item_id,
|
||||
step=step_kind,
|
||||
event_type="step_skipped",
|
||||
actor=data.sub,
|
||||
message=f"Step {step_kind} skipped by user",
|
||||
))
|
||||
|
||||
item = redis_mgr.get_item(item.item_id)
|
||||
return StepActionResponse(
|
||||
ok=True,
|
||||
item=item,
|
||||
step=step_kind,
|
||||
status="skipped",
|
||||
message=f"Skipped {step_kind}",
|
||||
)
|
||||
|
||||
# Build context
|
||||
context = WorkflowContext(
|
||||
mode=WorkflowMode.GUIDED_MANUAL,
|
||||
queue_id="default",
|
||||
item_id=item.item_id,
|
||||
sample_id=item.sample_id,
|
||||
current_state=WorkflowStateKind(runtime.current_state) if runtime.current_state else None,
|
||||
current_step_index=item.current_step_index,
|
||||
)
|
||||
|
||||
# Run the step
|
||||
try:
|
||||
result = runner.run_current_step(context)
|
||||
except Exception as e:
|
||||
return StepActionResponse(
|
||||
ok=False,
|
||||
item=redis_mgr.get_item(item.item_id),
|
||||
step=item.steps[item.current_step_index].kind,
|
||||
status="failed",
|
||||
message=str(e),
|
||||
)
|
||||
|
||||
# Get updated item
|
||||
item = redis_mgr.get_item(item.item_id)
|
||||
|
||||
# Check if completed
|
||||
if item.current_step_index >= len(item.steps):
|
||||
item = runner.complete_item(item.item_id, QueueItemStatus.COMPLETED)
|
||||
return StepActionResponse(
|
||||
ok=True,
|
||||
item=item,
|
||||
step=result.state.value,
|
||||
status="completed",
|
||||
message="All steps completed",
|
||||
)
|
||||
|
||||
return StepActionResponse(
|
||||
ok=True,
|
||||
item=item,
|
||||
step=result.state.value,
|
||||
status=result.status.value,
|
||||
message=result.message,
|
||||
)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Events endpoints
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.get("/events", response_model=EventListResponse)
|
||||
async def get_events(
|
||||
limit: int = 100,
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Get recent workflow events."""
|
||||
data = parse_token(token)
|
||||
check_jwt_ro(get_cfg(), data)
|
||||
|
||||
events = redis_mgr.read_events(limit=limit)
|
||||
return EventListResponse(events=events)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# SSE stream
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
|
||||
async def workflow_event_stream(
|
||||
redis_mgr: WorkflowRedisManager,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""
|
||||
Server-Sent Events stream for live workflow updates.
|
||||
|
||||
Polls runtime state and streams changes.
|
||||
"""
|
||||
last_runtime_json = ""
|
||||
last_control_json = ""
|
||||
last_event_id = "0-0"
|
||||
|
||||
try:
|
||||
while True:
|
||||
# Check runtime state
|
||||
runtime = redis_mgr.get_runtime()
|
||||
runtime_json = runtime.model_dump_json()
|
||||
|
||||
if runtime_json != last_runtime_json:
|
||||
last_runtime_json = runtime_json
|
||||
yield f"event: runtime\ndata: {runtime_json}\n\n"
|
||||
|
||||
# Check control state
|
||||
control = redis_mgr.get_control()
|
||||
control_json = control.model_dump_json()
|
||||
|
||||
if control_json != last_control_json:
|
||||
last_control_json = control_json
|
||||
yield f"event: control\ndata: {control_json}\n\n"
|
||||
|
||||
# Check for new events (using Redis streams)
|
||||
try:
|
||||
events_key = redis_mgr._events_key()
|
||||
new_events = redis_mgr._client.xread(
|
||||
{events_key: last_event_id},
|
||||
count=10,
|
||||
block=0,
|
||||
)
|
||||
|
||||
if new_events:
|
||||
for _, messages in new_events:
|
||||
for msg_id, fields in messages:
|
||||
last_event_id = msg_id
|
||||
raw = fields.get("json")
|
||||
if raw:
|
||||
yield f"event: workflow_event\ndata: {raw}\n\n"
|
||||
except Exception:
|
||||
pass # Redis stream read failed, continue polling
|
||||
|
||||
await asyncio.sleep(0.2)
|
||||
|
||||
except asyncio.CancelledError:
|
||||
return
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Automation mode endpoints
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.get("/automation/status", response_model=AutomationStatusResponse)
|
||||
async def get_automation_status(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
loop: "AutomationLoop" = Depends(get_automation_loop),
|
||||
):
|
||||
"""Get automation loop status."""
|
||||
data = parse_token(token)
|
||||
check_jwt_ro(get_cfg(), data)
|
||||
|
||||
return AutomationStatusResponse(
|
||||
enabled=loop.is_enabled,
|
||||
running=loop.is_running,
|
||||
runtime=redis_mgr.get_runtime(),
|
||||
control=redis_mgr.get_control(),
|
||||
)
|
||||
|
||||
|
||||
@router.post("/automation/start", response_model=AutomationStatusResponse)
|
||||
async def start_automation(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
loop: "AutomationLoop" = Depends(get_automation_loop),
|
||||
):
|
||||
"""Start automation mode - processes queue automatically."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
loop.start()
|
||||
|
||||
return AutomationStatusResponse(
|
||||
enabled=loop.is_enabled,
|
||||
running=loop.is_running,
|
||||
runtime=redis_mgr.get_runtime(),
|
||||
control=redis_mgr.get_control(),
|
||||
)
|
||||
|
||||
|
||||
@router.post("/automation/stop", response_model=AutomationStatusResponse)
|
||||
async def stop_automation(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
loop: "AutomationLoop" = Depends(get_automation_loop),
|
||||
):
|
||||
"""Stop automation mode - completes current step then stops."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
loop.stop()
|
||||
|
||||
return AutomationStatusResponse(
|
||||
enabled=loop.is_enabled,
|
||||
running=loop.is_running,
|
||||
runtime=redis_mgr.get_runtime(),
|
||||
control=redis_mgr.get_control(),
|
||||
)
|
||||
|
||||
@router.get("/sse")
|
||||
async def workflow_sse(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""
|
||||
SSE endpoint for live workflow updates.
|
||||
|
||||
Events:
|
||||
- runtime: RuntimeState changes
|
||||
- control: ControlState changes
|
||||
- workflow_event: Individual workflow events
|
||||
"""
|
||||
data = parse_token(token)
|
||||
check_jwt_ro(get_cfg(), data)
|
||||
|
||||
return StreamingResponse(
|
||||
workflow_event_stream(redis_mgr),
|
||||
media_type="text/event-stream",
|
||||
headers={
|
||||
"Cache-Control": "no-cache",
|
||||
"Connection": "keep-alive",
|
||||
"Access-Control-Allow-Origin": "*",
|
||||
"Access-Control-Allow-Headers": "Cache-Control",
|
||||
},
|
||||
)
|
||||
@@ -1,461 +0,0 @@
|
||||
import asyncio
|
||||
from typing import Callable
|
||||
|
||||
import redis
|
||||
|
||||
from aare.common.automation_models import (
|
||||
WorkflowEvent,
|
||||
QueueItemStatus,
|
||||
QueueItem,
|
||||
ControlState,
|
||||
WorkflowStateKind,
|
||||
WorkflowContext,
|
||||
StateResult,
|
||||
StateDefinition,
|
||||
WorkflowMode,
|
||||
)
|
||||
from aare.common.automation_queue_manager import (
|
||||
WorkflowRedisManager,
|
||||
build_default_steps
|
||||
)
|
||||
|
||||
from aare.common.automation_workflow import (
|
||||
can_transition,
|
||||
StateHandler,
|
||||
STATE_REGISTRY,
|
||||
HANDLER_REGISTRY
|
||||
)
|
||||
|
||||
|
||||
class PersistentWorkflowRunner:
|
||||
"""
|
||||
Orchestrates workflow execution with Redis persistence.
|
||||
|
||||
Combines:
|
||||
- State handlers (from automation_workflow)
|
||||
- Redis persistence (from automation_queue_manager)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
redis_manager: WorkflowRedisManager,
|
||||
registry: dict[WorkflowStateKind, StateDefinition] | None = None,
|
||||
handlers: dict[WorkflowStateKind, StateHandler] | None = None,
|
||||
):
|
||||
self._redis = redis_manager
|
||||
self._registry = registry or STATE_REGISTRY
|
||||
self._handlers = handlers or HANDLER_REGISTRY
|
||||
|
||||
@property
|
||||
def beamline(self) -> str:
|
||||
return self._redis._bl
|
||||
|
||||
def get_handler(self, state: WorkflowStateKind) -> StateHandler:
|
||||
try:
|
||||
return self._handlers[state]
|
||||
except KeyError as exc:
|
||||
raise KeyError(f"No handler registered for state: {state}") from exc
|
||||
|
||||
def start_item(self, item_id: str) -> QueueItem:
|
||||
item = self._redis.get_item(item_id)
|
||||
if item is None:
|
||||
raise KeyError(f"Item not found: {item_id}")
|
||||
|
||||
item = self._redis.update_item(item_id, {
|
||||
"status": QueueItemStatus.RUNNING,
|
||||
"current_step_index": 0,
|
||||
})
|
||||
|
||||
self._redis.patch_runtime({
|
||||
"running": True,
|
||||
"paused": False,
|
||||
"current_item_id": item_id,
|
||||
"current_step_index": 0,
|
||||
"current_state": None,
|
||||
"last_error": None,
|
||||
})
|
||||
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self.beamline,
|
||||
item_id=item_id,
|
||||
event_type="item_started",
|
||||
message=f"Started processing {item.sample_name}",
|
||||
))
|
||||
|
||||
return item
|
||||
|
||||
def run_current_step(self, context: WorkflowContext) -> StateResult:
|
||||
item = self._redis.get_item(context.item_id)
|
||||
if item is None:
|
||||
raise KeyError(f"Item not found: {context.item_id}")
|
||||
|
||||
step_index = item.current_step_index
|
||||
if step_index >= len(item.steps):
|
||||
raise RuntimeError("No more steps to run")
|
||||
|
||||
step_record = item.steps[step_index]
|
||||
state_kind = WorkflowStateKind(step_record.kind)
|
||||
|
||||
# Check transition is allowed
|
||||
if context.current_state is not None:
|
||||
if not can_transition(context.current_state, state_kind, context.mode):
|
||||
raise RuntimeError(
|
||||
f"Transition not allowed: {context.current_state} -> {state_kind}"
|
||||
)
|
||||
|
||||
# Mark step as running
|
||||
self._redis.update_step(context.item_id, step_index, status="running")
|
||||
self._redis.patch_runtime({
|
||||
"current_state": state_kind.value,
|
||||
"current_step_index": step_index,
|
||||
})
|
||||
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self.beamline,
|
||||
item_id=context.item_id,
|
||||
step=state_kind.value,
|
||||
event_type="step_started",
|
||||
message=f"Starting {state_kind.value}",
|
||||
))
|
||||
|
||||
# Execute handler
|
||||
handler = self.get_handler(state_kind)
|
||||
|
||||
try:
|
||||
result = handler.execute(context)
|
||||
except Exception as e:
|
||||
self._redis.update_step(
|
||||
context.item_id,
|
||||
step_index,
|
||||
status="failed",
|
||||
error_detail=str(e),
|
||||
)
|
||||
self._redis.patch_runtime({"last_error": str(e)})
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self.beamline,
|
||||
item_id=context.item_id,
|
||||
step=state_kind.value,
|
||||
event_type="step_failed",
|
||||
message=str(e),
|
||||
))
|
||||
raise
|
||||
|
||||
# Mark step as completed
|
||||
self._redis.update_step(
|
||||
context.item_id,
|
||||
step_index,
|
||||
status=result.status.value,
|
||||
message=result.message,
|
||||
)
|
||||
|
||||
# Advance step index
|
||||
self._redis.update_item(context.item_id, {
|
||||
"current_step_index": step_index + 1,
|
||||
})
|
||||
|
||||
context.current_state = state_kind
|
||||
context.current_step_index = step_index + 1
|
||||
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self.beamline,
|
||||
item_id=context.item_id,
|
||||
step=state_kind.value,
|
||||
event_type="step_completed",
|
||||
message=result.message,
|
||||
payload=result.payload,
|
||||
))
|
||||
|
||||
return result
|
||||
|
||||
def complete_item(self, item_id: str, status: QueueItemStatus) -> QueueItem:
|
||||
item = self._redis.update_item(item_id, {"status": status})
|
||||
|
||||
self._redis.patch_runtime({
|
||||
"running": False,
|
||||
"current_item_id": None,
|
||||
"current_state": None,
|
||||
})
|
||||
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self.beamline,
|
||||
item_id=item_id,
|
||||
event_type="item_completed",
|
||||
message=f"Item finished with status {status.value}",
|
||||
))
|
||||
|
||||
if status in (QueueItemStatus.COMPLETED, QueueItemStatus.ABORTED, QueueItemStatus.SKIPPED):
|
||||
self._redis.delete_item(item_id)
|
||||
return item
|
||||
|
||||
return item
|
||||
|
||||
def check_control(self) -> ControlState:
|
||||
return self._redis.get_control()
|
||||
|
||||
def should_pause(self) -> bool:
|
||||
return self.check_control().pause_requested
|
||||
|
||||
def should_abort(self) -> bool:
|
||||
return self.check_control().abort_requested
|
||||
|
||||
def should_skip(self) -> bool:
|
||||
return self.check_control().skip_requested
|
||||
|
||||
|
||||
class AutomationLoop:
|
||||
"""
|
||||
Background task that drives fully automated workflow execution.
|
||||
|
||||
Polls control state and processes the queue automatically.
|
||||
|
||||
IMPORTANT: This loop does NOT auto-start. It must be explicitly started
|
||||
via the start() method (triggered by the "Start Automation" button).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
runner: PersistentWorkflowRunner,
|
||||
redis_manager: WorkflowRedisManager,
|
||||
poll_interval: float = 0.3,
|
||||
step_delay: float = 0.5, # Delay between steps for control checks
|
||||
):
|
||||
self._runner = runner
|
||||
self._redis = redis_manager
|
||||
self._poll_interval = poll_interval
|
||||
self._step_delay = step_delay
|
||||
self._task: asyncio.Task | None = None
|
||||
self._enabled = False
|
||||
self._on_step_complete: Callable[[str, str, StateResult], None] | None = None
|
||||
|
||||
@property
|
||||
def is_running(self) -> bool:
|
||||
return self._task is not None and not self._task.done()
|
||||
|
||||
@property
|
||||
def is_enabled(self) -> bool:
|
||||
return self._enabled
|
||||
|
||||
def set_step_callback(self, callback: Callable[[str, str, StateResult], None]) -> None:
|
||||
"""Set callback for step completion: callback(item_id, step_name, result)"""
|
||||
self._on_step_complete = callback
|
||||
|
||||
def start(self) -> None:
|
||||
"""Start the automation loop. Must be explicitly called."""
|
||||
if self._task is not None and not self._task.done():
|
||||
return # Already running
|
||||
|
||||
self._enabled = True
|
||||
self._task = asyncio.create_task(self._run_loop())
|
||||
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self._runner.beamline,
|
||||
event_type="automation_started",
|
||||
message="Automation mode enabled",
|
||||
))
|
||||
|
||||
def stop(self) -> None:
|
||||
"""Stop the automation loop gracefully."""
|
||||
self._enabled = False
|
||||
|
||||
# Cancel the task if it exists
|
||||
if self._task is not None and not self._task.done():
|
||||
self._task.cancel()
|
||||
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self._runner.beamline,
|
||||
event_type="automation_stopped",
|
||||
message="Automation mode disabled",
|
||||
))
|
||||
|
||||
async def _run_loop(self) -> None:
|
||||
"""Main automation loop."""
|
||||
while self._enabled:
|
||||
try:
|
||||
# Check control state FIRST before doing anything
|
||||
control = self._runner.check_control()
|
||||
runtime = self._redis.get_runtime()
|
||||
|
||||
# Handle abort immediately
|
||||
if control.abort_requested:
|
||||
if runtime.current_item_id:
|
||||
self._runner.complete_item(runtime.current_item_id, QueueItemStatus.ABORTED)
|
||||
# Clear all control flags including pause
|
||||
self._redis.clear_control()
|
||||
self.stop()
|
||||
await asyncio.sleep(self._poll_interval)
|
||||
continue
|
||||
|
||||
# Handle pause
|
||||
if control.pause_requested:
|
||||
if runtime.running and not runtime.paused:
|
||||
self._redis.patch_runtime({"paused": True})
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self._runner.beamline,
|
||||
item_id=runtime.current_item_id,
|
||||
event_type="workflow_paused",
|
||||
message="Workflow paused by user request",
|
||||
))
|
||||
await asyncio.sleep(self._poll_interval)
|
||||
continue
|
||||
|
||||
# Handle resume
|
||||
if control.resume_requested and runtime.paused:
|
||||
self._redis.patch_runtime({"paused": False})
|
||||
self._redis.request_control(
|
||||
{"resume_requested": False},
|
||||
requested_by="automation_loop",
|
||||
)
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self._runner.beamline,
|
||||
item_id=runtime.current_item_id,
|
||||
event_type="workflow_resumed",
|
||||
message="Workflow resumed",
|
||||
))
|
||||
|
||||
# Don't process if paused
|
||||
if runtime.paused:
|
||||
await asyncio.sleep(self._poll_interval)
|
||||
continue
|
||||
|
||||
await self._tick()
|
||||
|
||||
except asyncio.CancelledError:
|
||||
# Loop was cancelled (stop() was called)
|
||||
break
|
||||
except Exception as e:
|
||||
self._redis.patch_runtime({"last_error": str(e)})
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self._runner.beamline,
|
||||
event_type="automation_error",
|
||||
message=f"Automation error: {e}",
|
||||
))
|
||||
await asyncio.sleep(2.0)
|
||||
|
||||
await asyncio.sleep(self._poll_interval)
|
||||
|
||||
async def _tick(self) -> None:
|
||||
"""Single iteration of the automation loop - process one step."""
|
||||
runtime = self._redis.get_runtime()
|
||||
control = self._runner.check_control()
|
||||
|
||||
# If nothing running, try to start next item
|
||||
if not runtime.running or not runtime.current_item_id:
|
||||
next_item = self._redis.get_next_pending_item()
|
||||
if next_item is None:
|
||||
return # Queue empty, nothing to do
|
||||
|
||||
self._runner.start_item(next_item.item_id)
|
||||
# Add delay after starting to allow UI to catch up
|
||||
await asyncio.sleep(self._step_delay)
|
||||
runtime = self._redis.get_runtime()
|
||||
|
||||
# Re-check control after potential start
|
||||
control = self._runner.check_control()
|
||||
if control.abort_requested or control.pause_requested:
|
||||
return # Let main loop handle it
|
||||
|
||||
# Process current item
|
||||
item = self._redis.get_item(runtime.current_item_id)
|
||||
if item is None:
|
||||
self._redis.patch_runtime({"running": False, "current_item_id": None})
|
||||
return
|
||||
|
||||
# Check if item is complete
|
||||
if item.current_step_index >= len(item.steps):
|
||||
self._runner.complete_item(item.item_id, QueueItemStatus.COMPLETED)
|
||||
# Clear next_sample_requested if set
|
||||
if control.next_sample_requested:
|
||||
self._redis.request_control(
|
||||
{"next_sample_requested": False},
|
||||
requested_by="automation_loop",
|
||||
)
|
||||
return
|
||||
|
||||
# Handle skip request (skip current step)
|
||||
if control.skip_requested:
|
||||
step_kind = item.steps[item.current_step_index].kind
|
||||
self._redis.update_step(item.item_id, item.current_step_index, status="skipped")
|
||||
self._redis.update_item(item.item_id, {"current_step_index": item.current_step_index + 1})
|
||||
self._redis.request_control({"skip_requested": False}, requested_by="automation_loop")
|
||||
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self._runner.beamline,
|
||||
item_id=item.item_id,
|
||||
step=step_kind,
|
||||
event_type="step_skipped",
|
||||
message=f"Step {step_kind} skipped",
|
||||
))
|
||||
return
|
||||
|
||||
# Build context and run step
|
||||
context = WorkflowContext(
|
||||
mode=WorkflowMode.AUTOMATION,
|
||||
queue_id="default",
|
||||
item_id=item.item_id,
|
||||
sample_id=item.sample_id,
|
||||
current_state=WorkflowStateKind(runtime.current_state) if runtime.current_state else None,
|
||||
current_step_index=item.current_step_index,
|
||||
)
|
||||
|
||||
# Run the step (this is blocking in the async context)
|
||||
step_name = item.steps[item.current_step_index].kind
|
||||
result = await asyncio.to_thread(self._runner.run_current_step, context)
|
||||
|
||||
# Add delay after step to allow control checks
|
||||
await asyncio.sleep(self._step_delay)
|
||||
|
||||
# Notify callback if set
|
||||
if self._on_step_complete:
|
||||
self._on_step_complete(item.item_id, step_name, result)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
client = redis.Redis(host="localhost", port=6379, db=0, decode_responses=True)
|
||||
redis_mgr = WorkflowRedisManager(client, beamline="x10sa")
|
||||
|
||||
# Create a queue item
|
||||
item = QueueItem(
|
||||
item_id="",
|
||||
beamline="x10sa",
|
||||
sample_id=123,
|
||||
sample_name="lysozyme_01",
|
||||
owner_pgroup="p12345",
|
||||
created_by="user@example.com",
|
||||
steps=build_default_steps(),
|
||||
)
|
||||
item = redis_mgr.create_item(item)
|
||||
|
||||
# Create runner
|
||||
runner = PersistentWorkflowRunner(
|
||||
redis_manager=redis_mgr,
|
||||
registry=STATE_REGISTRY,
|
||||
handlers=HANDLER_REGISTRY,
|
||||
)
|
||||
|
||||
# Start processing
|
||||
runner.start_item(item.item_id)
|
||||
|
||||
# Build context
|
||||
context = WorkflowContext(
|
||||
mode=WorkflowMode.GUIDED_MANUAL,
|
||||
queue_id="default",
|
||||
item_id=item.item_id,
|
||||
sample_id=item.sample_id,
|
||||
)
|
||||
|
||||
# Run each step (in guided mode, user triggers each one)
|
||||
while context.current_step_index < len(item.steps):
|
||||
if runner.should_pause():
|
||||
print("Paused by user")
|
||||
break
|
||||
|
||||
if runner.should_abort():
|
||||
runner.complete_item(item.item_id, QueueItemStatus.ABORTED)
|
||||
break
|
||||
|
||||
result = runner.run_current_step(context)
|
||||
print(f"Step {result.state.value}: {result.status.value}")
|
||||
|
||||
# Mark complete if all steps done
|
||||
if context.current_step_index >= len(item.steps):
|
||||
runner.complete_item(item.item_id, QueueItemStatus.COMPLETED)
|
||||
+1
-27
@@ -42,11 +42,6 @@ from aare.common.exception_handler import (
|
||||
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")
|
||||
|
||||
# OAuth2 setup
|
||||
@@ -56,10 +51,6 @@ oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token")
|
||||
bl = None
|
||||
cfg = None
|
||||
daq = None
|
||||
workflow_redis_manager = None
|
||||
workflow_runner = None
|
||||
|
||||
USE_SIMULATED_WORKFLOW = os.getenv("WORKFLOW_SIMULATION", "0") == "1"
|
||||
|
||||
_all_pgroups_cache: dict[str, tuple[list[str], float]] = {}
|
||||
_ALL_PGROUPS_TTL_S = 60.0 # adjust TTL as needed
|
||||
@@ -86,7 +77,7 @@ async def lifespan(application: FastAPI):
|
||||
All stateful / connection-opening initialisation belongs here so that
|
||||
each worker gets its own fresh Redis, BEC, EPICS, and TELL connections.
|
||||
"""
|
||||
global bl, cfg, daq, workflow_redis_manager, workflow_runner
|
||||
global bl, cfg, daq
|
||||
|
||||
await asyncio.sleep(random.uniform(0.5, 3.0))
|
||||
|
||||
@@ -97,21 +88,6 @@ async def lifespan(application: FastAPI):
|
||||
cfg = BeamlineConfig(bl)
|
||||
daq = AareDAQ(cfg, bl)
|
||||
|
||||
# ── Workflow system ──
|
||||
# workflow_redis_manager = WorkflowRedisManager(
|
||||
# client=cfg._BeamlineConfig__client, # Reuse this worker's 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)
|
||||
|
||||
try:
|
||||
cfg.reset_automation_progress()
|
||||
except Exception as e:
|
||||
@@ -137,8 +113,6 @@ async def lifespan(application: FastAPI):
|
||||
app = FastAPI(lifespan=lifespan)
|
||||
register_exception_handlers(app)
|
||||
|
||||
app.include_router(workflow_router)
|
||||
|
||||
def _required_recovery_code() -> str:
|
||||
"""
|
||||
Get the required recovery confirmation code from environment variables.
|
||||
|
||||
@@ -61,11 +61,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, AutomationProgressWidget
|
||||
from aare.gui.threads.workflow_sse_client import WorkflowSSEClient
|
||||
from aare.gui.panels.automation_panel import AutomationProgressWidget
|
||||
|
||||
logger = setup_logger("aareGUI")
|
||||
WORKFLOW_SSE_TEST = False
|
||||
|
||||
class MainWindow(QMainWindow):
|
||||
sample_geometry = Signal(SampleGeometryModel)
|
||||
@@ -362,17 +360,6 @@ class MainWindow(QMainWindow):
|
||||
self.addDockWidget(Qt.DockWidgetArea.RightDockWidgetArea, self.prediction_metrics_dock)
|
||||
self.prediction_metrics_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)
|
||||
|
||||
@@ -647,49 +634,6 @@ class MainWindow(QMainWindow):
|
||||
|
||||
register_tutorials(self, self.tutorial_manager)
|
||||
|
||||
# Workflow SSE client
|
||||
|
||||
if self.__base_url is not None and WORKFLOW_SSE_TEST is True:
|
||||
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
|
||||
if WORKFLOW_SSE_TEST is True:
|
||||
self.workflow_panel.request_queue_refresh.connect(self.daq.workflow_load_queue)
|
||||
self.workflow_panel.request_add_sample.connect(self.daq.workflow_add_sample)
|
||||
self.workflow_panel.request_add_samples.connect(self.daq.workflow_add_samples)
|
||||
self.workflow_panel.request_delete_item.connect(self.daq.workflow_delete_item)
|
||||
self.workflow_panel.request_move_item.connect(self.daq.workflow_move_item)
|
||||
self.workflow_panel.request_clear_queue.connect(self.daq.workflow_clear_queue)
|
||||
self.workflow_panel.request_start_item.connect(self.daq.workflow_start_item)
|
||||
self.workflow_panel.request_next_step.connect(self.daq.workflow_next_step)
|
||||
self.workflow_panel.request_pause.connect(self.daq.workflow_pause)
|
||||
self.workflow_panel.request_resume.connect(self.daq.workflow_resume)
|
||||
self.workflow_panel.request_abort.connect(self.daq.workflow_abort)
|
||||
self.workflow_panel.request_skip.connect(self.daq.workflow_skip)
|
||||
self.workflow_panel.request_start_automation.connect(self.daq.workflow_start_automation)
|
||||
self.workflow_panel.request_stop_automation.connect(self.daq.workflow_stop_automation)
|
||||
self.workflow_panel.request_skip_sample.connect(self.daq.workflow_skip_sample)
|
||||
|
||||
# DAQ worker -> workflow panel
|
||||
self.daq.workflow_queue_loaded.connect(self.workflow_panel.update_queue)
|
||||
self.daq.workflow_item_updated.connect(self.workflow_panel.update_current_item)
|
||||
|
||||
# Forward sample list to workflow panel for "Add All" feature
|
||||
self.daq.spreadsheet.connect(
|
||||
lambda slist: self.workflow_panel.update_available_samples(slist.s)
|
||||
)
|
||||
|
||||
# Initial load
|
||||
QTimer.singleShot(1000, self.daq.workflow_load_queue)
|
||||
|
||||
def _setup_global_shortcuts(self) -> None:
|
||||
self._shortcut_manual_sample = QAction("Raise Manual Sample Dock", self)
|
||||
self._shortcut_manual_sample.setShortcut(QKeySequence("Ctrl+M"))
|
||||
@@ -718,13 +662,6 @@ class MainWindow(QMainWindow):
|
||||
)
|
||||
self.addAction(self._shortcut_raise_job_list)
|
||||
|
||||
self._shortcut_toggle_workflow = QAction("Toggle workflow panel", self)
|
||||
self._shortcut_toggle_workflow.setShortcut(QKeySequence("Ctrl+W"))
|
||||
self._shortcut_toggle_workflow.triggered.connect(
|
||||
lambda: self.workflow_dock.setVisible(not self.workflow_dock.isVisible())
|
||||
)
|
||||
self.addAction(self._shortcut_toggle_workflow)
|
||||
|
||||
self._shortcut_toggle_target_stability = QAction("Toggle target stability panel", self)
|
||||
self._shortcut_toggle_target_stability.setShortcut(QKeySequence("Ctrl+Shift+T"))
|
||||
self._shortcut_toggle_target_stability.triggered.connect(
|
||||
@@ -933,13 +870,6 @@ class MainWindow(QMainWindow):
|
||||
self.prediction_metrics_dock.visibilityChanged.connect(show_prediction_metrics_action.setChecked)
|
||||
view_menu.addAction(show_prediction_metrics_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)
|
||||
@@ -1485,9 +1415,6 @@ class MainWindow(QMainWindow):
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to stop _remote_close_timer: {e}")
|
||||
|
||||
if hasattr(self, "workflow_sse") and self.workflow_sse is not None:
|
||||
self.workflow_sse.disconnect()
|
||||
|
||||
try:
|
||||
if hasattr(self, "daq") and self.daq is not None:
|
||||
self.daq.cleanup()
|
||||
|
||||
@@ -1,120 +1,31 @@
|
||||
"""
|
||||
Workflow automation panel for queue management and step control.
|
||||
AutomationProgressWidget
|
||||
|
||||
Displays:
|
||||
- Queue items with status (supports drag-drop from sample list)
|
||||
- Current step progress
|
||||
- Control buttons (Start Guided/Start Automation/Pause/Resume/Abort/Skip)
|
||||
- Live status from SSE
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
import time
|
||||
|
||||
from PySide6.QtCore import Qt, Signal, Slot, QTimer, QMimeData
|
||||
from PySide6.QtGui import QColor, QDragEnterEvent, QDropEvent, QKeySequence, QShortcut
|
||||
from PySide6.QtCore import Slot
|
||||
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,
|
||||
AutomationProgress,
|
||||
StepStatus,
|
||||
WorkflowStateKind,
|
||||
)
|
||||
from aare.common.models import SampleShortInfo, SampleShortInfoList
|
||||
|
||||
from aare.common.logger_config import setup_logger
|
||||
|
||||
logger = setup_logger("aareGUI")
|
||||
|
||||
|
||||
class StepProgressWidget(QWidget):
|
||||
"""Shows progress through workflow steps."""
|
||||
|
||||
def __init__(self, parent: QWidget | None = None):
|
||||
super().__init__(parent)
|
||||
self._steps: list[WorkflowStepRecord] = []
|
||||
self._current_index = 0
|
||||
self._setup_ui()
|
||||
|
||||
def _setup_ui(self) -> None:
|
||||
layout = QVBoxLayout(self)
|
||||
layout.setContentsMargins(4, 4, 4, 4)
|
||||
layout.setSpacing(2)
|
||||
|
||||
self._step_labels: list[QLabel] = []
|
||||
|
||||
self._container = QWidget()
|
||||
self._container_layout = QVBoxLayout(self._container)
|
||||
self._container_layout.setContentsMargins(0, 0, 0, 0)
|
||||
self._container_layout.setSpacing(2)
|
||||
layout.addWidget(self._container)
|
||||
|
||||
def set_steps(self, steps: list[WorkflowStepRecord], current_index: int) -> None:
|
||||
"""Update the step display."""
|
||||
self._steps = steps
|
||||
self._current_index = current_index
|
||||
|
||||
for lbl in self._step_labels:
|
||||
lbl.deleteLater()
|
||||
self._step_labels.clear()
|
||||
|
||||
for i, step in enumerate(steps):
|
||||
lbl = QLabel(f"{i + 1}. {step.kind}")
|
||||
lbl.setStyleSheet(self._style_for_step(i, step.status))
|
||||
self._container_layout.addWidget(lbl)
|
||||
self._step_labels.append(lbl)
|
||||
|
||||
def update_step(self, index: int, status: str) -> None:
|
||||
"""Update a single step's status."""
|
||||
if 0 <= index < len(self._step_labels):
|
||||
self._step_labels[index].setStyleSheet(self._style_for_step(index, status))
|
||||
|
||||
def clear_steps(self) -> None:
|
||||
"""Clear all step labels."""
|
||||
for lbl in self._step_labels:
|
||||
lbl.deleteLater()
|
||||
self._step_labels.clear()
|
||||
self._steps = []
|
||||
|
||||
def _style_for_step(self, index: int, status: str) -> str:
|
||||
"""Get stylesheet for step based on status."""
|
||||
base = "padding: 4px; border-radius: 3px; "
|
||||
|
||||
if status == "success":
|
||||
return base + "background-color: #90EE90; color: #006400;"
|
||||
elif status == "running":
|
||||
return base + "background-color: #87CEEB; color: #00008B; font-weight: bold;"
|
||||
elif status == "failed":
|
||||
return base + "background-color: #FFB6C1; color: #8B0000;"
|
||||
elif status == "skipped":
|
||||
return base + "background-color: #D3D3D3; color: #696969; text-decoration: line-through;"
|
||||
elif status == "paused":
|
||||
return base + "background-color: #FFE4B5; color: #8B4513;"
|
||||
else: # pending
|
||||
return base + "background-color: #F0F0F0; color: #808080;"
|
||||
|
||||
|
||||
class AutomationProgressWidget(QWidget):
|
||||
"""Compact fixed-step widget for DAQ automation progress."""
|
||||
|
||||
@@ -123,6 +34,7 @@ class AutomationProgressWidget(QWidget):
|
||||
self._progress: AutomationProgress | None = None
|
||||
self._labels: dict[WorkflowStateKind, QLabel] = {}
|
||||
self._title_label: QLabel | None = None
|
||||
self._stats_label: QLabel | None = None
|
||||
self._setup_ui()
|
||||
self.clear()
|
||||
|
||||
@@ -135,6 +47,11 @@ class AutomationProgressWidget(QWidget):
|
||||
self._title_label.setStyleSheet("font-size: 14px; font-weight: bold;")
|
||||
layout.addWidget(self._title_label)
|
||||
|
||||
self._stats_label = QLabel()
|
||||
self._stats_label.setStyleSheet("color: #555; font-size: 11px; margin-bottom: 4px;")
|
||||
self._stats_label.setWordWrap(True)
|
||||
layout.addWidget(self._stats_label)
|
||||
|
||||
for step in (
|
||||
WorkflowStateKind.MOUNT,
|
||||
WorkflowStateKind.LOOP_CENTRE,
|
||||
@@ -161,6 +78,8 @@ class AutomationProgressWidget(QWidget):
|
||||
],
|
||||
finished=False,
|
||||
success=None,
|
||||
samples_in_queue=0,
|
||||
current_sample_name="None"
|
||||
)
|
||||
self.set_progress(empty)
|
||||
|
||||
@@ -208,12 +127,17 @@ class AutomationProgressWidget(QWidget):
|
||||
|
||||
@Slot(object)
|
||||
def set_progress(self, progress: AutomationProgress) -> None:
|
||||
logger.info(
|
||||
f"[AutomationProgressWidget] current_step={progress.current_step} "
|
||||
f"finished={progress.finished} success={progress.success}"
|
||||
)
|
||||
self._progress = progress
|
||||
|
||||
# Update Queue Stats
|
||||
if self._stats_label:
|
||||
avg_time = f"{progress.avg_time_per_sample:.1f}s" if progress.avg_time_per_sample > 0 else "N/A"
|
||||
stats_text = (
|
||||
f"<b>Current:</b> {progress.current_sample_name or 'None'}<br/>"
|
||||
f"<b>Queue:</b> {progress.samples_in_queue} samples | <b>Avg:</b> {avg_time}"
|
||||
)
|
||||
self._stats_label.setText(stats_text)
|
||||
|
||||
for step_state in progress.steps:
|
||||
label = self._labels.get(step_state.step)
|
||||
if label is None:
|
||||
@@ -221,8 +145,17 @@ class AutomationProgressWidget(QWidget):
|
||||
|
||||
title = self._label_for_step(step_state.step)
|
||||
icon = self._icon_for_status(step_state.status)
|
||||
|
||||
# Calculate duration
|
||||
duration_str = ""
|
||||
if step_state.started_at:
|
||||
end = step_state.completed_at or time.time()
|
||||
duration_str = f" ({end - step_state.started_at:.1f}s)"
|
||||
|
||||
error_str = f" <br/><small>Error: {step_state.error_code}</small>" if step_state.error_code else ""
|
||||
message = f" — {step_state.message}" if step_state.message else ""
|
||||
label.setText(f"{icon} {title}{message}")
|
||||
|
||||
label.setText(f"{icon} <b>{title}</b>{duration_str}{message}{error_str}")
|
||||
label.setStyleSheet(self._style_for_status(step_state.status))
|
||||
|
||||
if self._title_label is not None:
|
||||
@@ -236,636 +169,4 @@ class AutomationProgressWidget(QWidget):
|
||||
elif progress.current_step:
|
||||
self._title_label.setText(f"Automation progress — {progress.current_step}")
|
||||
else:
|
||||
self._title_label.setText("Automation progress")
|
||||
|
||||
|
||||
class DraggableQueueListWidget(QListWidget):
|
||||
"""
|
||||
QListWidget that accepts drops from TellSamplePanel.
|
||||
|
||||
Supports:
|
||||
- Drag-drop samples from tell_sample_panel
|
||||
- Internal reordering via drag
|
||||
- Delete key to remove items
|
||||
"""
|
||||
|
||||
samples_dropped = Signal(list) # list[SampleShortInfo]
|
||||
item_reordered = Signal(str, int) # item_id, new_index
|
||||
delete_requested = Signal(list) # list[item_ids]
|
||||
start_item_requested = Signal(str) # item_id
|
||||
|
||||
def __init__(self, parent: QWidget | None = None):
|
||||
super().__init__(parent)
|
||||
|
||||
self.setAcceptDrops(True)
|
||||
self.setDragEnabled(True)
|
||||
self.setDragDropMode(QAbstractItemView.DragDropMode.DragDrop)
|
||||
self.setDefaultDropAction(Qt.DropAction.MoveAction)
|
||||
self.setSelectionMode(QAbstractItemView.SelectionMode.ExtendedSelection)
|
||||
self.setContextMenuPolicy(Qt.ContextMenuPolicy.CustomContextMenu)
|
||||
self.customContextMenuRequested.connect(self._show_context_menu)
|
||||
|
||||
# Delete shortcut
|
||||
self._delete_shortcut = QShortcut(QKeySequence.StandardKey.Delete, self)
|
||||
self._delete_shortcut.activated.connect(self._on_delete_pressed)
|
||||
|
||||
# Store item_id -> row mapping
|
||||
self._item_ids: list[str] = []
|
||||
|
||||
def set_item_ids(self, item_ids: list[str]) -> None:
|
||||
"""Track item IDs for reordering."""
|
||||
self._item_ids = item_ids
|
||||
|
||||
def dragEnterEvent(self, event: QDragEnterEvent) -> None:
|
||||
"""Accept drops from sample panels."""
|
||||
mime = event.mimeData()
|
||||
if mime.hasText():
|
||||
try:
|
||||
text = mime.text()
|
||||
if text.startswith("{") or text.startswith("["):
|
||||
event.acceptProposedAction()
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
if event.source() == self:
|
||||
event.acceptProposedAction()
|
||||
return
|
||||
event.ignore()
|
||||
|
||||
def dragMoveEvent(self, event) -> None:
|
||||
"""Show drop indicator."""
|
||||
if event.mimeData().hasText() or event.source() == self:
|
||||
event.acceptProposedAction()
|
||||
else:
|
||||
event.ignore()
|
||||
|
||||
def dropEvent(self, event: QDropEvent) -> None:
|
||||
"""Handle drop - either samples from panel or internal reorder."""
|
||||
mime = event.mimeData()
|
||||
|
||||
if mime.hasText():
|
||||
text = mime.text()
|
||||
try:
|
||||
sample_list = SampleShortInfoList.model_validate_json(text)
|
||||
if sample_list.s:
|
||||
self.samples_dropped.emit(sample_list.s)
|
||||
event.acceptProposedAction()
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
try:
|
||||
sample = SampleShortInfo.model_validate_json(text)
|
||||
self.samples_dropped.emit([sample])
|
||||
event.acceptProposedAction()
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if event.source() == self:
|
||||
drop_row = self.indexAt(event.position().toPoint()).row()
|
||||
if drop_row < 0:
|
||||
drop_row = self.count()
|
||||
|
||||
selected = self.selectedItems()
|
||||
if selected and self._item_ids:
|
||||
for item in selected:
|
||||
row = self.row(item)
|
||||
if 0 <= row < len(self._item_ids):
|
||||
item_id = self._item_ids[row]
|
||||
self.item_reordered.emit(item_id, drop_row)
|
||||
|
||||
event.acceptProposedAction()
|
||||
return
|
||||
|
||||
event.ignore()
|
||||
|
||||
def _on_delete_pressed(self) -> None:
|
||||
"""Handle delete key press."""
|
||||
selected = self.selectedItems()
|
||||
if not selected:
|
||||
return
|
||||
|
||||
item_ids = []
|
||||
for item in selected:
|
||||
row = self.row(item)
|
||||
if 0 <= row < len(self._item_ids):
|
||||
item_ids.append(self._item_ids[row])
|
||||
|
||||
if item_ids:
|
||||
self.delete_requested.emit(item_ids)
|
||||
|
||||
def _show_context_menu(self, pos) -> None:
|
||||
"""Show context menu for queue items."""
|
||||
item = self.itemAt(pos)
|
||||
if not item:
|
||||
return
|
||||
|
||||
row = self.row(item)
|
||||
if row < 0 or row >= len(self._item_ids):
|
||||
return
|
||||
|
||||
menu = QMenu(self)
|
||||
|
||||
delete_action = menu.addAction("🗑 Remove from queue")
|
||||
start_action = menu.addAction("▶ Start this item")
|
||||
|
||||
menu.addSeparator()
|
||||
move_top_action = menu.addAction("⬆ Move to top")
|
||||
move_bottom_action = menu.addAction("⬇ Move to bottom")
|
||||
|
||||
action = menu.exec_(self.mapToGlobal(pos))
|
||||
|
||||
if action == delete_action:
|
||||
self.delete_requested.emit([self._item_ids[row]])
|
||||
elif action == start_action:
|
||||
self.start_item_requested.emit(self._item_ids[row])
|
||||
elif action == move_top_action:
|
||||
self.item_reordered.emit(self._item_ids[row], 0)
|
||||
elif action == move_bottom_action:
|
||||
self.item_reordered.emit(self._item_ids[row], 999999)
|
||||
|
||||
|
||||
class WorkflowPanel(QWidget):
|
||||
"""
|
||||
Main workflow panel combining queue display and controls.
|
||||
|
||||
Supports:
|
||||
- Drag-drop samples from TELL sample panel
|
||||
- Queue management (delete, reorder, clear, add all)
|
||||
- Step-by-step guided mode (explicit start)
|
||||
- Full automation mode (explicit start)
|
||||
"""
|
||||
|
||||
# Signals for DAQ worker
|
||||
request_queue_refresh = Signal()
|
||||
request_add_sample = Signal(object) # SampleShortInfo
|
||||
request_add_samples = Signal(list) # list[SampleShortInfo]
|
||||
request_delete_item = Signal(str) # item_id
|
||||
request_move_item = Signal(str, int) # item_id, new_order_index
|
||||
request_clear_queue = Signal()
|
||||
request_start_item = Signal(str) # item_id
|
||||
request_next_step = Signal()
|
||||
request_pause = Signal()
|
||||
request_resume = Signal()
|
||||
request_abort = Signal()
|
||||
request_skip = Signal() # Skip current step
|
||||
request_skip_sample = Signal() # Skip entire current sample
|
||||
request_start_automation = Signal()
|
||||
request_stop_automation = Signal()
|
||||
request_start_guided = Signal(str) # Start guided mode with item_id
|
||||
|
||||
def __init__(self, parent: QWidget | None = None):
|
||||
super().__init__(parent)
|
||||
self._queue_items: list[QueueItem] = []
|
||||
self._all_samples: list[SampleShortInfo] = []
|
||||
self._runtime: RuntimeState | None = None
|
||||
self._control: ControlState | None = None
|
||||
self._automation_enabled = False
|
||||
self._was_paused_before_automation = False # Track pause state before automation
|
||||
self._setup_ui()
|
||||
self._setup_shortcuts()
|
||||
|
||||
def _setup_ui(self) -> None:
|
||||
layout = QVBoxLayout(self)
|
||||
layout.setContentsMargins(8, 8, 8, 8)
|
||||
layout.setSpacing(8)
|
||||
|
||||
# === Status Section ===
|
||||
status_group = QGroupBox("Current Status")
|
||||
status_layout = QVBoxLayout(status_group)
|
||||
|
||||
self._status_label = QLabel("⏹️ Idle")
|
||||
self._status_label.setStyleSheet("font-size: 14px; font-weight: bold;")
|
||||
status_layout.addWidget(self._status_label)
|
||||
|
||||
self._current_item_label = QLabel("No item running")
|
||||
status_layout.addWidget(self._current_item_label)
|
||||
|
||||
self._step_progress = StepProgressWidget()
|
||||
status_layout.addWidget(self._step_progress)
|
||||
|
||||
layout.addWidget(status_group)
|
||||
|
||||
# === Start Buttons (explicit start required) ===
|
||||
start_group = QGroupBox("Start Processing")
|
||||
start_layout = QHBoxLayout(start_group)
|
||||
|
||||
self._start_guided_btn = QPushButton("▶ Start Guided Mode")
|
||||
self._start_guided_btn.setStyleSheet("background-color: #4CAF50; color: white; font-weight: bold; padding: 8px;")
|
||||
self._start_guided_btn.setToolTip("Start processing the first pending sample in guided (step-by-step) mode")
|
||||
self._start_guided_btn.clicked.connect(self._on_start_guided)
|
||||
|
||||
self._start_auto_btn = QPushButton("▶▶ Start Automation")
|
||||
self._start_auto_btn.setStyleSheet("background-color: #2196F3; color: white; font-weight: bold; padding: 8px;")
|
||||
self._start_auto_btn.setToolTip("Start fully automated processing of all samples")
|
||||
self._start_auto_btn.clicked.connect(self._on_start_automation)
|
||||
|
||||
self._stop_btn = QPushButton("⏹ Stop")
|
||||
self._stop_btn.setStyleSheet("background-color: #9E9E9E; color: white; font-weight: bold; padding: 8px;")
|
||||
self._stop_btn.setToolTip("Stop automation mode (current step will complete)")
|
||||
self._stop_btn.clicked.connect(self._on_stop)
|
||||
self._stop_btn.setVisible(False)
|
||||
|
||||
start_layout.addWidget(self._start_guided_btn)
|
||||
start_layout.addWidget(self._start_auto_btn)
|
||||
start_layout.addWidget(self._stop_btn)
|
||||
|
||||
layout.addWidget(start_group)
|
||||
|
||||
# === Control Buttons ===
|
||||
controls_group = QGroupBox("Step Controls")
|
||||
controls_layout = QVBoxLayout(controls_group)
|
||||
|
||||
# Step controls row 1
|
||||
step_layout = QHBoxLayout()
|
||||
|
||||
self._next_btn = QPushButton("Next Step")
|
||||
self._next_btn.setStyleSheet("background-color: #4CAF50; color: white;")
|
||||
self._next_btn.setToolTip("Execute the next step (guided mode only)")
|
||||
self._next_btn.clicked.connect(self.request_next_step.emit)
|
||||
|
||||
self._skip_btn = QPushButton("Skip Step")
|
||||
self._skip_btn.setStyleSheet("background-color: #FF9800; color: white;")
|
||||
self._skip_btn.setToolTip("Skip the current step and move to the next")
|
||||
self._skip_btn.clicked.connect(self.request_skip.emit)
|
||||
|
||||
self._skip_sample_btn = QPushButton("Skip Sample")
|
||||
self._skip_sample_btn.setStyleSheet("background-color: #FF5722; color: white;")
|
||||
self._skip_sample_btn.setToolTip("Skip the entire current sample and move to the next")
|
||||
self._skip_sample_btn.clicked.connect(self.request_skip_sample.emit)
|
||||
|
||||
step_layout.addWidget(self._next_btn)
|
||||
step_layout.addWidget(self._skip_btn)
|
||||
step_layout.addWidget(self._skip_sample_btn)
|
||||
|
||||
controls_layout.addLayout(step_layout)
|
||||
|
||||
# Step controls row 2
|
||||
control_layout2 = QHBoxLayout()
|
||||
|
||||
self._pause_btn = QPushButton("⏸ Pause")
|
||||
self._pause_btn.setStyleSheet("background-color: #2196F3; color: white;")
|
||||
self._pause_btn.clicked.connect(self._on_pause_resume)
|
||||
|
||||
self._abort_btn = QPushButton("🛑 Abort")
|
||||
self._abort_btn.setStyleSheet("background-color: #F44336; color: white;")
|
||||
self._abort_btn.clicked.connect(self.request_abort.emit)
|
||||
|
||||
control_layout2.addWidget(self._pause_btn)
|
||||
control_layout2.addWidget(self._abort_btn)
|
||||
|
||||
controls_layout.addLayout(control_layout2)
|
||||
layout.addWidget(controls_group)
|
||||
|
||||
# === Queue Section ===
|
||||
queue_group = QGroupBox("Queue (drag samples here)")
|
||||
queue_layout = QVBoxLayout(queue_group)
|
||||
|
||||
self._drop_hint = QLabel("💡 Drag samples from the Sample List to add them")
|
||||
self._drop_hint.setStyleSheet("color: #666; font-style: italic;")
|
||||
queue_layout.addWidget(self._drop_hint)
|
||||
|
||||
self._queue_list = DraggableQueueListWidget()
|
||||
self._queue_list.setMinimumHeight(150)
|
||||
self._queue_list.samples_dropped.connect(self._on_samples_dropped)
|
||||
self._queue_list.delete_requested.connect(self._on_delete_requested)
|
||||
self._queue_list.item_reordered.connect(self._on_item_reordered)
|
||||
self._queue_list.start_item_requested.connect(self._on_start_specific_item)
|
||||
queue_layout.addWidget(self._queue_list)
|
||||
|
||||
# Queue management buttons
|
||||
queue_btn_layout = QHBoxLayout()
|
||||
|
||||
self._add_all_btn = QPushButton("➕ Add All")
|
||||
self._add_all_btn.clicked.connect(self._on_add_all_clicked)
|
||||
self._add_all_btn.setToolTip("Add all samples from the Sample List")
|
||||
|
||||
self._remove_selected_btn = QPushButton("🗑 Remove")
|
||||
self._remove_selected_btn.clicked.connect(self._on_remove_selected)
|
||||
|
||||
self._clear_btn = QPushButton("✖ Clear")
|
||||
self._clear_btn.clicked.connect(self._on_clear_queue)
|
||||
|
||||
self._refresh_btn = QPushButton("🔄")
|
||||
self._refresh_btn.setFixedWidth(40)
|
||||
self._refresh_btn.setToolTip("Refresh queue")
|
||||
self._refresh_btn.clicked.connect(self.request_queue_refresh.emit)
|
||||
|
||||
queue_btn_layout.addWidget(self._add_all_btn)
|
||||
queue_btn_layout.addWidget(self._remove_selected_btn)
|
||||
queue_btn_layout.addWidget(self._clear_btn)
|
||||
queue_btn_layout.addStretch()
|
||||
queue_btn_layout.addWidget(self._refresh_btn)
|
||||
queue_layout.addLayout(queue_btn_layout)
|
||||
|
||||
layout.addWidget(queue_group)
|
||||
|
||||
# Initial state
|
||||
self._update_button_states()
|
||||
|
||||
def _setup_shortcuts(self) -> None:
|
||||
self._pause_resume_shortcut = QShortcut(QKeySequence(Qt.Key.Key_Space), self)
|
||||
self._pause_resume_shortcut.setContext(Qt.ShortcutContext.WidgetWithChildrenShortcut)
|
||||
self._pause_resume_shortcut.activated.connect(self._on_pause_resume)
|
||||
|
||||
self._refresh_shortcut = QShortcut(QKeySequence("Ctrl+R"), self)
|
||||
self._refresh_shortcut.setContext(Qt.ShortcutContext.WidgetWithChildrenShortcut)
|
||||
self._refresh_shortcut.activated.connect(self.request_queue_refresh.emit)
|
||||
|
||||
self._clear_shortcut = QShortcut(QKeySequence("Ctrl+Delete"), self)
|
||||
self._clear_shortcut.setContext(Qt.ShortcutContext.WidgetWithChildrenShortcut)
|
||||
self._clear_shortcut.activated.connect(self._on_clear_queue)
|
||||
|
||||
def _on_start_guided(self) -> None:
|
||||
"""Start guided mode with the first pending sample."""
|
||||
pending = [i for i in self._queue_items if i.status == QueueItemStatus.PENDING]
|
||||
if not pending:
|
||||
QMessageBox.information(self, "No Samples", "No pending samples in the queue.")
|
||||
return
|
||||
|
||||
# Start the first pending item
|
||||
self.request_start_item.emit(pending[0].item_id)
|
||||
|
||||
def _on_start_specific_item(self, item_id: str) -> None:
|
||||
"""Start guided mode with a specific sample."""
|
||||
self.request_start_item.emit(item_id)
|
||||
|
||||
def _on_start_automation(self) -> None:
|
||||
"""Start full automation mode."""
|
||||
pending = [i for i in self._queue_items if i.status == QueueItemStatus.PENDING]
|
||||
is_running = self._runtime is not None and self._runtime.running
|
||||
|
||||
if not pending and not is_running:
|
||||
QMessageBox.information(self, "No Samples", "No pending samples in the queue.")
|
||||
return
|
||||
|
||||
# Remember if we were paused before starting automation
|
||||
self._was_paused_before_automation = (
|
||||
self._control is not None and self._control.pause_requested
|
||||
)
|
||||
|
||||
self._automation_enabled = True
|
||||
self.request_start_automation.emit()
|
||||
self._update_button_states()
|
||||
|
||||
def _on_stop(self) -> None:
|
||||
"""Stop automation mode."""
|
||||
self._automation_enabled = False
|
||||
self.request_stop_automation.emit()
|
||||
|
||||
# Restore pause state if it was paused before automation started
|
||||
if self._was_paused_before_automation:
|
||||
self.request_pause.emit()
|
||||
|
||||
self._update_button_states()
|
||||
|
||||
def _on_pause_resume(self) -> None:
|
||||
"""Toggle pause/resume."""
|
||||
if self._runtime and self._runtime.paused:
|
||||
self.request_resume.emit()
|
||||
else:
|
||||
self.request_pause.emit()
|
||||
|
||||
def _on_samples_dropped(self, samples: list[SampleShortInfo]) -> None:
|
||||
"""Handle samples dropped onto the queue."""
|
||||
logger.info(f"Adding {len(samples)} samples to workflow queue")
|
||||
for sample in samples:
|
||||
self.request_add_sample.emit(sample)
|
||||
QTimer.singleShot(500, self.request_queue_refresh.emit)
|
||||
|
||||
def _on_delete_requested(self, item_ids: list[str]) -> None:
|
||||
"""Handle delete request from queue list."""
|
||||
for item_id in item_ids:
|
||||
self.request_delete_item.emit(item_id)
|
||||
QTimer.singleShot(300, self.request_queue_refresh.emit)
|
||||
|
||||
def _on_item_reordered(self, item_id: str, new_index: int) -> None:
|
||||
"""Handle item reorder."""
|
||||
self.request_move_item.emit(item_id, new_index)
|
||||
QTimer.singleShot(300, self.request_queue_refresh.emit)
|
||||
|
||||
def _on_remove_selected(self) -> None:
|
||||
"""Remove selected items from queue."""
|
||||
selected = self._queue_list.selectedItems()
|
||||
if not selected:
|
||||
return
|
||||
|
||||
item_ids = self._queue_list._item_ids
|
||||
for item in selected:
|
||||
row = self._queue_list.row(item)
|
||||
if 0 <= row < len(item_ids):
|
||||
self.request_delete_item.emit(item_ids[row])
|
||||
|
||||
QTimer.singleShot(300, self.request_queue_refresh.emit)
|
||||
|
||||
def _on_clear_queue(self) -> None:
|
||||
"""Clear all items from queue."""
|
||||
if not self._queue_items:
|
||||
return
|
||||
|
||||
reply = QMessageBox.question(
|
||||
self,
|
||||
"Clear Queue",
|
||||
"Remove all non-running items from the queue?",
|
||||
QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No,
|
||||
)
|
||||
|
||||
if reply == QMessageBox.StandardButton.Yes:
|
||||
self.request_clear_queue.emit()
|
||||
|
||||
QTimer.singleShot(200, self._clear_local_queue)
|
||||
|
||||
QTimer.singleShot(500, self.request_queue_refresh.emit)
|
||||
|
||||
def _clear_local_queue(self) -> None:
|
||||
"""Clear local queue state (called after server clear)."""
|
||||
# Keep only running items
|
||||
self._queue_items = [i for i in self._queue_items if i.status == QueueItemStatus.RUNNING]
|
||||
|
||||
# Update GUI
|
||||
self._queue_list.clear()
|
||||
item_ids = []
|
||||
for item in self._queue_items:
|
||||
status_emoji = "▶️"
|
||||
display_text = f"{status_emoji} {item.sample_name or item.item_id}"
|
||||
list_item = QListWidgetItem(display_text)
|
||||
list_item.setBackground(QColor("#E6F3FF"))
|
||||
self._queue_list.addItem(list_item)
|
||||
item_ids.append(item.item_id)
|
||||
|
||||
self._queue_list.set_item_ids(item_ids)
|
||||
self._update_button_states()
|
||||
|
||||
def _on_add_all_clicked(self) -> None:
|
||||
"""Add all available samples to the queue."""
|
||||
if not self._all_samples:
|
||||
QMessageBox.information(
|
||||
self,
|
||||
"No Samples",
|
||||
"No samples available to add. Load samples in the Sample List first.",
|
||||
)
|
||||
return
|
||||
|
||||
reply = QMessageBox.question(
|
||||
self,
|
||||
"Add All Samples",
|
||||
f"Add all {len(self._all_samples)} samples to the workflow queue?",
|
||||
QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No,
|
||||
)
|
||||
|
||||
if reply == QMessageBox.StandardButton.Yes:
|
||||
self.request_add_samples.emit(self._all_samples)
|
||||
QTimer.singleShot(500, self.request_queue_refresh.emit)
|
||||
|
||||
def _update_button_states(self) -> None:
|
||||
"""Update button enabled/disabled states based on current state."""
|
||||
is_running = self._runtime is not None and self._runtime.running
|
||||
is_paused = self._runtime is not None and self._runtime.paused
|
||||
has_pending = any(i.status == QueueItemStatus.PENDING for i in self._queue_items)
|
||||
|
||||
self._start_guided_btn.setVisible(not is_running and not self._automation_enabled)
|
||||
|
||||
self._start_auto_btn.setVisible(not self._automation_enabled)
|
||||
self._stop_btn.setVisible(self._automation_enabled)
|
||||
|
||||
self._start_guided_btn.setEnabled(has_pending)
|
||||
self._start_auto_btn.setEnabled(has_pending or is_running)
|
||||
|
||||
self._next_btn.setVisible(not self._automation_enabled)
|
||||
self._next_btn.setEnabled(is_running and not is_paused)
|
||||
|
||||
# Skip buttons - enabled when running
|
||||
self._skip_btn.setEnabled(is_running)
|
||||
self._skip_sample_btn.setEnabled(is_running)
|
||||
self._abort_btn.setEnabled(is_running)
|
||||
|
||||
# Pause/Resume button
|
||||
if is_paused:
|
||||
self._pause_btn.setText("▶ Resume")
|
||||
self._pause_btn.setStyleSheet("background-color: #4CAF50; color: white;")
|
||||
else:
|
||||
self._pause_btn.setText("⏸ Pause")
|
||||
self._pause_btn.setStyleSheet("background-color: #2196F3; color: white;")
|
||||
self._pause_btn.setEnabled(is_running)
|
||||
|
||||
# Drop hint
|
||||
has_items = len(self._queue_items) > 0
|
||||
self._drop_hint.setVisible(not has_items)
|
||||
|
||||
# === Public Update Methods ===
|
||||
|
||||
@Slot(list)
|
||||
def update_queue(self, items: list[QueueItem]) -> None:
|
||||
"""Update queue display."""
|
||||
self._queue_items = items
|
||||
self._queue_list.clear()
|
||||
|
||||
item_ids = []
|
||||
for item in items:
|
||||
status_emoji = {
|
||||
QueueItemStatus.PENDING: "⏳",
|
||||
QueueItemStatus.RUNNING: "▶️",
|
||||
QueueItemStatus.COMPLETED: "✅",
|
||||
QueueItemStatus.FAILED: "❌",
|
||||
QueueItemStatus.ABORTED: "🛑",
|
||||
QueueItemStatus.SKIPPED: "⏭️",
|
||||
}.get(item.status, "❓")
|
||||
|
||||
display_text = f"{status_emoji} {item.sample_name or item.item_id}"
|
||||
list_item = QListWidgetItem(display_text)
|
||||
|
||||
if item.status == QueueItemStatus.RUNNING:
|
||||
list_item.setBackground(QColor("#E6F3FF"))
|
||||
elif item.status == QueueItemStatus.COMPLETED:
|
||||
list_item.setBackground(QColor("#E6FFE6"))
|
||||
elif item.status == QueueItemStatus.FAILED:
|
||||
list_item.setBackground(QColor("#FFE6E6"))
|
||||
|
||||
self._queue_list.addItem(list_item)
|
||||
item_ids.append(item.item_id)
|
||||
|
||||
self._queue_list.set_item_ids(item_ids)
|
||||
self._update_button_states()
|
||||
|
||||
@Slot(object)
|
||||
def update_runtime(self, runtime: RuntimeState) -> None:
|
||||
"""Update from runtime state."""
|
||||
self._runtime = runtime
|
||||
|
||||
if runtime.running:
|
||||
if runtime.paused:
|
||||
self._status_label.setText("⏸️ Paused")
|
||||
self._status_label.setStyleSheet(
|
||||
"font-size: 14px; font-weight: bold; color: #FF9800;"
|
||||
)
|
||||
else:
|
||||
mode_str = "Automation" if self._automation_enabled else "Guided"
|
||||
self._status_label.setText(f"▶️ Running ({mode_str})")
|
||||
self._status_label.setStyleSheet(
|
||||
"font-size: 14px; font-weight: bold; color: #4CAF50;"
|
||||
)
|
||||
|
||||
self._current_item_label.setText(
|
||||
f"Item: {runtime.current_item_id or 'Unknown'}"
|
||||
)
|
||||
else:
|
||||
self._status_label.setText("⏹️ Idle")
|
||||
self._status_label.setStyleSheet(
|
||||
"font-size: 14px; font-weight: bold; color: #757575;"
|
||||
)
|
||||
self._current_item_label.setText("No item running")
|
||||
self._step_progress.clear_steps()
|
||||
|
||||
# Reset automation enabled when nothing is running
|
||||
if self._automation_enabled and not runtime.running:
|
||||
# Check if queue has more pending items
|
||||
pending = [i for i in self._queue_items if i.status == QueueItemStatus.PENDING]
|
||||
if not pending:
|
||||
self._automation_enabled = False
|
||||
|
||||
self._update_button_states()
|
||||
|
||||
@Slot(object)
|
||||
def update_control(self, control: ControlState) -> None:
|
||||
"""Update from control state."""
|
||||
self._control = control
|
||||
self._update_button_states()
|
||||
|
||||
@Slot(object)
|
||||
def update_current_item(self, item: QueueItem) -> None:
|
||||
"""Update step progress for current item."""
|
||||
self._step_progress.set_steps(item.steps, item.current_step_index)
|
||||
|
||||
for i, step in enumerate(item.steps):
|
||||
self._step_progress.update_step(i, step.status)
|
||||
|
||||
@Slot(bool)
|
||||
def update_automation_enabled(self, enabled: bool) -> None:
|
||||
"""Update automation mode state."""
|
||||
self._automation_enabled = enabled
|
||||
self._update_button_states()
|
||||
|
||||
@Slot(list)
|
||||
def update_available_samples(self, samples: list[SampleShortInfo]) -> None:
|
||||
"""Update the cached sample list for 'Add All' functionality."""
|
||||
self._all_samples = samples
|
||||
|
||||
@Slot(str, str)
|
||||
def on_workflow_event(self, event_type: str, message: str) -> None:
|
||||
"""Handle workflow events from SSE stream."""
|
||||
logger.debug(f"Workflow event: {event_type} - {message}")
|
||||
|
||||
if event_type in ("item_started", "item_completed", "step_completed", "step_skipped", "sample_skipped"):
|
||||
self.request_queue_refresh.emit()
|
||||
|
||||
# Detect automation stop
|
||||
if event_type == "automation_stopped":
|
||||
self._automation_enabled = False
|
||||
|
||||
# Restore pause state if it was paused before automation started
|
||||
if self._was_paused_before_automation:
|
||||
self.request_pause.emit()
|
||||
self._was_paused_before_automation = False
|
||||
|
||||
self._update_button_states()
|
||||
self._title_label.setText("Automation progress")
|
||||
@@ -73,14 +73,6 @@ class DAQWorker(QObject):
|
||||
last_error_payload_changed = Signal(dict)
|
||||
last_error_payloads_changed = Signal(list)
|
||||
|
||||
# Workflow signals
|
||||
workflow_queue_loaded = Signal(list) # list[QueueItem]
|
||||
workflow_runtime_changed = Signal(object) # RuntimeState
|
||||
workflow_control_changed = Signal(object) # ControlState
|
||||
workflow_item_updated = Signal(object) # QueueItem
|
||||
workflow_automation_status = Signal(bool) # enabled
|
||||
workflow_event = Signal(object) # WorkflowEvent
|
||||
|
||||
baton_status_changed = Signal(BatonStatus)
|
||||
baton_request_result = Signal(dict)
|
||||
baton_response_result = Signal(dict)
|
||||
@@ -1751,175 +1743,6 @@ class DAQWorker(QObject):
|
||||
except Exception as e:
|
||||
logger.warning(f"Error ending session on close: {e}")
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Workflow API methods
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
@Slot()
|
||||
def workflow_load_queue(self):
|
||||
"""Load workflow queue."""
|
||||
if self.__base_url is None:
|
||||
return
|
||||
|
||||
request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/queue"))
|
||||
request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode())
|
||||
reply = self.__net_manager.get(request)
|
||||
reply.finished.connect(lambda: self._handle_workflow_queue_response(reply))
|
||||
|
||||
def _handle_workflow_queue_response(self, reply: QNetworkReply):
|
||||
try:
|
||||
response_data = self.handle_response(reply)
|
||||
data = json.loads(response_data)
|
||||
from aare.common.automation_models import QueueItem, QueueItemStatus
|
||||
items = [QueueItem.model_validate(i) for i in data.get("items", [])]
|
||||
self.workflow_queue_loaded.emit(items)
|
||||
|
||||
# Also emit the currently running item for step progress
|
||||
for item in items:
|
||||
if item.status == QueueItemStatus.RUNNING:
|
||||
self.workflow_item_updated.emit(item)
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to load workflow queue: {e}")
|
||||
|
||||
@Slot(object) # SampleShortInfo
|
||||
def workflow_add_sample(self, sample: "SampleShortInfo"):
|
||||
"""Add a sample to the workflow queue."""
|
||||
if self.__base_url is None:
|
||||
logger.info(f"POST /workflow/queue: {sample.sample_name}")
|
||||
return
|
||||
|
||||
from aare.common.automation_models import CreateQueueItemRequest
|
||||
request_data = CreateQueueItemRequest(
|
||||
sample_id=sample.db_id,
|
||||
sample_name=sample.sample_name,
|
||||
priority=int(sample.priority or 100),
|
||||
)
|
||||
|
||||
request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/queue"))
|
||||
request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode())
|
||||
request.setRawHeader(b"Content-Type", b"application/json")
|
||||
reply = self.__net_manager.post(request, QByteArray(request_data.model_dump_json().encode()))
|
||||
reply.finished.connect(lambda: self.handle_req_response(reply))
|
||||
|
||||
@Slot(list) # list[SampleShortInfo]
|
||||
def workflow_add_samples(self, samples: list):
|
||||
"""Add multiple samples to the workflow queue."""
|
||||
for sample in samples:
|
||||
self.workflow_add_sample(sample)
|
||||
|
||||
@Slot(str)
|
||||
def workflow_delete_item(self, item_id: str):
|
||||
"""Delete an item from the workflow queue."""
|
||||
if self.__base_url is None:
|
||||
logger.info(f"DELETE /workflow/queue/{item_id}")
|
||||
return
|
||||
|
||||
self.generic_delete(f"workflow/queue/{item_id}")
|
||||
|
||||
@Slot(str, int)
|
||||
def workflow_move_item(self, item_id: str, new_order_index: int):
|
||||
"""Move/reorder an item in the workflow queue."""
|
||||
if self.__base_url is None:
|
||||
logger.info(f"POST /workflow/queue/{item_id}/move order={new_order_index}")
|
||||
return
|
||||
|
||||
body = json.dumps({"new_order_index": new_order_index})
|
||||
self.generic_post(f"workflow/queue/{item_id}/move", body)
|
||||
|
||||
@Slot()
|
||||
def workflow_clear_queue(self):
|
||||
"""Clear all non-running items from the queue via server."""
|
||||
if self.__base_url is None:
|
||||
return
|
||||
|
||||
request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/queue/clear"))
|
||||
request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode())
|
||||
request.setRawHeader(b"Content-Type", b"application/json")
|
||||
reply = self.__net_manager.deleteResource(request)
|
||||
reply.finished.connect(lambda: self.handle_req_response(reply))
|
||||
|
||||
@Slot()
|
||||
def workflow_skip_sample(self):
|
||||
"""Skip the entire current sample."""
|
||||
self.generic_post("workflow/control/skip_sample")
|
||||
|
||||
def _handle_clear_queue_response(self, reply: QNetworkReply):
|
||||
try:
|
||||
response_data = self.handle_response(reply)
|
||||
data = json.loads(response_data)
|
||||
from aare.common.automation_models import QueueItem, QueueItemStatus
|
||||
items = [QueueItem.model_validate(i) for i in data.get("items", [])]
|
||||
|
||||
# Delete all non-running items
|
||||
for item in items:
|
||||
if item.status != QueueItemStatus.RUNNING:
|
||||
self.workflow_delete_item(item.item_id)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to clear workflow queue: {e}")
|
||||
|
||||
@Slot(str)
|
||||
def workflow_start_item(self, item_id: str):
|
||||
"""Start processing a queue item."""
|
||||
self.generic_post(f"workflow/start/{item_id}")
|
||||
QTimer.singleShot(500, self.workflow_load_queue)
|
||||
|
||||
@Slot()
|
||||
def workflow_next_step(self):
|
||||
"""Run next step in guided mode."""
|
||||
if self.__base_url is None:
|
||||
return
|
||||
|
||||
request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/next"))
|
||||
request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode())
|
||||
request.setRawHeader(b"Content-Type", b"application/json")
|
||||
reply = self.__net_manager.post(request, QByteArray(b""))
|
||||
reply.finished.connect(lambda: self._handle_workflow_step_response(reply))
|
||||
|
||||
def _handle_workflow_step_response(self, reply: QNetworkReply):
|
||||
try:
|
||||
response_data = self.handle_response(reply)
|
||||
data = json.loads(response_data)
|
||||
if data.get("item"):
|
||||
from aare.common.automation_models import QueueItem
|
||||
item = QueueItem.model_validate(data["item"])
|
||||
self.workflow_item_updated.emit(item)
|
||||
# Also refresh queue
|
||||
self.workflow_load_queue()
|
||||
except Exception as e:
|
||||
logger.error(f"Workflow step error: {e}")
|
||||
self.http_error.emit(str(e))
|
||||
|
||||
@Slot()
|
||||
def workflow_pause(self):
|
||||
"""Request pause."""
|
||||
self.generic_post("workflow/control/pause")
|
||||
|
||||
@Slot()
|
||||
def workflow_resume(self):
|
||||
"""Request resume."""
|
||||
self.generic_post("workflow/control/resume")
|
||||
|
||||
@Slot()
|
||||
def workflow_abort(self):
|
||||
"""Request abort."""
|
||||
self.generic_post("workflow/control/abort")
|
||||
|
||||
@Slot()
|
||||
def workflow_skip(self):
|
||||
"""Request skip."""
|
||||
self.generic_post("workflow/control/skip")
|
||||
|
||||
@Slot()
|
||||
def workflow_start_automation(self):
|
||||
"""Start automation mode."""
|
||||
self.generic_post("workflow/automation/start")
|
||||
|
||||
@Slot()
|
||||
def workflow_stop_automation(self):
|
||||
"""Stop automation mode."""
|
||||
self.generic_post("workflow/automation/stop")
|
||||
|
||||
def _handle_gui_sessions_response(self, reply: QNetworkReply):
|
||||
try:
|
||||
response_data = self.handle_response(reply)
|
||||
|
||||
@@ -1,175 +0,0 @@
|
||||
"""
|
||||
SSE client for workflow events.
|
||||
|
||||
Connects to the /workflow/sse endpoint and emits signals for:
|
||||
- Runtime state changes
|
||||
- Control state changes
|
||||
- Individual workflow events
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
from PySide6.QtCore import QObject, Signal, Slot, QTimer
|
||||
from PySide6.QtNetwork import QNetworkAccessManager, QNetworkRequest, QNetworkReply
|
||||
from PySide6.QtCore import QUrl, QByteArray
|
||||
|
||||
from aare.common.automation_models import RuntimeState, ControlState, WorkflowEvent
|
||||
from aare.common.logger_config import setup_logger
|
||||
|
||||
logger = setup_logger("aareGUI")
|
||||
|
||||
|
||||
class WorkflowSSEClient(QObject):
|
||||
"""
|
||||
SSE client that subscribes to workflow events.
|
||||
|
||||
Emits signals when state changes are received.
|
||||
"""
|
||||
|
||||
# Signals
|
||||
runtime_changed = Signal(object) # RuntimeState
|
||||
control_changed = Signal(object) # ControlState
|
||||
workflow_event = Signal(object) # WorkflowEvent
|
||||
connected = Signal()
|
||||
disconnected = Signal()
|
||||
error = Signal(str)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
token: str,
|
||||
parent: QObject | None = None,
|
||||
):
|
||||
super().__init__(parent)
|
||||
|
||||
self._base_url = base_url
|
||||
self._token = token
|
||||
self._manager = QNetworkAccessManager(self)
|
||||
self._reply: QNetworkReply | None = None
|
||||
self._buffer = ""
|
||||
|
||||
# Reconnection
|
||||
self._reconnect_timer = QTimer(self)
|
||||
self._reconnect_timer.setInterval(5000) # 5 seconds
|
||||
self._reconnect_timer.timeout.connect(self.connect)
|
||||
self._should_reconnect = False
|
||||
|
||||
def connect(self) -> None:
|
||||
"""Start SSE connection."""
|
||||
if self._reply is not None:
|
||||
return # Already connected
|
||||
|
||||
self._should_reconnect = True
|
||||
|
||||
url = QUrl(f"{self._base_url}/workflow/sse")
|
||||
|
||||
request = QNetworkRequest(url)
|
||||
request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode())
|
||||
request.setRawHeader(b"Accept", b"text/event-stream")
|
||||
request.setRawHeader(b"Cache-Control", b"no-cache")
|
||||
|
||||
self._reply = self._manager.get(request)
|
||||
self._reply.readyRead.connect(self._on_data_ready)
|
||||
self._reply.finished.connect(self._on_finished)
|
||||
self._reply.errorOccurred.connect(self._on_error)
|
||||
|
||||
self._reconnect_timer.stop()
|
||||
logger.debug("Workflow SSE: connecting...")
|
||||
|
||||
def disconnect(self) -> None:
|
||||
"""Stop SSE connection."""
|
||||
self._should_reconnect = False
|
||||
self._reconnect_timer.stop()
|
||||
|
||||
if self._reply is not None:
|
||||
try:
|
||||
self._reply.abort()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
self._reply.deleteLater()
|
||||
except Exception:
|
||||
pass
|
||||
self._reply = None
|
||||
|
||||
self.disconnected.emit()
|
||||
|
||||
@Slot()
|
||||
def _on_data_ready(self) -> None:
|
||||
"""Handle incoming SSE data."""
|
||||
if self._reply is None:
|
||||
return
|
||||
|
||||
try:
|
||||
data = self._reply.readAll().data().decode("utf-8")
|
||||
self._buffer += data
|
||||
|
||||
# Process complete events (separated by double newlines)
|
||||
while "\n\n" in self._buffer:
|
||||
event_data, self._buffer = self._buffer.split("\n\n", 1)
|
||||
self._parse_event(event_data)
|
||||
except Exception as e:
|
||||
logger.warning(f"Workflow SSE data read error: {e}")
|
||||
|
||||
def _parse_event(self, event_data: str) -> None:
|
||||
"""Parse a single SSE event."""
|
||||
event_type = "message"
|
||||
data_lines = []
|
||||
|
||||
for line in event_data.split("\n"):
|
||||
if line.startswith("event:"):
|
||||
event_type = line[6:].strip()
|
||||
elif line.startswith("data:"):
|
||||
data_lines.append(line[5:].strip())
|
||||
|
||||
if not data_lines:
|
||||
return
|
||||
|
||||
data_str = "\n".join(data_lines)
|
||||
|
||||
try:
|
||||
if event_type == "runtime":
|
||||
runtime = RuntimeState.model_validate_json(data_str)
|
||||
self.runtime_changed.emit(runtime)
|
||||
elif event_type == "control":
|
||||
control = ControlState.model_validate_json(data_str)
|
||||
self.control_changed.emit(control)
|
||||
elif event_type == "workflow_event":
|
||||
event = WorkflowEvent.model_validate_json(data_str)
|
||||
self.workflow_event.emit(event)
|
||||
except Exception as e:
|
||||
logger.warning(f"Workflow SSE: failed to parse {event_type}: {e}")
|
||||
|
||||
@Slot()
|
||||
def _on_finished(self) -> None:
|
||||
"""Handle connection finished."""
|
||||
if self._reply is not None:
|
||||
try:
|
||||
self._reply.deleteLater()
|
||||
except Exception:
|
||||
pass
|
||||
self._reply = None
|
||||
|
||||
self._buffer = ""
|
||||
self.disconnected.emit()
|
||||
|
||||
# Reconnect if desired
|
||||
if self._should_reconnect:
|
||||
logger.debug("Workflow SSE: disconnected, will reconnect...")
|
||||
self._reconnect_timer.start()
|
||||
|
||||
@Slot(QNetworkReply.NetworkError)
|
||||
def _on_error(self, error: QNetworkReply.NetworkError) -> None:
|
||||
"""Handle connection error."""
|
||||
error_msg = ""
|
||||
if self._reply is not None:
|
||||
try:
|
||||
error_msg = self._reply.errorString()
|
||||
except Exception:
|
||||
error_msg = str(error)
|
||||
else:
|
||||
error_msg = str(error)
|
||||
logger.warning(f"Workflow SSE error: {error_msg}")
|
||||
self.error.emit(error_msg)
|
||||
@@ -1,84 +0,0 @@
|
||||
import pytest
|
||||
from unittest.mock import MagicMock
|
||||
import json
|
||||
from aare.common.automation_queue_manager import WorkflowRedisManager
|
||||
from aare.common.automation_models import (
|
||||
QueueItem,
|
||||
WorkflowEvent,
|
||||
RuntimeState,
|
||||
QueueItemStatus,
|
||||
WorkflowStepRecord,
|
||||
WorkflowStateKind,
|
||||
)
|
||||
|
||||
@pytest.fixture
|
||||
def mock_redis():
|
||||
mock = MagicMock()
|
||||
# mock pipeline
|
||||
pipeline = MagicMock()
|
||||
mock.pipeline.return_value = pipeline
|
||||
pipeline.execute.return_value = []
|
||||
return mock
|
||||
|
||||
@pytest.fixture
|
||||
def manager(mock_redis):
|
||||
return WorkflowRedisManager(mock_redis, beamline="x10sa")
|
||||
|
||||
def test_keys(manager):
|
||||
assert manager._key("test") == "x10sa:workflow:test"
|
||||
assert manager._item_key("123") == "x10sa:workflow:item:123"
|
||||
assert manager._queue_key() == "x10sa:workflow:queue"
|
||||
|
||||
def test_create_item(manager, mock_redis):
|
||||
item = QueueItem(item_id="item1", beamline="x10sa", owner_pgroup="p12345", samples=[], status=QueueItemStatus.PENDING, order_index=1, steps=[])
|
||||
# Mock _next_item_id
|
||||
mock_redis.incr.return_value = 1
|
||||
mock_redis.zcard.return_value = 0
|
||||
|
||||
manager.create_item(item)
|
||||
|
||||
# Check if item was set in pipeline
|
||||
pipeline = mock_redis.pipeline.return_value
|
||||
pipeline.set.assert_called()
|
||||
pipeline.zadd.assert_called()
|
||||
|
||||
def test_get_item(manager, mock_redis):
|
||||
item = QueueItem(item_id="item1", beamline="x10sa", owner_pgroup="p12345", samples=[], status=QueueItemStatus.PENDING, order_index=1, steps=[])
|
||||
mock_redis.get.return_value = item.model_dump_json()
|
||||
|
||||
ret_item = manager.get_item("item1")
|
||||
assert ret_item.item_id == "item1"
|
||||
mock_redis.get.assert_called_with("x10sa:workflow:item:item1")
|
||||
|
||||
def test_delete_item(manager, mock_redis):
|
||||
manager.delete_item("item1")
|
||||
pipeline = mock_redis.pipeline.return_value
|
||||
pipeline.delete.assert_called_with("x10sa:workflow:item:item1")
|
||||
pipeline.zrem.assert_called_with("x10sa:workflow:queue", "item1")
|
||||
|
||||
def test_get_runtime(manager, mock_redis):
|
||||
runtime = RuntimeState(running=False)
|
||||
mock_redis.get.return_value = runtime.model_dump_json()
|
||||
|
||||
ret_runtime = manager.get_runtime()
|
||||
assert ret_runtime.running == False
|
||||
|
||||
def test_patch_runtime(manager, mock_redis):
|
||||
runtime = RuntimeState(running=False)
|
||||
mock_redis.get.return_value = runtime.model_dump_json()
|
||||
|
||||
manager.patch_runtime({"running": True})
|
||||
mock_redis.set.assert_called()
|
||||
# verify the content of set
|
||||
args, kwargs = mock_redis.set.call_args
|
||||
sent_data = json.loads(args[1])
|
||||
assert sent_data["running"] == True
|
||||
|
||||
def test_append_event(manager, mock_redis):
|
||||
event = WorkflowEvent(event_type="info", message="test event", beamline="x10sa")
|
||||
manager.append_event(event)
|
||||
mock_redis.xadd.assert_called()
|
||||
# verify args
|
||||
args, kwargs = mock_redis.xadd.call_args
|
||||
assert args[0] == "x10sa:workflow:events"
|
||||
assert "json" in args[1]
|
||||
@@ -1,105 +0,0 @@
|
||||
import pytest
|
||||
from unittest.mock import MagicMock, patch
|
||||
from aare.common.automation_workflow import (
|
||||
STATE_REGISTRY, HANDLER_REGISTRY, SIMULATED_HANDLER_REGISTRY,
|
||||
get_state_definition, get_allowed_next_states, can_transition,
|
||||
MountHandler, LoopCentreHandler, RasterHandler, DataCollectionHandler,
|
||||
WorkflowRunner, WorkflowContext, WorkflowMode, WorkflowStateKind, StepStatus
|
||||
)
|
||||
|
||||
def test_get_state_definition():
|
||||
defn = get_state_definition(WorkflowStateKind.MOUNT)
|
||||
assert defn.kind == WorkflowStateKind.MOUNT
|
||||
|
||||
with pytest.raises(KeyError, match="Unknown workflow state"):
|
||||
get_state_definition("invalid_state")
|
||||
|
||||
def test_get_allowed_next_states():
|
||||
# Mount -> Loop Centre
|
||||
next_states = get_allowed_next_states(WorkflowStateKind.MOUNT, mode=WorkflowMode.AUTOMATION)
|
||||
assert WorkflowStateKind.LOOP_CENTRE in next_states
|
||||
|
||||
# Loop Centre -> Raster or Data Collection (depending on mode)
|
||||
next_states_auto = get_allowed_next_states(WorkflowStateKind.LOOP_CENTRE, mode=WorkflowMode.AUTOMATION)
|
||||
assert WorkflowStateKind.RASTER in next_states_auto
|
||||
assert WorkflowStateKind.DATA_COLLECTION not in next_states_auto
|
||||
|
||||
next_states_manual = get_allowed_next_states(WorkflowStateKind.LOOP_CENTRE, mode=WorkflowMode.GUIDED_MANUAL)
|
||||
assert WorkflowStateKind.RASTER in next_states_manual
|
||||
assert WorkflowStateKind.DATA_COLLECTION in next_states_manual
|
||||
|
||||
def test_can_transition():
|
||||
assert can_transition(WorkflowStateKind.MOUNT, WorkflowStateKind.LOOP_CENTRE) is True
|
||||
assert can_transition(WorkflowStateKind.MOUNT, WorkflowStateKind.RASTER) is False
|
||||
|
||||
@pytest.fixture
|
||||
def context():
|
||||
return WorkflowContext(
|
||||
mode=WorkflowMode.AUTOMATION,
|
||||
queue_id="q1",
|
||||
item_id="item1",
|
||||
sample_id=123
|
||||
)
|
||||
|
||||
def test_handlers_execute(context):
|
||||
handlers = [
|
||||
MountHandler(),
|
||||
LoopCentreHandler(),
|
||||
RasterHandler(),
|
||||
DataCollectionHandler()
|
||||
]
|
||||
|
||||
for handler in handlers:
|
||||
res = handler.execute(context)
|
||||
assert res.status == StepStatus.SUCCESS
|
||||
assert context.current_state == handler.state_kind
|
||||
|
||||
def test_handler_abort(context):
|
||||
context.abort_requested = True
|
||||
handler = MountHandler()
|
||||
with pytest.raises(RuntimeError, match="Abort requested"):
|
||||
handler.execute(context)
|
||||
|
||||
def test_workflow_runner(context):
|
||||
runner = WorkflowRunner()
|
||||
|
||||
# Start at None, go to Mount
|
||||
res = runner.run_state(context, WorkflowStateKind.MOUNT)
|
||||
assert res.state == WorkflowStateKind.MOUNT
|
||||
assert context.current_step_index == 1
|
||||
|
||||
# Mount -> Loop Centre
|
||||
res = runner.run_state(context, WorkflowStateKind.LOOP_CENTRE)
|
||||
assert res.state == WorkflowStateKind.LOOP_CENTRE
|
||||
|
||||
# Invalid transition: Loop Centre -> Mount (in reverse)
|
||||
with pytest.raises(RuntimeError, match="Transition not allowed"):
|
||||
runner.run_state(context, WorkflowStateKind.MOUNT)
|
||||
|
||||
def test_state_handler_helpers(context):
|
||||
handler = MountHandler()
|
||||
assert handler.definition().kind == WorkflowStateKind.MOUNT
|
||||
|
||||
context.current_state = None
|
||||
assert handler.can_run(context) is True
|
||||
context.current_state = WorkflowStateKind.MOUNT
|
||||
assert handler.can_run(context) is True
|
||||
context.current_state = WorkflowStateKind.RASTER
|
||||
assert handler.can_run(context) is False
|
||||
|
||||
def test_workflow_runner_error_cases():
|
||||
runner = WorkflowRunner()
|
||||
with pytest.raises(KeyError, match="No handler registered"):
|
||||
runner.get_handler("non_existent_state")
|
||||
|
||||
@patch("time.sleep", return_value=None)
|
||||
def test_all_simulated_handlers(mock_sleep, context):
|
||||
for kind, handler in SIMULATED_HANDLER_REGISTRY.items():
|
||||
context.abort_requested = False
|
||||
res = handler.execute(context)
|
||||
assert res.status == StepStatus.SUCCESS
|
||||
assert "SIMULATION" in res.message
|
||||
|
||||
context.abort_requested = True
|
||||
with pytest.raises(RuntimeError, match="Abort requested"):
|
||||
handler.execute(context)
|
||||
@@ -1,139 +0,0 @@
|
||||
import pytest
|
||||
import asyncio
|
||||
from unittest.mock import MagicMock, patch
|
||||
from aare.daq.automation_runner import PersistentWorkflowRunner, AutomationLoop
|
||||
from aare.common.automation_models import (
|
||||
QueueItem, QueueItemStatus, WorkflowContext, WorkflowMode,
|
||||
WorkflowStateKind, StateResult, StepStatus, ControlState, RuntimeState,
|
||||
StepState
|
||||
)
|
||||
|
||||
@pytest.fixture
|
||||
def mock_redis():
|
||||
m = MagicMock()
|
||||
m._bl = "x10sa"
|
||||
m.get_runtime.return_value = RuntimeState()
|
||||
m.get_control.return_value = ControlState()
|
||||
return m
|
||||
|
||||
@pytest.fixture
|
||||
def runner(mock_redis):
|
||||
return PersistentWorkflowRunner(redis_manager=mock_redis)
|
||||
|
||||
def test_runner_init(runner, mock_redis):
|
||||
assert runner.beamline == "x10sa"
|
||||
assert runner._redis == mock_redis
|
||||
|
||||
def test_start_item(runner, mock_redis):
|
||||
item = QueueItem(
|
||||
item_id="item1", beamline="x10sa", sample_id=1, sample_name="S1",
|
||||
owner_pgroup="p1", created_by="u", steps=[]
|
||||
)
|
||||
mock_redis.get_item.return_value = item
|
||||
mock_redis.update_item.return_value = item
|
||||
|
||||
started_item = runner.start_item("item1")
|
||||
|
||||
assert started_item == item
|
||||
mock_redis.update_item.assert_called_once()
|
||||
mock_redis.patch_runtime.assert_called_once()
|
||||
mock_redis.append_event.assert_called_once()
|
||||
|
||||
def test_run_current_step_success(runner, mock_redis):
|
||||
step = MagicMock()
|
||||
step.kind = "mount"
|
||||
item = MagicMock(item_id="item1", steps=[step], current_step_index=0)
|
||||
mock_redis.get_item.return_value = item
|
||||
|
||||
handler = MagicMock()
|
||||
handler.execute.return_value = StateResult(
|
||||
state=WorkflowStateKind.MOUNT, status=StepStatus.SUCCESS, message="OK"
|
||||
)
|
||||
runner._handlers = {WorkflowStateKind.MOUNT: handler}
|
||||
|
||||
context = WorkflowContext(
|
||||
mode=WorkflowMode.AUTOMATION, queue_id="q", item_id="item1", sample_id=1
|
||||
)
|
||||
|
||||
res = runner.run_current_step(context)
|
||||
|
||||
assert res.status == StepStatus.SUCCESS
|
||||
mock_redis.update_step.assert_called()
|
||||
mock_redis.update_item.assert_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_automation_loop_start_stop(runner, mock_redis, mocker):
|
||||
loop = AutomationLoop(runner, mock_redis)
|
||||
|
||||
# Patch _run_loop to be a non-coroutine to avoid "never awaited" warning
|
||||
mocker.patch.object(loop, "_run_loop", return_value=None)
|
||||
|
||||
with patch("asyncio.create_task") as mock_task:
|
||||
mock_task.return_value = MagicMock()
|
||||
loop.start()
|
||||
assert loop.is_enabled is True
|
||||
mock_task.assert_called_once()
|
||||
|
||||
loop.stop()
|
||||
assert loop.is_enabled is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_automation_loop_tick_no_items(runner, mock_redis):
|
||||
loop = AutomationLoop(runner, mock_redis)
|
||||
mock_redis.get_runtime.return_value = RuntimeState(running=False)
|
||||
mock_redis.get_next_pending_item.return_value = None
|
||||
runner.check_control = MagicMock(return_value=ControlState())
|
||||
|
||||
# Ensure runner.start_item is a regular Mock, not AsyncMock
|
||||
runner.start_item = MagicMock()
|
||||
|
||||
await loop._tick()
|
||||
|
||||
mock_redis.get_next_pending_item.assert_called_once()
|
||||
mock_redis.patch_runtime.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_automation_loop_tick_process_step(runner, mock_redis):
|
||||
loop = AutomationLoop(runner, mock_redis)
|
||||
# 1. Mock runtime and control so it thinks an item is running
|
||||
mock_redis.get_runtime.return_value = RuntimeState(
|
||||
running=True, current_item_id="item1", current_state="mount"
|
||||
)
|
||||
runner.check_control = MagicMock(return_value=ControlState())
|
||||
|
||||
# 2. Mock item and its steps
|
||||
step = MagicMock()
|
||||
step.kind = "loop_centre"
|
||||
item = MagicMock(item_id="item1", steps=[MagicMock(kind="mount"), step], current_step_index=1, sample_id=123)
|
||||
mock_redis.get_item.return_value = item
|
||||
|
||||
# Ensure runner.run_current_step is a regular Mock
|
||||
runner.run_current_step = MagicMock(return_value=StateResult(
|
||||
state=WorkflowStateKind.LOOP_CENTRE, status=StepStatus.SUCCESS
|
||||
))
|
||||
|
||||
await loop._tick()
|
||||
|
||||
runner.run_current_step.assert_called_once()
|
||||
assert runner.run_current_step.call_args[0][0].item_id == "item1"
|
||||
|
||||
def test_runner_error_handling(runner, mock_redis):
|
||||
step = MagicMock()
|
||||
step.kind = "mount"
|
||||
item = MagicMock(item_id="item1", steps=[step], current_step_index=0)
|
||||
mock_redis.get_item.return_value = item
|
||||
|
||||
handler = MagicMock()
|
||||
handler.execute.side_effect = Exception("Hardware failure")
|
||||
runner._handlers = {WorkflowStateKind.MOUNT: handler}
|
||||
|
||||
context = WorkflowContext(
|
||||
mode=WorkflowMode.AUTOMATION, queue_id="q", item_id="item1", sample_id=1
|
||||
)
|
||||
|
||||
with pytest.raises(Exception, match="Hardware failure"):
|
||||
runner.run_current_step(context)
|
||||
|
||||
mock_redis.update_step.assert_called_with(
|
||||
"item1", 0, status="failed", error_detail="Hardware failure"
|
||||
)
|
||||
@@ -17,7 +17,6 @@ def test_main_window_init(qtbot, mock_ui_state):
|
||||
patch("aare.gui.main_window.PredictionSubscriber"), \
|
||||
patch("aare.gui.main_window.VideoThread"), \
|
||||
patch("aare.gui.main_window.JFJochDBusClient"), \
|
||||
patch("aare.gui.main_window.WorkflowSSEClient"), \
|
||||
patch("aare.gui.main_window.jwt.decode") as mock_jwt:
|
||||
|
||||
mock_get.return_value.status_code = 200
|
||||
@@ -51,7 +50,6 @@ def test_main_window_mount_view(qtbot, mock_ui_state):
|
||||
patch("aare.gui.main_window.PredictionSubscriber"), \
|
||||
patch("aare.gui.main_window.VideoThread"), \
|
||||
patch("aare.gui.main_window.JFJochDBusClient"), \
|
||||
patch("aare.gui.main_window.WorkflowSSEClient"), \
|
||||
patch("aare.gui.main_window.jwt.decode") as mock_jwt:
|
||||
|
||||
mock_jwt.return_value = {"sub": "testuser", "staff": True, "pgroups": ["p123"], "session": 15}
|
||||
@@ -80,7 +78,6 @@ def test_mark_user_interaction_reports_backend(qtbot, mock_ui_state):
|
||||
patch("aare.gui.main_window.PredictionSubscriber"), \
|
||||
patch("aare.gui.main_window.VideoThread"), \
|
||||
patch("aare.gui.main_window.JFJochDBusClient"), \
|
||||
patch("aare.gui.main_window.WorkflowSSEClient"), \
|
||||
patch("aare.gui.main_window.jwt.decode") as mock_jwt:
|
||||
|
||||
mock_jwt.return_value = {"sub": "testuser", "staff": True, "pgroups": ["p123"], "session": 15}
|
||||
@@ -110,7 +107,6 @@ def test_automation_activity_refreshes_idle_timestamp(qtbot, mock_ui_state):
|
||||
patch("aare.gui.main_window.PredictionSubscriber"), \
|
||||
patch("aare.gui.main_window.VideoThread"), \
|
||||
patch("aare.gui.main_window.JFJochDBusClient"), \
|
||||
patch("aare.gui.main_window.WorkflowSSEClient"), \
|
||||
patch("aare.gui.main_window.jwt.decode") as mock_jwt:
|
||||
|
||||
from aare.common.models import (
|
||||
@@ -199,7 +195,6 @@ def test_idle_timeout_closes_when_inactive_and_not_running(qtbot, mock_ui_state)
|
||||
patch("aare.gui.main_window.PredictionSubscriber"), \
|
||||
patch("aare.gui.main_window.VideoThread"), \
|
||||
patch("aare.gui.main_window.JFJochDBusClient"), \
|
||||
patch("aare.gui.main_window.WorkflowSSEClient"), \
|
||||
patch("aare.gui.main_window.jwt.decode") as mock_jwt:
|
||||
|
||||
mock_jwt.return_value = {"sub": "testuser", "staff": True, "pgroups": ["p123"], "session": 15}
|
||||
@@ -233,7 +228,6 @@ def test_idle_timeout_does_not_close_while_automation_active(qtbot, mock_ui_stat
|
||||
patch("aare.gui.main_window.PredictionSubscriber"), \
|
||||
patch("aare.gui.main_window.VideoThread"), \
|
||||
patch("aare.gui.main_window.JFJochDBusClient"), \
|
||||
patch("aare.gui.main_window.WorkflowSSEClient"), \
|
||||
patch("aare.gui.main_window.jwt.decode") as mock_jwt:
|
||||
|
||||
mock_jwt.return_value = {"sub": "testuser", "staff": True, "pgroups": ["p123"], "session": 15}
|
||||
|
||||
Reference in New Issue
Block a user