Automation 2.0: WIP added endpoints to backend, frontend connected and panel made, now debugging
This commit is contained in:
@@ -3,7 +3,10 @@ from __future__ import annotations
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
import time
|
||||
import uuid
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
class WorkflowMode(str, Enum):
|
||||
FLEXIBLE_MANUAL = "flexible_manual"
|
||||
@@ -64,93 +67,116 @@ class WorkflowContext:
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
STATE_REGISTRY: dict[WorkflowStateKind, StateDefinition] = {
|
||||
WorkflowStateKind.MOUNT: StateDefinition(
|
||||
kind=WorkflowStateKind.MOUNT,
|
||||
description="Mount the sample",
|
||||
transitions=(
|
||||
TransitionRule(
|
||||
to_state=WorkflowStateKind.LOOP_CENTRE,
|
||||
allowed_modes=frozenset({
|
||||
WorkflowMode.FLEXIBLE_MANUAL,
|
||||
WorkflowMode.GUIDED_MANUAL,
|
||||
WorkflowMode.AUTOMATION,
|
||||
}),
|
||||
),
|
||||
),
|
||||
),
|
||||
WorkflowStateKind.LOOP_CENTRE: StateDefinition(
|
||||
kind=WorkflowStateKind.LOOP_CENTRE,
|
||||
description="Centre the loop",
|
||||
transitions=(
|
||||
TransitionRule(
|
||||
to_state=WorkflowStateKind.RASTER,
|
||||
allowed_modes=frozenset({
|
||||
WorkflowMode.FLEXIBLE_MANUAL,
|
||||
WorkflowMode.GUIDED_MANUAL,
|
||||
WorkflowMode.AUTOMATION,
|
||||
}),
|
||||
),
|
||||
TransitionRule(
|
||||
to_state=WorkflowStateKind.DATA_COLLECTION,
|
||||
allowed_modes=frozenset({
|
||||
WorkflowMode.FLEXIBLE_MANUAL,
|
||||
WorkflowMode.GUIDED_MANUAL,
|
||||
}),
|
||||
optional=True,
|
||||
),
|
||||
),
|
||||
),
|
||||
WorkflowStateKind.RASTER: StateDefinition(
|
||||
kind=WorkflowStateKind.RASTER,
|
||||
description="Run raster scan",
|
||||
transitions=(
|
||||
TransitionRule(
|
||||
to_state=WorkflowStateKind.DATA_COLLECTION,
|
||||
allowed_modes=frozenset({
|
||||
WorkflowMode.FLEXIBLE_MANUAL,
|
||||
WorkflowMode.GUIDED_MANUAL,
|
||||
WorkflowMode.AUTOMATION,
|
||||
}),
|
||||
),
|
||||
),
|
||||
),
|
||||
WorkflowStateKind.DATA_COLLECTION: StateDefinition(
|
||||
kind=WorkflowStateKind.DATA_COLLECTION,
|
||||
description="Collect diffraction data",
|
||||
transitions=(),
|
||||
),
|
||||
}
|
||||
class QueueItemStatus(str, Enum):
|
||||
PENDING = "pending"
|
||||
RUNNING = "running"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
SKIPPED = "skipped"
|
||||
ABORTED = "aborted"
|
||||
|
||||
|
||||
def get_state_definition(kind: WorkflowStateKind) -> StateDefinition:
|
||||
try:
|
||||
return STATE_REGISTRY[kind]
|
||||
except KeyError as exc:
|
||||
raise KeyError(f"Unknown workflow state: {kind}") from exc
|
||||
class WorkflowStepRecord(BaseModel):
|
||||
kind: str
|
||||
status: str = "pending"
|
||||
message: str = ""
|
||||
started_at: float | None = None
|
||||
completed_at: float | None = None
|
||||
error_detail: str | None = None
|
||||
|
||||
|
||||
def get_allowed_next_states(
|
||||
kind: WorkflowStateKind,
|
||||
mode: WorkflowMode | None = None,
|
||||
) -> list[WorkflowStateKind]:
|
||||
definition = get_state_definition(kind)
|
||||
out: list[WorkflowStateKind] = []
|
||||
|
||||
for transition in definition.transitions:
|
||||
if mode is None:
|
||||
out.append(transition.to_state)
|
||||
continue
|
||||
|
||||
if not transition.allowed_modes or mode in transition.allowed_modes:
|
||||
out.append(transition.to_state)
|
||||
|
||||
return out
|
||||
class QueueItem(BaseModel):
|
||||
item_id: str
|
||||
beamline: str
|
||||
sample_id: int | None = None
|
||||
sample_name: str = ""
|
||||
owner_pgroup: str = ""
|
||||
created_by: str = ""
|
||||
created_at: float = Field(default_factory=time.time)
|
||||
priority: int = 100
|
||||
order_index: int = 0
|
||||
status: QueueItemStatus = QueueItemStatus.PENDING
|
||||
steps: list[WorkflowStepRecord] = Field(default_factory=list)
|
||||
current_step_index: int = 0
|
||||
recipe: dict[str, Any] = Field(default_factory=dict)
|
||||
metadata: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
def can_transition(
|
||||
from_state: WorkflowStateKind,
|
||||
to_state: WorkflowStateKind,
|
||||
mode: WorkflowMode | None = None,
|
||||
) -> bool:
|
||||
return to_state in get_allowed_next_states(from_state, mode=mode)
|
||||
class RuntimeState(BaseModel):
|
||||
running: bool = False
|
||||
paused: bool = False
|
||||
current_queue_id: str = ""
|
||||
current_item_id: str | None = None
|
||||
current_state: str | None = None
|
||||
current_step_index: int = 0
|
||||
last_error: str | None = None
|
||||
last_update: float = Field(default_factory=time.time)
|
||||
|
||||
|
||||
class ControlState(BaseModel):
|
||||
pause_requested: bool = False
|
||||
resume_requested: bool = False
|
||||
abort_requested: bool = False
|
||||
skip_requested: bool = False
|
||||
next_sample_requested: bool = False
|
||||
requested_by: str | None = None
|
||||
requested_at: float | None = None
|
||||
|
||||
|
||||
class WorkflowEvent(BaseModel):
|
||||
event_id: str = Field(default_factory=lambda: str(uuid.uuid4()))
|
||||
beamline: str = ""
|
||||
item_id: str | None = None
|
||||
step: str | None = None
|
||||
event_type: str = ""
|
||||
timestamp: float = Field(default_factory=time.time)
|
||||
actor: str = ""
|
||||
message: str = ""
|
||||
payload: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class CreateQueueItemRequest(BaseModel):
|
||||
sample_id: int | None = None
|
||||
sample_name: str = ""
|
||||
priority: int = 100
|
||||
recipe: dict = Field(default_factory=dict)
|
||||
steps: list[str] | None = None # If None, use default steps
|
||||
|
||||
|
||||
class MoveItemRequest(BaseModel):
|
||||
new_order_index: int
|
||||
|
||||
|
||||
class QueueListResponse(BaseModel):
|
||||
items: list[QueueItem]
|
||||
total: int
|
||||
|
||||
|
||||
class RuntimeResponse(BaseModel):
|
||||
runtime: RuntimeState
|
||||
control: ControlState
|
||||
|
||||
|
||||
class ControlActionResponse(BaseModel):
|
||||
ok: bool
|
||||
control: ControlState
|
||||
message: str = ""
|
||||
|
||||
|
||||
class StepActionResponse(BaseModel):
|
||||
ok: bool
|
||||
item: QueueItem | None = None
|
||||
step: str | None = None
|
||||
status: str = ""
|
||||
message: str = ""
|
||||
|
||||
|
||||
class EventListResponse(BaseModel):
|
||||
events: list[WorkflowEvent]
|
||||
|
||||
|
||||
class AutomationStatusResponse(BaseModel):
|
||||
enabled: bool
|
||||
running: bool
|
||||
runtime: RuntimeState
|
||||
control: ControlState
|
||||
@@ -0,0 +1,240 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
import redis
|
||||
|
||||
from aare.common.automation_models import (
|
||||
QueueItem,
|
||||
WorkflowEvent,
|
||||
ControlState,
|
||||
RuntimeState,
|
||||
QueueItemStatus,
|
||||
WorkflowStepRecord,
|
||||
WorkflowStateKind,
|
||||
)
|
||||
|
||||
|
||||
def build_default_steps() -> list[WorkflowStepRecord]:
|
||||
return [
|
||||
WorkflowStepRecord(kind=WorkflowStateKind.MOUNT.value),
|
||||
WorkflowStepRecord(kind=WorkflowStateKind.LOOP_CENTRE.value),
|
||||
WorkflowStepRecord(kind=WorkflowStateKind.RASTER.value),
|
||||
WorkflowStepRecord(kind=WorkflowStateKind.DATA_COLLECTION.value),
|
||||
]
|
||||
|
||||
|
||||
class WorkflowRedisManager:
|
||||
def __init__(self, client: redis.Redis, beamline: str):
|
||||
self._client = client
|
||||
self._bl = beamline.lower()
|
||||
|
||||
def _key(self, suffix: str) -> str:
|
||||
return f"{self._bl}:workflow:{suffix}"
|
||||
|
||||
def _item_key(self, item_id: str) -> str:
|
||||
return self._key(f"item:{item_id}")
|
||||
|
||||
def _queue_key(self) -> str:
|
||||
return self._key("queue")
|
||||
|
||||
def _runtime_key(self) -> str:
|
||||
return self._key("runtime")
|
||||
|
||||
def _control_key(self) -> str:
|
||||
return self._key("control")
|
||||
|
||||
def _events_key(self) -> str:
|
||||
return self._key("events")
|
||||
|
||||
def _next_item_id(self) -> str:
|
||||
n = int(self._client.incr(self._key("item_seq")))
|
||||
return f"wf_{int(time.time())}_{n:06d}"
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Queue item operations
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
def create_item(self, item: QueueItem) -> QueueItem:
|
||||
if not item.item_id:
|
||||
item.item_id = self._next_item_id()
|
||||
|
||||
pipe = self._client.pipeline(transaction=True)
|
||||
pipe.set(self._item_key(item.item_id), item.model_dump_json())
|
||||
pipe.zadd(self._queue_key(), {item.item_id: float(item.order_index)})
|
||||
pipe.execute()
|
||||
|
||||
self.append_event(WorkflowEvent(
|
||||
beamline=self._bl,
|
||||
item_id=item.item_id,
|
||||
event_type="item_created",
|
||||
message="Queue item created",
|
||||
payload={"status": item.status.value},
|
||||
))
|
||||
|
||||
return item
|
||||
|
||||
def get_item(self, item_id: str) -> QueueItem | None:
|
||||
raw = self._client.get(self._item_key(item_id))
|
||||
if raw is None:
|
||||
return None
|
||||
return QueueItem.model_validate_json(raw)
|
||||
|
||||
def update_item(self, item_id: str, patch: dict[str, Any]) -> QueueItem:
|
||||
item = self.get_item(item_id)
|
||||
if item is None:
|
||||
raise KeyError(f"Queue item not found: {item_id}")
|
||||
|
||||
updated = item.model_copy(update=patch)
|
||||
self._client.set(self._item_key(item_id), updated.model_dump_json())
|
||||
return updated
|
||||
|
||||
def delete_item(self, item_id: str) -> None:
|
||||
pipe = self._client.pipeline(transaction=True)
|
||||
pipe.delete(self._item_key(item_id))
|
||||
pipe.zrem(self._queue_key(), item_id)
|
||||
pipe.execute()
|
||||
|
||||
def list_queue_order(self) -> list[str]:
|
||||
return [str(x) for x in self._client.zrange(self._queue_key(), 0, -1)]
|
||||
|
||||
def list_items(self, *, include_finished: bool = True) -> list[QueueItem]:
|
||||
out: list[QueueItem] = []
|
||||
for item_id in self.list_queue_order():
|
||||
item = self.get_item(item_id)
|
||||
if item is None:
|
||||
continue
|
||||
if not include_finished and item.status in {
|
||||
QueueItemStatus.COMPLETED,
|
||||
QueueItemStatus.FAILED,
|
||||
QueueItemStatus.SKIPPED,
|
||||
QueueItemStatus.ABORTED,
|
||||
}:
|
||||
continue
|
||||
out.append(item)
|
||||
return out
|
||||
|
||||
def get_next_pending_item(self) -> QueueItem | None:
|
||||
for item_id in self.list_queue_order():
|
||||
item = self.get_item(item_id)
|
||||
if item is not None and item.status == QueueItemStatus.PENDING:
|
||||
return item
|
||||
return None
|
||||
|
||||
def move_item(self, item_id: str, new_order_index: int) -> None:
|
||||
if self.get_item(item_id) is None:
|
||||
raise KeyError(f"Queue item not found: {item_id}")
|
||||
self._client.zadd(self._queue_key(), {item_id: float(new_order_index)})
|
||||
self.update_item(item_id, {"order_index": new_order_index})
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Runtime state
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
def get_runtime(self) -> RuntimeState:
|
||||
raw = self._client.get(self._runtime_key())
|
||||
if raw is None:
|
||||
return RuntimeState()
|
||||
return RuntimeState.model_validate_json(raw)
|
||||
|
||||
def set_runtime(self, runtime: RuntimeState) -> RuntimeState:
|
||||
runtime.last_update = time.time()
|
||||
self._client.set(self._runtime_key(), runtime.model_dump_json())
|
||||
return runtime
|
||||
|
||||
def patch_runtime(self, patch: dict[str, Any]) -> RuntimeState:
|
||||
runtime = self.get_runtime()
|
||||
updated = runtime.model_copy(update=patch)
|
||||
return self.set_runtime(updated)
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Control state
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
def get_control(self) -> ControlState:
|
||||
raw = self._client.get(self._control_key())
|
||||
if raw is None:
|
||||
return ControlState()
|
||||
return ControlState.model_validate_json(raw)
|
||||
|
||||
def request_control(self, patch: dict[str, Any], *, requested_by: str) -> ControlState:
|
||||
control = self.get_control()
|
||||
updated = control.model_copy(update={
|
||||
**patch,
|
||||
"requested_by": requested_by,
|
||||
"requested_at": time.time(),
|
||||
})
|
||||
self._client.set(self._control_key(), updated.model_dump_json())
|
||||
|
||||
self.append_event(WorkflowEvent(
|
||||
beamline=self._bl,
|
||||
item_id=self.get_runtime().current_item_id,
|
||||
event_type="control_requested",
|
||||
actor=requested_by,
|
||||
payload=patch,
|
||||
))
|
||||
|
||||
return updated
|
||||
|
||||
def clear_control(self) -> ControlState:
|
||||
control = ControlState()
|
||||
self._client.set(self._control_key(), control.model_dump_json())
|
||||
return control
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Event log
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
def append_event(self, event: WorkflowEvent) -> None:
|
||||
payload = event.model_dump_json()
|
||||
self._client.xadd(self._events_key(), {"json": payload}, maxlen=5000, approximate=True)
|
||||
|
||||
def read_events(self, *, limit: int = 200) -> list[WorkflowEvent]:
|
||||
rows = self._client.xrevrange(self._events_key(), count=limit)
|
||||
out: list[WorkflowEvent] = []
|
||||
for _, fields in reversed(rows):
|
||||
raw = fields.get("json")
|
||||
if raw:
|
||||
out.append(WorkflowEvent.model_validate_json(raw))
|
||||
return out
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Step status updates
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
def update_step(
|
||||
self,
|
||||
item_id: str,
|
||||
step_index: int,
|
||||
*,
|
||||
status: str | None = None,
|
||||
message: str | None = None,
|
||||
error_detail: str | None = None,
|
||||
) -> QueueItem:
|
||||
item = self.get_item(item_id)
|
||||
if item is None:
|
||||
raise KeyError(f"Queue item not found: {item_id}")
|
||||
|
||||
if not (0 <= step_index < len(item.steps)):
|
||||
raise IndexError(f"Step index out of range: {step_index}")
|
||||
|
||||
step = item.steps[step_index]
|
||||
|
||||
if status is not None:
|
||||
step.status = status
|
||||
if status == "running" and step.started_at is None:
|
||||
step.started_at = time.time()
|
||||
elif status in ("completed", "failed", "skipped"):
|
||||
step.completed_at = time.time()
|
||||
|
||||
if message is not None:
|
||||
step.message = message
|
||||
|
||||
if error_detail is not None:
|
||||
step.error_detail = error_detail
|
||||
|
||||
item.steps[step_index] = step
|
||||
self._client.set(self._item_key(item_id), item.model_dump_json())
|
||||
|
||||
return item
|
||||
@@ -0,0 +1,351 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from aare.common.automation_models import (
|
||||
WorkflowStateKind,
|
||||
WorkflowMode,
|
||||
StepStatus,
|
||||
StateResult,
|
||||
WorkflowContext,
|
||||
TransitionRule,
|
||||
StateDefinition,
|
||||
)
|
||||
|
||||
|
||||
STATE_REGISTRY: dict[WorkflowStateKind, StateDefinition] = {
|
||||
WorkflowStateKind.MOUNT: StateDefinition(
|
||||
kind=WorkflowStateKind.MOUNT,
|
||||
description="Mount the sample",
|
||||
transitions=(
|
||||
TransitionRule(
|
||||
to_state=WorkflowStateKind.LOOP_CENTRE,
|
||||
allowed_modes=frozenset({
|
||||
WorkflowMode.FLEXIBLE_MANUAL,
|
||||
WorkflowMode.GUIDED_MANUAL,
|
||||
WorkflowMode.AUTOMATION,
|
||||
}),
|
||||
),
|
||||
),
|
||||
),
|
||||
WorkflowStateKind.LOOP_CENTRE: StateDefinition(
|
||||
kind=WorkflowStateKind.LOOP_CENTRE,
|
||||
description="Centre the loop",
|
||||
transitions=(
|
||||
TransitionRule(
|
||||
to_state=WorkflowStateKind.RASTER,
|
||||
allowed_modes=frozenset({
|
||||
WorkflowMode.FLEXIBLE_MANUAL,
|
||||
WorkflowMode.GUIDED_MANUAL,
|
||||
WorkflowMode.AUTOMATION,
|
||||
}),
|
||||
),
|
||||
TransitionRule(
|
||||
to_state=WorkflowStateKind.DATA_COLLECTION,
|
||||
allowed_modes=frozenset({
|
||||
WorkflowMode.FLEXIBLE_MANUAL,
|
||||
WorkflowMode.GUIDED_MANUAL,
|
||||
}),
|
||||
optional=True,
|
||||
),
|
||||
),
|
||||
),
|
||||
WorkflowStateKind.RASTER: StateDefinition(
|
||||
kind=WorkflowStateKind.RASTER,
|
||||
description="Run raster scan",
|
||||
transitions=(
|
||||
TransitionRule(
|
||||
to_state=WorkflowStateKind.DATA_COLLECTION,
|
||||
allowed_modes=frozenset({
|
||||
WorkflowMode.FLEXIBLE_MANUAL,
|
||||
WorkflowMode.GUIDED_MANUAL,
|
||||
WorkflowMode.AUTOMATION,
|
||||
}),
|
||||
),
|
||||
),
|
||||
),
|
||||
WorkflowStateKind.DATA_COLLECTION: StateDefinition(
|
||||
kind=WorkflowStateKind.DATA_COLLECTION,
|
||||
description="Collect diffraction data",
|
||||
transitions=(),
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def get_state_definition(kind: WorkflowStateKind) -> StateDefinition:
|
||||
try:
|
||||
return STATE_REGISTRY[kind]
|
||||
except KeyError as exc:
|
||||
raise KeyError(f"Unknown workflow state: {kind}") from exc
|
||||
|
||||
|
||||
def get_allowed_next_states(
|
||||
kind: WorkflowStateKind,
|
||||
mode: WorkflowMode | None = None,
|
||||
) -> list[WorkflowStateKind]:
|
||||
definition = get_state_definition(kind)
|
||||
out: list[WorkflowStateKind] = []
|
||||
|
||||
for transition in definition.transitions:
|
||||
if mode is None:
|
||||
out.append(transition.to_state)
|
||||
continue
|
||||
|
||||
if not transition.allowed_modes or mode in transition.allowed_modes:
|
||||
out.append(transition.to_state)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
def can_transition(
|
||||
from_state: WorkflowStateKind,
|
||||
to_state: WorkflowStateKind,
|
||||
mode: WorkflowMode | None = None,
|
||||
) -> bool:
|
||||
return to_state in get_allowed_next_states(from_state, mode=mode)
|
||||
|
||||
|
||||
class StateHandler(ABC):
|
||||
state_kind: WorkflowStateKind
|
||||
|
||||
def __init__(self, registry: dict[WorkflowStateKind, StateDefinition] | None = None):
|
||||
self._registry = registry or STATE_REGISTRY
|
||||
|
||||
def definition(self) -> StateDefinition:
|
||||
return get_state_definition(self.state_kind)
|
||||
|
||||
def can_run(self, context: WorkflowContext) -> bool:
|
||||
return context.current_state in (None, self.state_kind)
|
||||
|
||||
@abstractmethod
|
||||
def validate(self, context: WorkflowContext) -> None:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def execute(self, context: WorkflowContext) -> StateResult:
|
||||
pass
|
||||
|
||||
|
||||
class MountHandler(StateHandler):
|
||||
state_kind = WorkflowStateKind.MOUNT
|
||||
|
||||
def validate(self, context: WorkflowContext) -> None:
|
||||
if context.abort_requested:
|
||||
raise RuntimeError("Abort requested; cannot mount.")
|
||||
|
||||
def execute(self, context: WorkflowContext) -> StateResult:
|
||||
self.validate(context)
|
||||
context.current_state = WorkflowStateKind.MOUNT
|
||||
context.last_message = "Sample mounted"
|
||||
return StateResult(
|
||||
state=self.state_kind,
|
||||
status=StepStatus.SUCCESS,
|
||||
message="Sample mounted successfully.",
|
||||
payload={"mounted": True},
|
||||
)
|
||||
|
||||
|
||||
class LoopCentreHandler(StateHandler):
|
||||
state_kind = WorkflowStateKind.LOOP_CENTRE
|
||||
|
||||
def validate(self, context: WorkflowContext) -> None:
|
||||
if context.abort_requested:
|
||||
raise RuntimeError("Abort requested; cannot loop-centre.")
|
||||
|
||||
def execute(self, context: WorkflowContext) -> StateResult:
|
||||
self.validate(context)
|
||||
context.current_state = WorkflowStateKind.LOOP_CENTRE
|
||||
context.last_message = "Loop centred"
|
||||
return StateResult(
|
||||
state=self.state_kind,
|
||||
status=StepStatus.SUCCESS,
|
||||
message="Loop centring completed.",
|
||||
payload={"centred": True},
|
||||
)
|
||||
|
||||
|
||||
class RasterHandler(StateHandler):
|
||||
state_kind = WorkflowStateKind.RASTER
|
||||
|
||||
def validate(self, context: WorkflowContext) -> None:
|
||||
if context.abort_requested:
|
||||
raise RuntimeError("Abort requested; cannot raster.")
|
||||
|
||||
def execute(self, context: WorkflowContext) -> StateResult:
|
||||
self.validate(context)
|
||||
context.current_state = WorkflowStateKind.RASTER
|
||||
context.last_message = "Raster completed"
|
||||
return StateResult(
|
||||
state=self.state_kind,
|
||||
status=StepStatus.SUCCESS,
|
||||
message="Raster scan completed.",
|
||||
payload={"best_spot_found": True},
|
||||
)
|
||||
|
||||
|
||||
class DataCollectionHandler(StateHandler):
|
||||
state_kind = WorkflowStateKind.DATA_COLLECTION
|
||||
|
||||
def validate(self, context: WorkflowContext) -> None:
|
||||
if context.abort_requested:
|
||||
raise RuntimeError("Abort requested; cannot collect data.")
|
||||
|
||||
def execute(self, context: WorkflowContext) -> StateResult:
|
||||
self.validate(context)
|
||||
context.current_state = WorkflowStateKind.DATA_COLLECTION
|
||||
context.last_message = "Data collected"
|
||||
return StateResult(
|
||||
state=self.state_kind,
|
||||
status=StepStatus.SUCCESS,
|
||||
message="Data collection completed.",
|
||||
payload={"frames_collected": 1},
|
||||
)
|
||||
|
||||
|
||||
class WorkflowRunner:
|
||||
"""Simple in-memory runner (no persistence)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
registry: dict[WorkflowStateKind, StateDefinition] | None = None,
|
||||
handlers: dict[WorkflowStateKind, StateHandler] | None = None,
|
||||
):
|
||||
self._registry = registry or STATE_REGISTRY
|
||||
self._handlers = handlers or HANDLER_REGISTRY
|
||||
|
||||
def get_handler(self, state: WorkflowStateKind) -> StateHandler:
|
||||
try:
|
||||
return self._handlers[state]
|
||||
except KeyError as exc:
|
||||
raise KeyError(f"No handler registered for state: {state}") from exc
|
||||
|
||||
def can_move_to(
|
||||
self,
|
||||
current: WorkflowStateKind,
|
||||
next_state: WorkflowStateKind,
|
||||
mode: WorkflowMode,
|
||||
) -> bool:
|
||||
return can_transition(current, next_state, mode)
|
||||
|
||||
def run_state(
|
||||
self,
|
||||
context: WorkflowContext,
|
||||
state: WorkflowStateKind,
|
||||
) -> StateResult:
|
||||
if context.current_state is not None:
|
||||
if not self.can_move_to(context.current_state, state, context.mode):
|
||||
raise RuntimeError(
|
||||
f"Transition not allowed: {context.current_state} -> {state}"
|
||||
)
|
||||
|
||||
handler = self.get_handler(state)
|
||||
result = handler.execute(context)
|
||||
|
||||
context.current_state = state
|
||||
context.current_step_index += 1
|
||||
context.last_message = result.message
|
||||
|
||||
return result
|
||||
|
||||
|
||||
class SimulatedMountHandler(StateHandler):
|
||||
"""Simulated mount handler for testing - doesn't actually mount."""
|
||||
state_kind = WorkflowStateKind.MOUNT
|
||||
|
||||
def validate(self, context: WorkflowContext) -> None:
|
||||
if context.abort_requested:
|
||||
raise RuntimeError("Abort requested; cannot mount.")
|
||||
|
||||
def execute(self, context: WorkflowContext) -> StateResult:
|
||||
self.validate(context)
|
||||
import time
|
||||
time.sleep(0.5) # Simulate some work
|
||||
context.current_state = WorkflowStateKind.MOUNT
|
||||
context.last_message = "[SIMULATION] Sample would be mounted"
|
||||
return StateResult(
|
||||
state=self.state_kind,
|
||||
status=StepStatus.SUCCESS,
|
||||
message=f"[SIMULATION] Would mount sample_id={context.sample_id}",
|
||||
payload={"mounted": True, "simulated": True},
|
||||
)
|
||||
|
||||
|
||||
class SimulatedLoopCentreHandler(StateHandler):
|
||||
"""Simulated loop centre handler for testing."""
|
||||
state_kind = WorkflowStateKind.LOOP_CENTRE
|
||||
|
||||
def validate(self, context: WorkflowContext) -> None:
|
||||
if context.abort_requested:
|
||||
raise RuntimeError("Abort requested; cannot loop-centre.")
|
||||
|
||||
def execute(self, context: WorkflowContext) -> StateResult:
|
||||
self.validate(context)
|
||||
import time
|
||||
time.sleep(0.3)
|
||||
context.current_state = WorkflowStateKind.LOOP_CENTRE
|
||||
context.last_message = "[SIMULATION] Loop would be centred"
|
||||
return StateResult(
|
||||
state=self.state_kind,
|
||||
status=StepStatus.SUCCESS,
|
||||
message="[SIMULATION] Would run loop centering algorithm",
|
||||
payload={"centred": True, "simulated": True},
|
||||
)
|
||||
|
||||
|
||||
class SimulatedRasterHandler(StateHandler):
|
||||
"""Simulated raster handler for testing."""
|
||||
state_kind = WorkflowStateKind.RASTER
|
||||
|
||||
def validate(self, context: WorkflowContext) -> None:
|
||||
if context.abort_requested:
|
||||
raise RuntimeError("Abort requested; cannot raster.")
|
||||
|
||||
def execute(self, context: WorkflowContext) -> StateResult:
|
||||
self.validate(context)
|
||||
import time
|
||||
time.sleep(0.4)
|
||||
context.current_state = WorkflowStateKind.RASTER
|
||||
context.last_message = "[SIMULATION] Raster scan would be performed"
|
||||
return StateResult(
|
||||
state=self.state_kind,
|
||||
status=StepStatus.SUCCESS,
|
||||
message="[SIMULATION] Would run raster scan, find best diffraction spot",
|
||||
payload={"best_spot_found": True, "simulated": True},
|
||||
)
|
||||
|
||||
|
||||
class SimulatedDataCollectionHandler(StateHandler):
|
||||
"""Simulated data collection handler for testing."""
|
||||
state_kind = WorkflowStateKind.DATA_COLLECTION
|
||||
|
||||
def validate(self, context: WorkflowContext) -> None:
|
||||
if context.abort_requested:
|
||||
raise RuntimeError("Abort requested; cannot collect data.")
|
||||
|
||||
def execute(self, context: WorkflowContext) -> StateResult:
|
||||
self.validate(context)
|
||||
import time
|
||||
time.sleep(0.5)
|
||||
context.current_state = WorkflowStateKind.DATA_COLLECTION
|
||||
context.last_message = "[SIMULATION] Data collection would be performed"
|
||||
return StateResult(
|
||||
state=self.state_kind,
|
||||
status=StepStatus.SUCCESS,
|
||||
message="[SIMULATION] Would collect 1800 frames at 0.2° oscillation",
|
||||
payload={"frames_collected": 1800, "simulated": True},
|
||||
)
|
||||
|
||||
|
||||
# Simulated handler registry for testing
|
||||
SIMULATED_HANDLER_REGISTRY: dict[WorkflowStateKind, StateHandler] = {
|
||||
WorkflowStateKind.MOUNT: SimulatedMountHandler(STATE_REGISTRY),
|
||||
WorkflowStateKind.LOOP_CENTRE: SimulatedLoopCentreHandler(STATE_REGISTRY),
|
||||
WorkflowStateKind.RASTER: SimulatedRasterHandler(STATE_REGISTRY),
|
||||
WorkflowStateKind.DATA_COLLECTION: SimulatedDataCollectionHandler(STATE_REGISTRY),
|
||||
}
|
||||
|
||||
HANDLER_REGISTRY: dict[WorkflowStateKind, StateHandler] = {
|
||||
WorkflowStateKind.MOUNT: MountHandler(STATE_REGISTRY),
|
||||
WorkflowStateKind.LOOP_CENTRE: LoopCentreHandler(STATE_REGISTRY),
|
||||
WorkflowStateKind.RASTER: RasterHandler(STATE_REGISTRY),
|
||||
WorkflowStateKind.DATA_COLLECTION: DataCollectionHandler(STATE_REGISTRY),
|
||||
}
|
||||
@@ -5,8 +5,8 @@ from datetime import datetime, timedelta, UTC
|
||||
from typing import List
|
||||
|
||||
import jwt
|
||||
from fastapi import HTTPException, status
|
||||
from fastapi.security import OAuth2PasswordRequestForm
|
||||
from fastapi import Depends, HTTPException, status
|
||||
from fastapi.security import OAuth2PasswordRequestForm, OAuth2PasswordBearer
|
||||
from pydantic import BaseModel
|
||||
|
||||
from aare.daq.config import BeamlineConfig
|
||||
@@ -25,6 +25,8 @@ SESSION_EXPIRE_SECONDS = 60 * 10
|
||||
STAFF_GROUP = "unx-MXgroup"
|
||||
SUPER_USERS = ["e10019", "e11206", "e18147"]
|
||||
|
||||
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token")
|
||||
|
||||
class TokenData(BaseModel):
|
||||
sub: str # Username
|
||||
pgroups: List[str]
|
||||
@@ -54,7 +56,7 @@ def authenticate_user(cfg: BeamlineConfig, form_data: OAuth2PasswordRequestForm)
|
||||
return create_access_token(token)
|
||||
|
||||
|
||||
def parse_token(token: str) -> TokenData:
|
||||
def parse_token(token: str = Depends(oauth2_scheme)) -> TokenData:
|
||||
try:
|
||||
payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
|
||||
token = TokenData(**payload)
|
||||
|
||||
@@ -0,0 +1,694 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from typing import AsyncGenerator, TYPE_CHECKING
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from pydantic import BaseModel, Field
|
||||
from starlette.responses import StreamingResponse
|
||||
|
||||
from aare.common.automation_models import (
|
||||
QueueItem,
|
||||
QueueItemStatus,
|
||||
WorkflowContext,
|
||||
WorkflowEvent,
|
||||
WorkflowMode,
|
||||
WorkflowStepRecord,
|
||||
RuntimeState,
|
||||
ControlState,
|
||||
WorkflowStateKind, EventListResponse, StepActionResponse, ControlActionResponse, RuntimeResponse,
|
||||
CreateQueueItemRequest, QueueListResponse, MoveItemRequest, AutomationStatusResponse,
|
||||
)
|
||||
from aare.common.automation_queue_manager import (
|
||||
WorkflowRedisManager,
|
||||
build_default_steps,
|
||||
)
|
||||
from aare.common.automation_workflow import STATE_REGISTRY, HANDLER_REGISTRY
|
||||
from aare.daq.automation_runner import PersistentWorkflowRunner
|
||||
from aare.daq.auth import parse_token, check_jwt_rw, check_jwt_ro, oauth2_scheme
|
||||
from aare.common.models import TokenData
|
||||
from aare.daq.config import BeamlineConfig
|
||||
|
||||
from aare.daq.automation_runner import AutomationLoop
|
||||
|
||||
router = APIRouter(prefix="/workflow", tags=["workflow"])
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Dependency: get managers
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
# These will be set up when the router is included
|
||||
_redis_manager: WorkflowRedisManager | None = None
|
||||
_runner: PersistentWorkflowRunner | None = None
|
||||
_cfg: BeamlineConfig | None = None
|
||||
_automation_loop: AutomationLoop | None = None
|
||||
|
||||
|
||||
def set_workflow_dependencies(
|
||||
redis_manager: WorkflowRedisManager,
|
||||
runner: PersistentWorkflowRunner,
|
||||
cfg: BeamlineConfig,
|
||||
) -> None:
|
||||
global _redis_manager, _runner, _cfg, _automation_loop
|
||||
_redis_manager = redis_manager
|
||||
_runner = runner
|
||||
_cfg = cfg
|
||||
_automation_loop = AutomationLoop(runner, redis_manager)
|
||||
|
||||
|
||||
def get_automation_loop() -> AutomationLoop:
|
||||
if _automation_loop is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="Automation loop not initialized",
|
||||
)
|
||||
return _automation_loop
|
||||
|
||||
|
||||
def get_cfg() -> BeamlineConfig:
|
||||
if _cfg is None:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail="Workflow system not initialized",
|
||||
)
|
||||
return _cfg
|
||||
|
||||
|
||||
def get_redis_manager() -> WorkflowRedisManager:
|
||||
if _redis_manager is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="Workflow system not initialized",
|
||||
)
|
||||
return _redis_manager
|
||||
|
||||
|
||||
def get_runner() -> PersistentWorkflowRunner:
|
||||
if _runner is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="Workflow runner not initialized",
|
||||
)
|
||||
return _runner
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Queue management endpoints
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.get("/queue", response_model=QueueListResponse)
|
||||
async def list_queue(
|
||||
include_finished: bool = False,
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""List all items in the workflow queue."""
|
||||
data = parse_token(token)
|
||||
check_jwt_ro(get_cfg(), data)
|
||||
items = redis_mgr.list_items(include_finished=include_finished)
|
||||
|
||||
# Filter by pgroup if not staff
|
||||
if not data.staff:
|
||||
items = [i for i in items if i.owner_pgroup in data.pgroups]
|
||||
|
||||
return QueueListResponse(items=items, total=len(items))
|
||||
|
||||
|
||||
@router.post("/queue", response_model=QueueItem)
|
||||
async def create_queue_item(
|
||||
request: CreateQueueItemRequest,
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Add a new item to the workflow queue."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
# Build steps
|
||||
if request.steps:
|
||||
steps = [
|
||||
WorkflowStepRecord(kind=s)
|
||||
for s in request.steps
|
||||
if s in [sk.value for sk in WorkflowStateKind]
|
||||
]
|
||||
else:
|
||||
steps = build_default_steps()
|
||||
|
||||
item = QueueItem(
|
||||
item_id="",
|
||||
beamline=redis_mgr._bl,
|
||||
sample_id=request.sample_id,
|
||||
sample_name=request.sample_name,
|
||||
owner_pgroup=data.pgroups[0] if data.pgroups else "",
|
||||
created_by=data.sub,
|
||||
priority=request.priority,
|
||||
order_index=int(asyncio.get_event_loop().time() * 1000),
|
||||
steps=steps,
|
||||
recipe=request.recipe,
|
||||
)
|
||||
|
||||
item = redis_mgr.create_item(item)
|
||||
return item
|
||||
|
||||
|
||||
@router.get("/queue/{item_id}", response_model=QueueItem)
|
||||
async def get_queue_item(
|
||||
item_id: str,
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Get a single queue item by ID."""
|
||||
data = parse_token(token)
|
||||
check_jwt_ro(get_cfg(), data)
|
||||
|
||||
item = redis_mgr.get_item(item_id)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=404, detail="Item not found")
|
||||
|
||||
# Check access
|
||||
if not data.staff and item.owner_pgroup not in data.pgroups:
|
||||
raise HTTPException(status_code=403, detail="Access denied")
|
||||
|
||||
return item
|
||||
|
||||
|
||||
@router.delete("/queue/{item_id}")
|
||||
async def delete_queue_item(
|
||||
item_id: str,
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Remove an item from the queue."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
item = redis_mgr.get_item(item_id)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=404, detail="Item not found")
|
||||
|
||||
# Check access
|
||||
if not data.staff and item.owner_pgroup not in data.pgroups:
|
||||
raise HTTPException(status_code=403, detail="Access denied")
|
||||
|
||||
# Don't allow deleting running items
|
||||
if item.status == QueueItemStatus.RUNNING:
|
||||
raise HTTPException(status_code=409, detail="Cannot delete running item")
|
||||
|
||||
redis_mgr.delete_item(item_id)
|
||||
return {"ok": True, "message": f"Item {item_id} deleted"}
|
||||
|
||||
|
||||
@router.post("/queue/{item_id}/move")
|
||||
async def move_queue_item(
|
||||
item_id: str,
|
||||
request: MoveItemRequest,
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Reorder an item in the queue."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
item = redis_mgr.get_item(item_id)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=404, detail="Item not found")
|
||||
|
||||
if not data.staff and item.owner_pgroup not in data.pgroups:
|
||||
raise HTTPException(status_code=403, detail="Access denied")
|
||||
|
||||
redis_mgr.move_item(item_id, request.new_order_index)
|
||||
return {"ok": True, "message": f"Item {item_id} moved"}
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Runtime and control endpoints
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.get("/runtime", response_model=RuntimeResponse)
|
||||
async def get_runtime(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Get current runtime and control state."""
|
||||
data = parse_token(token)
|
||||
check_jwt_ro(get_cfg(), data)
|
||||
|
||||
return RuntimeResponse(
|
||||
runtime=redis_mgr.get_runtime(),
|
||||
control=redis_mgr.get_control(),
|
||||
)
|
||||
|
||||
|
||||
@router.get("/control", response_model=ControlState)
|
||||
async def get_control(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Get current control state."""
|
||||
data = parse_token(token)
|
||||
check_jwt_ro(get_cfg(), data)
|
||||
return redis_mgr.get_control()
|
||||
|
||||
|
||||
@router.post("/control/pause", response_model=ControlActionResponse)
|
||||
async def request_pause(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Request workflow pause after current step."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
control = redis_mgr.request_control(
|
||||
{"pause_requested": True, "resume_requested": False},
|
||||
requested_by=data.sub,
|
||||
)
|
||||
|
||||
return ControlActionResponse(
|
||||
ok=True,
|
||||
control=control,
|
||||
message="Pause requested",
|
||||
)
|
||||
|
||||
|
||||
@router.post("/control/resume", response_model=ControlActionResponse)
|
||||
async def request_resume(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Resume paused workflow."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
control = redis_mgr.request_control(
|
||||
{"pause_requested": False, "resume_requested": True},
|
||||
requested_by=data.sub,
|
||||
)
|
||||
|
||||
return ControlActionResponse(
|
||||
ok=True,
|
||||
control=control,
|
||||
message="Resume requested",
|
||||
)
|
||||
|
||||
|
||||
@router.post("/control/abort", response_model=ControlActionResponse)
|
||||
async def request_abort(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Abort current workflow execution."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
control = redis_mgr.request_control(
|
||||
{"abort_requested": True},
|
||||
requested_by=data.sub,
|
||||
)
|
||||
|
||||
return ControlActionResponse(
|
||||
ok=True,
|
||||
control=control,
|
||||
message="Abort requested",
|
||||
)
|
||||
|
||||
|
||||
@router.post("/control/skip", response_model=ControlActionResponse)
|
||||
async def request_skip(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Skip current step."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
control = redis_mgr.request_control(
|
||||
{"skip_requested": True},
|
||||
requested_by=data.sub,
|
||||
)
|
||||
|
||||
return ControlActionResponse(
|
||||
ok=True,
|
||||
control=control,
|
||||
message="Skip requested",
|
||||
)
|
||||
|
||||
|
||||
@router.post("/control/clear", response_model=ControlActionResponse)
|
||||
async def clear_control(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Clear all control requests."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
control = redis_mgr.clear_control()
|
||||
|
||||
return ControlActionResponse(
|
||||
ok=True,
|
||||
control=control,
|
||||
message="Control state cleared",
|
||||
)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Execution endpoints
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.post("/start/{item_id}", response_model=StepActionResponse)
|
||||
async def start_item(
|
||||
item_id: str,
|
||||
mode: WorkflowMode = WorkflowMode.GUIDED_MANUAL,
|
||||
token: str = Depends(oauth2_scheme),
|
||||
runner: PersistentWorkflowRunner = Depends(get_runner),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Start processing a queue item."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
item = redis_mgr.get_item(item_id)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=404, detail="Item not found")
|
||||
|
||||
if not data.staff and item.owner_pgroup not in data.pgroups:
|
||||
raise HTTPException(status_code=403, detail="Access denied")
|
||||
|
||||
if item.status == QueueItemStatus.RUNNING:
|
||||
raise HTTPException(status_code=409, detail="Item already running")
|
||||
|
||||
# Check if another item is running
|
||||
runtime = redis_mgr.get_runtime()
|
||||
if runtime.running:
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
detail=f"Another item is running: {runtime.current_item_id}",
|
||||
)
|
||||
|
||||
# Clear any stale control requests
|
||||
redis_mgr.clear_control()
|
||||
|
||||
# Start the item
|
||||
item = runner.start_item(item_id)
|
||||
|
||||
return StepActionResponse(
|
||||
ok=True,
|
||||
item=item,
|
||||
step=item.steps[0].kind if item.steps else None,
|
||||
status="started",
|
||||
message=f"Started processing {item.sample_name}",
|
||||
)
|
||||
|
||||
|
||||
@router.post("/next", response_model=StepActionResponse)
|
||||
async def run_next_step(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
runner: PersistentWorkflowRunner = Depends(get_runner),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""
|
||||
Run the next step in guided mode.
|
||||
|
||||
This is the main endpoint for guided manual operation.
|
||||
User clicks "Next" and this executes one step.
|
||||
"""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
runtime = redis_mgr.get_runtime()
|
||||
|
||||
if not runtime.running or not runtime.current_item_id:
|
||||
raise HTTPException(status_code=409, detail="No item is currently running")
|
||||
|
||||
item = redis_mgr.get_item(runtime.current_item_id)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=404, detail="Running item not found")
|
||||
|
||||
if not data.staff and item.owner_pgroup not in data.pgroups:
|
||||
raise HTTPException(status_code=403, detail="Access denied")
|
||||
|
||||
# Check if all steps completed
|
||||
if item.current_step_index >= len(item.steps):
|
||||
item = runner.complete_item(item.item_id, QueueItemStatus.COMPLETED)
|
||||
return StepActionResponse(
|
||||
ok=True,
|
||||
item=item,
|
||||
status="completed",
|
||||
message="All steps completed",
|
||||
)
|
||||
|
||||
# Check for abort
|
||||
if runner.should_abort():
|
||||
item = runner.complete_item(item.item_id, QueueItemStatus.ABORTED)
|
||||
redis_mgr.clear_control()
|
||||
return StepActionResponse(
|
||||
ok=True,
|
||||
item=item,
|
||||
status="aborted",
|
||||
message="Workflow aborted by user",
|
||||
)
|
||||
|
||||
# Check for skip
|
||||
control = runner.check_control()
|
||||
if control.skip_requested:
|
||||
step_index = item.current_step_index
|
||||
step_kind = item.steps[step_index].kind
|
||||
|
||||
redis_mgr.update_step(item.item_id, step_index, status="skipped")
|
||||
redis_mgr.update_item(item.item_id, {"current_step_index": step_index + 1})
|
||||
redis_mgr.request_control({"skip_requested": False}, requested_by=data.sub)
|
||||
|
||||
redis_mgr.append_event(WorkflowEvent(
|
||||
beamline=runner.beamline,
|
||||
item_id=item.item_id,
|
||||
step=step_kind,
|
||||
event_type="step_skipped",
|
||||
actor=data.sub,
|
||||
message=f"Step {step_kind} skipped by user",
|
||||
))
|
||||
|
||||
item = redis_mgr.get_item(item.item_id)
|
||||
return StepActionResponse(
|
||||
ok=True,
|
||||
item=item,
|
||||
step=step_kind,
|
||||
status="skipped",
|
||||
message=f"Skipped {step_kind}",
|
||||
)
|
||||
|
||||
# Build context
|
||||
context = WorkflowContext(
|
||||
mode=WorkflowMode.GUIDED_MANUAL,
|
||||
queue_id="default",
|
||||
item_id=item.item_id,
|
||||
sample_id=item.sample_id,
|
||||
current_state=WorkflowStateKind(runtime.current_state) if runtime.current_state else None,
|
||||
current_step_index=item.current_step_index,
|
||||
)
|
||||
|
||||
# Run the step
|
||||
try:
|
||||
result = runner.run_current_step(context)
|
||||
except Exception as e:
|
||||
return StepActionResponse(
|
||||
ok=False,
|
||||
item=redis_mgr.get_item(item.item_id),
|
||||
step=item.steps[item.current_step_index].kind,
|
||||
status="failed",
|
||||
message=str(e),
|
||||
)
|
||||
|
||||
# Get updated item
|
||||
item = redis_mgr.get_item(item.item_id)
|
||||
|
||||
# Check if completed
|
||||
if item.current_step_index >= len(item.steps):
|
||||
item = runner.complete_item(item.item_id, QueueItemStatus.COMPLETED)
|
||||
return StepActionResponse(
|
||||
ok=True,
|
||||
item=item,
|
||||
step=result.state.value,
|
||||
status="completed",
|
||||
message="All steps completed",
|
||||
)
|
||||
|
||||
return StepActionResponse(
|
||||
ok=True,
|
||||
item=item,
|
||||
step=result.state.value,
|
||||
status=result.status.value,
|
||||
message=result.message,
|
||||
)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Events endpoints
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.get("/events", response_model=EventListResponse)
|
||||
async def get_events(
|
||||
limit: int = 100,
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""Get recent workflow events."""
|
||||
data = parse_token(token)
|
||||
check_jwt_ro(get_cfg(), data)
|
||||
|
||||
events = redis_mgr.read_events(limit=limit)
|
||||
return EventListResponse(events=events)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# SSE stream
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
|
||||
async def workflow_event_stream(
|
||||
redis_mgr: WorkflowRedisManager,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""
|
||||
Server-Sent Events stream for live workflow updates.
|
||||
|
||||
Polls runtime state and streams changes.
|
||||
"""
|
||||
last_runtime_json = ""
|
||||
last_control_json = ""
|
||||
last_event_id = "0-0"
|
||||
|
||||
try:
|
||||
while True:
|
||||
# Check runtime state
|
||||
runtime = redis_mgr.get_runtime()
|
||||
runtime_json = runtime.model_dump_json()
|
||||
|
||||
if runtime_json != last_runtime_json:
|
||||
last_runtime_json = runtime_json
|
||||
yield f"event: runtime\ndata: {runtime_json}\n\n"
|
||||
|
||||
# Check control state
|
||||
control = redis_mgr.get_control()
|
||||
control_json = control.model_dump_json()
|
||||
|
||||
if control_json != last_control_json:
|
||||
last_control_json = control_json
|
||||
yield f"event: control\ndata: {control_json}\n\n"
|
||||
|
||||
# Check for new events (using Redis streams)
|
||||
try:
|
||||
events_key = redis_mgr._events_key()
|
||||
new_events = redis_mgr._client.xread(
|
||||
{events_key: last_event_id},
|
||||
count=10,
|
||||
block=0,
|
||||
)
|
||||
|
||||
if new_events:
|
||||
for _, messages in new_events:
|
||||
for msg_id, fields in messages:
|
||||
last_event_id = msg_id
|
||||
raw = fields.get("json")
|
||||
if raw:
|
||||
yield f"event: workflow_event\ndata: {raw}\n\n"
|
||||
except Exception:
|
||||
pass # Redis stream read failed, continue polling
|
||||
|
||||
await asyncio.sleep(0.2)
|
||||
|
||||
except asyncio.CancelledError:
|
||||
return
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Automation mode endpoints
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.get("/automation/status", response_model=AutomationStatusResponse)
|
||||
async def get_automation_status(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
loop: "AutomationLoop" = Depends(get_automation_loop),
|
||||
):
|
||||
"""Get automation loop status."""
|
||||
data = parse_token(token)
|
||||
check_jwt_ro(get_cfg(), data)
|
||||
|
||||
return AutomationStatusResponse(
|
||||
enabled=loop.is_enabled,
|
||||
running=loop.is_running,
|
||||
runtime=redis_mgr.get_runtime(),
|
||||
control=redis_mgr.get_control(),
|
||||
)
|
||||
|
||||
|
||||
@router.post("/automation/start", response_model=AutomationStatusResponse)
|
||||
async def start_automation(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
loop: "AutomationLoop" = Depends(get_automation_loop),
|
||||
):
|
||||
"""Start automation mode - processes queue automatically."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
loop.start()
|
||||
|
||||
return AutomationStatusResponse(
|
||||
enabled=loop.is_enabled,
|
||||
running=loop.is_running,
|
||||
runtime=redis_mgr.get_runtime(),
|
||||
control=redis_mgr.get_control(),
|
||||
)
|
||||
|
||||
|
||||
@router.post("/automation/stop", response_model=AutomationStatusResponse)
|
||||
async def stop_automation(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
loop: "AutomationLoop" = Depends(get_automation_loop),
|
||||
):
|
||||
"""Stop automation mode - completes current step then stops."""
|
||||
data = parse_token(token)
|
||||
check_jwt_rw(get_cfg(), data)
|
||||
|
||||
loop.stop()
|
||||
|
||||
return AutomationStatusResponse(
|
||||
enabled=loop.is_enabled,
|
||||
running=loop.is_running,
|
||||
runtime=redis_mgr.get_runtime(),
|
||||
control=redis_mgr.get_control(),
|
||||
)
|
||||
|
||||
@router.get("/sse")
|
||||
async def workflow_sse(
|
||||
token: str = Depends(oauth2_scheme),
|
||||
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
|
||||
):
|
||||
"""
|
||||
SSE endpoint for live workflow updates.
|
||||
|
||||
Events:
|
||||
- runtime: RuntimeState changes
|
||||
- control: ControlState changes
|
||||
- workflow_event: Individual workflow events
|
||||
"""
|
||||
data = parse_token(token)
|
||||
check_jwt_ro(get_cfg(), data)
|
||||
|
||||
return StreamingResponse(
|
||||
workflow_event_stream(redis_mgr),
|
||||
media_type="text/event-stream",
|
||||
headers={
|
||||
"Cache-Control": "no-cache",
|
||||
"Connection": "keep-alive",
|
||||
"Access-Control-Allow-Origin": "*",
|
||||
"Access-Control-Allow-Headers": "Cache-Control",
|
||||
},
|
||||
)
|
||||
@@ -0,0 +1,427 @@
|
||||
import asyncio
|
||||
from typing import Callable
|
||||
|
||||
import redis
|
||||
|
||||
from aare.common.automation_models import (
|
||||
WorkflowEvent,
|
||||
QueueItemStatus,
|
||||
QueueItem,
|
||||
ControlState,
|
||||
WorkflowStateKind,
|
||||
WorkflowContext,
|
||||
StateResult,
|
||||
StateDefinition,
|
||||
WorkflowMode,
|
||||
)
|
||||
from aare.common.automation_queue_manager import (
|
||||
WorkflowRedisManager,
|
||||
build_default_steps
|
||||
)
|
||||
|
||||
from aare.common.automation_workflow import (
|
||||
can_transition,
|
||||
StateHandler,
|
||||
STATE_REGISTRY,
|
||||
HANDLER_REGISTRY
|
||||
)
|
||||
|
||||
|
||||
class PersistentWorkflowRunner:
|
||||
"""
|
||||
Orchestrates workflow execution with Redis persistence.
|
||||
|
||||
Combines:
|
||||
- State handlers (from automation_workflow)
|
||||
- Redis persistence (from automation_queue_manager)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
redis_manager: WorkflowRedisManager,
|
||||
registry: dict[WorkflowStateKind, StateDefinition] | None = None,
|
||||
handlers: dict[WorkflowStateKind, StateHandler] | None = None,
|
||||
):
|
||||
self._redis = redis_manager
|
||||
self._registry = registry or STATE_REGISTRY
|
||||
self._handlers = handlers or HANDLER_REGISTRY
|
||||
|
||||
@property
|
||||
def beamline(self) -> str:
|
||||
return self._redis._bl
|
||||
|
||||
def get_handler(self, state: WorkflowStateKind) -> StateHandler:
|
||||
try:
|
||||
return self._handlers[state]
|
||||
except KeyError as exc:
|
||||
raise KeyError(f"No handler registered for state: {state}") from exc
|
||||
|
||||
def start_item(self, item_id: str) -> QueueItem:
|
||||
item = self._redis.get_item(item_id)
|
||||
if item is None:
|
||||
raise KeyError(f"Item not found: {item_id}")
|
||||
|
||||
item = self._redis.update_item(item_id, {
|
||||
"status": QueueItemStatus.RUNNING,
|
||||
"current_step_index": 0,
|
||||
})
|
||||
|
||||
self._redis.patch_runtime({
|
||||
"running": True,
|
||||
"paused": False,
|
||||
"current_item_id": item_id,
|
||||
"current_step_index": 0,
|
||||
"current_state": None,
|
||||
"last_error": None,
|
||||
})
|
||||
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self.beamline,
|
||||
item_id=item_id,
|
||||
event_type="item_started",
|
||||
message=f"Started processing {item.sample_name}",
|
||||
))
|
||||
|
||||
return item
|
||||
|
||||
def run_current_step(self, context: WorkflowContext) -> StateResult:
|
||||
item = self._redis.get_item(context.item_id)
|
||||
if item is None:
|
||||
raise KeyError(f"Item not found: {context.item_id}")
|
||||
|
||||
step_index = item.current_step_index
|
||||
if step_index >= len(item.steps):
|
||||
raise RuntimeError("No more steps to run")
|
||||
|
||||
step_record = item.steps[step_index]
|
||||
state_kind = WorkflowStateKind(step_record.kind)
|
||||
|
||||
# Check transition is allowed
|
||||
if context.current_state is not None:
|
||||
if not can_transition(context.current_state, state_kind, context.mode):
|
||||
raise RuntimeError(
|
||||
f"Transition not allowed: {context.current_state} -> {state_kind}"
|
||||
)
|
||||
|
||||
# Mark step as running
|
||||
self._redis.update_step(context.item_id, step_index, status="running")
|
||||
self._redis.patch_runtime({
|
||||
"current_state": state_kind.value,
|
||||
"current_step_index": step_index,
|
||||
})
|
||||
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self.beamline,
|
||||
item_id=context.item_id,
|
||||
step=state_kind.value,
|
||||
event_type="step_started",
|
||||
message=f"Starting {state_kind.value}",
|
||||
))
|
||||
|
||||
# Execute handler
|
||||
handler = self.get_handler(state_kind)
|
||||
|
||||
try:
|
||||
result = handler.execute(context)
|
||||
except Exception as e:
|
||||
self._redis.update_step(
|
||||
context.item_id,
|
||||
step_index,
|
||||
status="failed",
|
||||
error_detail=str(e),
|
||||
)
|
||||
self._redis.patch_runtime({"last_error": str(e)})
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self.beamline,
|
||||
item_id=context.item_id,
|
||||
step=state_kind.value,
|
||||
event_type="step_failed",
|
||||
message=str(e),
|
||||
))
|
||||
raise
|
||||
|
||||
# Mark step as completed
|
||||
self._redis.update_step(
|
||||
context.item_id,
|
||||
step_index,
|
||||
status=result.status.value,
|
||||
message=result.message,
|
||||
)
|
||||
|
||||
# Advance step index
|
||||
self._redis.update_item(context.item_id, {
|
||||
"current_step_index": step_index + 1,
|
||||
})
|
||||
|
||||
context.current_state = state_kind
|
||||
context.current_step_index = step_index + 1
|
||||
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self.beamline,
|
||||
item_id=context.item_id,
|
||||
step=state_kind.value,
|
||||
event_type="step_completed",
|
||||
message=result.message,
|
||||
payload=result.payload,
|
||||
))
|
||||
|
||||
return result
|
||||
|
||||
def complete_item(self, item_id: str, status: QueueItemStatus) -> QueueItem:
|
||||
item = self._redis.update_item(item_id, {"status": status})
|
||||
|
||||
self._redis.patch_runtime({
|
||||
"running": False,
|
||||
"current_item_id": None,
|
||||
"current_state": None,
|
||||
})
|
||||
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self.beamline,
|
||||
item_id=item_id,
|
||||
event_type="item_completed",
|
||||
message=f"Item finished with status {status.value}",
|
||||
))
|
||||
|
||||
return item
|
||||
|
||||
def check_control(self) -> ControlState:
|
||||
return self._redis.get_control()
|
||||
|
||||
def should_pause(self) -> bool:
|
||||
return self.check_control().pause_requested
|
||||
|
||||
def should_abort(self) -> bool:
|
||||
return self.check_control().abort_requested
|
||||
|
||||
def should_skip(self) -> bool:
|
||||
return self.check_control().skip_requested
|
||||
|
||||
|
||||
class AutomationLoop:
|
||||
"""
|
||||
Background task that drives fully automated workflow execution.
|
||||
|
||||
Polls control state and processes the queue automatically.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
runner: PersistentWorkflowRunner,
|
||||
redis_manager: WorkflowRedisManager,
|
||||
poll_interval: float = 0.5,
|
||||
):
|
||||
self._runner = runner
|
||||
self._redis = redis_manager
|
||||
self._poll_interval = poll_interval
|
||||
self._task: asyncio.Task | None = None
|
||||
self._enabled = False
|
||||
self._on_step_complete: Callable[[str, str, StateResult], None] | None = None
|
||||
|
||||
@property
|
||||
def is_running(self) -> bool:
|
||||
return self._task is not None and not self._task.done()
|
||||
|
||||
@property
|
||||
def is_enabled(self) -> bool:
|
||||
return self._enabled
|
||||
|
||||
def set_step_callback(self, callback: Callable[[str, str, StateResult], None]) -> None:
|
||||
"""Set callback for step completion: callback(item_id, step_name, result)"""
|
||||
self._on_step_complete = callback
|
||||
|
||||
def start(self) -> None:
|
||||
"""Start the automation loop."""
|
||||
if self._task is not None and not self._task.done():
|
||||
return # Already running
|
||||
|
||||
self._enabled = True
|
||||
self._task = asyncio.create_task(self._run_loop())
|
||||
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self._runner.beamline,
|
||||
event_type="automation_started",
|
||||
message="Automation mode enabled",
|
||||
))
|
||||
|
||||
def stop(self) -> None:
|
||||
"""Stop the automation loop gracefully."""
|
||||
self._enabled = False
|
||||
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self._runner.beamline,
|
||||
event_type="automation_stopped",
|
||||
message="Automation mode disabled",
|
||||
))
|
||||
|
||||
async def _run_loop(self) -> None:
|
||||
"""Main automation loop."""
|
||||
while self._enabled:
|
||||
try:
|
||||
await self._tick()
|
||||
except Exception as e:
|
||||
self._redis.patch_runtime({"last_error": str(e)})
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self._runner.beamline,
|
||||
event_type="automation_error",
|
||||
message=f"Automation error: {e}",
|
||||
))
|
||||
# Brief pause on error before retrying
|
||||
await asyncio.sleep(2.0)
|
||||
|
||||
await asyncio.sleep(self._poll_interval)
|
||||
|
||||
async def _tick(self) -> None:
|
||||
"""Single iteration of the automation loop."""
|
||||
control = self._runner.check_control()
|
||||
runtime = self._redis.get_runtime()
|
||||
|
||||
# Handle pause state
|
||||
if control.pause_requested:
|
||||
if runtime.running and not runtime.paused:
|
||||
self._redis.patch_runtime({"paused": True})
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self._runner.beamline,
|
||||
item_id=runtime.current_item_id,
|
||||
event_type="workflow_paused",
|
||||
message="Workflow paused by user request",
|
||||
))
|
||||
return
|
||||
|
||||
# Handle resume
|
||||
if control.resume_requested and runtime.paused:
|
||||
self._redis.patch_runtime({"paused": False})
|
||||
self._redis.request_control(
|
||||
{"resume_requested": False},
|
||||
requested_by="automation_loop",
|
||||
)
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self._runner.beamline,
|
||||
item_id=runtime.current_item_id,
|
||||
event_type="workflow_resumed",
|
||||
message="Workflow resumed",
|
||||
))
|
||||
|
||||
# Don't process if paused
|
||||
if runtime.paused:
|
||||
return
|
||||
|
||||
# Handle abort
|
||||
if control.abort_requested and runtime.current_item_id:
|
||||
self._runner.complete_item(runtime.current_item_id, QueueItemStatus.ABORTED)
|
||||
self._redis.clear_control()
|
||||
return
|
||||
|
||||
# If nothing running, try to start next item
|
||||
if not runtime.running or not runtime.current_item_id:
|
||||
next_item = self._redis.get_next_pending_item()
|
||||
if next_item is None:
|
||||
return # Queue empty, nothing to do
|
||||
|
||||
self._runner.start_item(next_item.item_id)
|
||||
runtime = self._redis.get_runtime()
|
||||
|
||||
# Process current item
|
||||
item = self._redis.get_item(runtime.current_item_id)
|
||||
if item is None:
|
||||
self._redis.patch_runtime({"running": False, "current_item_id": None})
|
||||
return
|
||||
|
||||
# Check if item is complete
|
||||
if item.current_step_index >= len(item.steps):
|
||||
self._runner.complete_item(item.item_id, QueueItemStatus.COMPLETED)
|
||||
|
||||
# Check for next sample request
|
||||
if control.next_sample_requested:
|
||||
self._redis.request_control(
|
||||
{"next_sample_requested": False},
|
||||
requested_by="automation_loop",
|
||||
)
|
||||
return
|
||||
|
||||
# Handle skip request
|
||||
if control.skip_requested:
|
||||
step_kind = item.steps[item.current_step_index].kind
|
||||
self._redis.update_step(item.item_id, item.current_step_index, status="skipped")
|
||||
self._redis.update_item(item.item_id, {"current_step_index": item.current_step_index + 1})
|
||||
self._redis.request_control({"skip_requested": False}, requested_by="automation_loop")
|
||||
|
||||
self._redis.append_event(WorkflowEvent(
|
||||
beamline=self._runner.beamline,
|
||||
item_id=item.item_id,
|
||||
step=step_kind,
|
||||
event_type="step_skipped",
|
||||
message=f"Step {step_kind} skipped",
|
||||
))
|
||||
return
|
||||
|
||||
# Build context and run step
|
||||
context = WorkflowContext(
|
||||
mode=WorkflowMode.AUTOMATION,
|
||||
queue_id="default",
|
||||
item_id=item.item_id,
|
||||
sample_id=item.sample_id,
|
||||
current_state=WorkflowStateKind(runtime.current_state) if runtime.current_state else None,
|
||||
current_step_index=item.current_step_index,
|
||||
)
|
||||
|
||||
# Run the step (this is blocking in the async context)
|
||||
result = await asyncio.to_thread(self._runner.run_current_step, context)
|
||||
|
||||
# Notify callback if set
|
||||
if self._on_step_complete:
|
||||
step_name = item.steps[item.current_step_index - 1].kind # -1 because index advanced
|
||||
self._on_step_complete(item.item_id, step_name, result)
|
||||
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
client = redis.Redis(host="localhost", port=6379, db=0, decode_responses=True)
|
||||
redis_mgr = WorkflowRedisManager(client, beamline="x10sa")
|
||||
|
||||
# Create a queue item
|
||||
item = QueueItem(
|
||||
item_id="",
|
||||
beamline="x10sa",
|
||||
sample_id=123,
|
||||
sample_name="lysozyme_01",
|
||||
owner_pgroup="p12345",
|
||||
created_by="user@example.com",
|
||||
steps=build_default_steps(),
|
||||
)
|
||||
item = redis_mgr.create_item(item)
|
||||
|
||||
# Create runner
|
||||
runner = PersistentWorkflowRunner(
|
||||
redis_manager=redis_mgr,
|
||||
registry=STATE_REGISTRY,
|
||||
handlers=HANDLER_REGISTRY,
|
||||
)
|
||||
|
||||
# Start processing
|
||||
runner.start_item(item.item_id)
|
||||
|
||||
# Build context
|
||||
context = WorkflowContext(
|
||||
mode=WorkflowMode.GUIDED_MANUAL,
|
||||
queue_id="default",
|
||||
item_id=item.item_id,
|
||||
sample_id=item.sample_id,
|
||||
)
|
||||
|
||||
# Run each step (in guided mode, user triggers each one)
|
||||
while context.current_step_index < len(item.steps):
|
||||
if runner.should_pause():
|
||||
print("Paused by user")
|
||||
break
|
||||
|
||||
if runner.should_abort():
|
||||
runner.complete_item(item.item_id, QueueItemStatus.ABORTED)
|
||||
break
|
||||
|
||||
result = runner.run_current_step(context)
|
||||
print(f"Step {result.state.value}: {result.status.value}")
|
||||
|
||||
# Mark complete if all steps done
|
||||
if context.current_step_index >= len(item.steps):
|
||||
runner.complete_item(item.item_id, QueueItemStatus.COMPLETED)
|
||||
@@ -34,6 +34,12 @@ from aare.common.exception_handler import (
|
||||
SampleException,
|
||||
UserRightsException,
|
||||
)
|
||||
|
||||
from aare.common.automation_queue_manager import WorkflowRedisManager
|
||||
from aare.common.automation_workflow import STATE_REGISTRY, HANDLER_REGISTRY, SIMULATED_HANDLER_REGISTRY
|
||||
from aare.daq.automation_runner import PersistentWorkflowRunner
|
||||
from aare.daq.automation_api_router import router as workflow_router, set_workflow_dependencies
|
||||
|
||||
logger = setup_logger("aareDAQ")
|
||||
app = FastAPI()
|
||||
register_exception_handlers(app)
|
||||
@@ -45,6 +51,31 @@ bl = mx_beamline()
|
||||
cfg = BeamlineConfig(bl)
|
||||
daq = AareDAQ(cfg, bl)
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Initialize workflow system
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
USE_SIMULATED_WORKFLOW = os.getenv("WORKFLOW_SIMULATION", "0") == "1"
|
||||
|
||||
workflow_redis_manager = WorkflowRedisManager(
|
||||
client=cfg._BeamlineConfig__client, # Reuse existing Redis connection
|
||||
beamline=bl.value,
|
||||
)
|
||||
|
||||
workflow_runner = PersistentWorkflowRunner(
|
||||
redis_manager=workflow_redis_manager,
|
||||
registry=STATE_REGISTRY,
|
||||
handlers=SIMULATED_HANDLER_REGISTRY if USE_SIMULATED_WORKFLOW else HANDLER_REGISTRY,
|
||||
)
|
||||
|
||||
if USE_SIMULATED_WORKFLOW:
|
||||
logger.warning("⚠️ Workflow system running in SIMULATION mode - no actual DAQ operations")
|
||||
|
||||
set_workflow_dependencies(workflow_redis_manager, workflow_runner, cfg)
|
||||
|
||||
# Include the workflow router
|
||||
app.include_router(workflow_router)
|
||||
|
||||
try:
|
||||
daq.sync_current_sample_from_tell(force=True)
|
||||
except Exception as e:
|
||||
|
||||
@@ -55,6 +55,9 @@ from aare.gui.widgets.status_bar import StatusBar
|
||||
from aare.gui.widgets.video_image import VideoGraphicsView
|
||||
from aare.gui.panels.fluorescence_panel import FluorescencePanel
|
||||
|
||||
from aare.gui.panels.automation_panel import WorkflowPanel
|
||||
from aare.gui.threads.workflow_sse_client import WorkflowSSEClient
|
||||
|
||||
logger = setup_logger("aareGUI")
|
||||
|
||||
class MainWindow(QMainWindow):
|
||||
@@ -288,6 +291,17 @@ class MainWindow(QMainWindow):
|
||||
self.addDockWidget(Qt.DockWidgetArea.RightDockWidgetArea, self.smargon_trace_dock)
|
||||
self.smargon_trace_dock.hide()
|
||||
|
||||
# === Workflow Panel ===
|
||||
self.workflow_panel = WorkflowPanel()
|
||||
self.workflow_dock = QDockWidget("Workflow", self)
|
||||
self.workflow_dock.setObjectName("workflow_dock")
|
||||
self.workflow_dock.setWidget(self.workflow_panel)
|
||||
self.workflow_dock.setAllowedAreas(
|
||||
Qt.DockWidgetArea.RightDockWidgetArea | Qt.DockWidgetArea.LeftDockWidgetArea
|
||||
)
|
||||
self.addDockWidget(Qt.DockWidgetArea.RightDockWidgetArea, self.workflow_dock)
|
||||
self.workflow_dock.hide() # Hidden by default
|
||||
|
||||
root_layout.addWidget(top_widget)
|
||||
self.setCentralWidget(root_widget)
|
||||
|
||||
@@ -497,6 +511,46 @@ class MainWindow(QMainWindow):
|
||||
|
||||
register_tutorials(self, self.tutorial_manager)
|
||||
|
||||
# Workflow SSE client
|
||||
if self.__base_url is not None:
|
||||
self.workflow_sse = WorkflowSSEClient(self.__base_url, self.__token, self)
|
||||
self.workflow_sse.runtime_changed.connect(self.workflow_panel.update_runtime)
|
||||
self.workflow_sse.control_changed.connect(self.workflow_panel.update_control)
|
||||
self.workflow_sse.workflow_event.connect(
|
||||
lambda e: self.workflow_panel.on_workflow_event(e.event_type, e.message)
|
||||
)
|
||||
self.workflow_sse.connect()
|
||||
else:
|
||||
self.workflow_sse = None
|
||||
|
||||
# Connect workflow panel signals to DAQ worker
|
||||
self.workflow_panel.request_queue_refresh.connect(self.daq.workflow_load_queue)
|
||||
self.workflow_panel.request_add_sample.connect(self.daq.workflow_add_sample)
|
||||
self.workflow_panel.request_add_samples.connect(self.daq.workflow_add_samples)
|
||||
self.workflow_panel.request_delete_item.connect(self.daq.workflow_delete_item)
|
||||
self.workflow_panel.request_move_item.connect(self.daq.workflow_move_item)
|
||||
self.workflow_panel.request_clear_queue.connect(self.daq.workflow_clear_queue)
|
||||
self.workflow_panel.request_start_item.connect(self.daq.workflow_start_item)
|
||||
self.workflow_panel.request_next_step.connect(self.daq.workflow_next_step)
|
||||
self.workflow_panel.request_pause.connect(self.daq.workflow_pause)
|
||||
self.workflow_panel.request_resume.connect(self.daq.workflow_resume)
|
||||
self.workflow_panel.request_abort.connect(self.daq.workflow_abort)
|
||||
self.workflow_panel.request_skip.connect(self.daq.workflow_skip)
|
||||
self.workflow_panel.request_start_automation.connect(self.daq.workflow_start_automation)
|
||||
self.workflow_panel.request_stop_automation.connect(self.daq.workflow_stop_automation)
|
||||
|
||||
# DAQ worker -> workflow panel
|
||||
self.daq.workflow_queue_loaded.connect(self.workflow_panel.update_queue)
|
||||
self.daq.workflow_item_updated.connect(self.workflow_panel.update_current_item)
|
||||
|
||||
# Forward sample list to workflow panel for "Add All" feature
|
||||
self.daq.spreadsheet.connect(
|
||||
lambda slist: self.workflow_panel.update_available_samples(slist.s)
|
||||
)
|
||||
|
||||
# Initial load
|
||||
QTimer.singleShot(1000, self.daq.workflow_load_queue)
|
||||
|
||||
@Slot(QPixmap)
|
||||
def _on_samcam_prediction_pixmap(self, pix: QPixmap) -> None:
|
||||
self._last_pred_image_ts = time.monotonic()
|
||||
@@ -596,6 +650,13 @@ class MainWindow(QMainWindow):
|
||||
)
|
||||
view_menu.addAction(show_smargon_trace_action)
|
||||
|
||||
show_workflow_action = QAction("Show Workflow Panel", self)
|
||||
show_workflow_action.setCheckable(True)
|
||||
show_workflow_action.setChecked(False)
|
||||
show_workflow_action.triggered.connect(lambda checked: self.workflow_dock.setVisible(checked))
|
||||
self.workflow_dock.visibilityChanged.connect(show_workflow_action.setChecked)
|
||||
view_menu.addAction(show_workflow_action)
|
||||
|
||||
show_log_action = QAction("Show Log", self)
|
||||
show_log_action.setCheckable(True)
|
||||
show_log_action.setChecked(False)
|
||||
@@ -770,6 +831,9 @@ class MainWindow(QMainWindow):
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to stop _samcam_source_timer: {e}")
|
||||
|
||||
if hasattr(self, "workflow_sse") and self.workflow_sse is not None:
|
||||
self.workflow_sse.disconnect()
|
||||
|
||||
for attr_name in (
|
||||
"camera_thread",
|
||||
"prediction_thread",
|
||||
|
||||
@@ -0,0 +1,640 @@
|
||||
"""
|
||||
Workflow automation panel for queue management and step control.
|
||||
|
||||
Displays:
|
||||
- Queue items with status (supports drag-drop from sample list)
|
||||
- Current step progress
|
||||
- Control buttons (Next/Pause/Resume/Abort/Skip)
|
||||
- Live status from SSE
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from PySide6.QtCore import Qt, Signal, Slot, QTimer, QMimeData
|
||||
from PySide6.QtGui import QColor, QDragEnterEvent, QDropEvent, QKeySequence, QShortcut
|
||||
from PySide6.QtWidgets import (
|
||||
QWidget,
|
||||
QVBoxLayout,
|
||||
QHBoxLayout,
|
||||
QLabel,
|
||||
QPushButton,
|
||||
QListWidget,
|
||||
QListWidgetItem,
|
||||
QProgressBar,
|
||||
QGroupBox,
|
||||
QFrame,
|
||||
QSizePolicy,
|
||||
QAbstractItemView,
|
||||
QMenu,
|
||||
QMessageBox,
|
||||
)
|
||||
|
||||
from aare.common.automation_models import (
|
||||
QueueItem,
|
||||
QueueItemStatus,
|
||||
RuntimeState,
|
||||
ControlState,
|
||||
WorkflowStepRecord,
|
||||
)
|
||||
from aare.common.models import SampleShortInfo, SampleShortInfoList
|
||||
from aare.common.logger_config import setup_logger
|
||||
|
||||
logger = setup_logger("aareGUI")
|
||||
|
||||
|
||||
class StepProgressWidget(QWidget):
|
||||
"""Shows progress through workflow steps."""
|
||||
|
||||
def __init__(self, parent: QWidget | None = None):
|
||||
super().__init__(parent)
|
||||
self._steps: list[WorkflowStepRecord] = []
|
||||
self._current_index = 0
|
||||
self._setup_ui()
|
||||
|
||||
def _setup_ui(self) -> None:
|
||||
layout = QVBoxLayout(self)
|
||||
layout.setContentsMargins(4, 4, 4, 4)
|
||||
layout.setSpacing(2)
|
||||
|
||||
self._step_labels: list[QLabel] = []
|
||||
|
||||
self._container = QWidget()
|
||||
self._container_layout = QVBoxLayout(self._container)
|
||||
self._container_layout.setContentsMargins(0, 0, 0, 0)
|
||||
self._container_layout.setSpacing(2)
|
||||
layout.addWidget(self._container)
|
||||
|
||||
def set_steps(self, steps: list[WorkflowStepRecord], current_index: int) -> None:
|
||||
"""Update the step display."""
|
||||
self._steps = steps
|
||||
self._current_index = current_index
|
||||
|
||||
for lbl in self._step_labels:
|
||||
lbl.deleteLater()
|
||||
self._step_labels.clear()
|
||||
|
||||
for i, step in enumerate(steps):
|
||||
lbl = QLabel(f"{i + 1}. {step.kind}")
|
||||
lbl.setStyleSheet(self._style_for_step(i, step.status))
|
||||
self._container_layout.addWidget(lbl)
|
||||
self._step_labels.append(lbl)
|
||||
|
||||
def update_step(self, index: int, status: str) -> None:
|
||||
"""Update a single step's status."""
|
||||
if 0 <= index < len(self._step_labels):
|
||||
self._step_labels[index].setStyleSheet(self._style_for_step(index, status))
|
||||
|
||||
def _style_for_step(self, index: int, status: str) -> str:
|
||||
"""Get stylesheet for step based on status."""
|
||||
base = "padding: 4px; border-radius: 3px; "
|
||||
|
||||
if status == "success":
|
||||
return base + "background-color: #90EE90; color: #006400;"
|
||||
elif status == "running":
|
||||
return base + "background-color: #87CEEB; color: #00008B; font-weight: bold;"
|
||||
elif status == "failed":
|
||||
return base + "background-color: #FFB6C1; color: #8B0000;"
|
||||
elif status == "skipped":
|
||||
return base + "background-color: #D3D3D3; color: #696969; text-decoration: line-through;"
|
||||
elif status == "paused":
|
||||
return base + "background-color: #FFE4B5; color: #8B4513;"
|
||||
else: # pending
|
||||
return base + "background-color: #F0F0F0; color: #808080;"
|
||||
|
||||
|
||||
class DraggableQueueListWidget(QListWidget):
|
||||
"""
|
||||
QListWidget that accepts drops from TellSamplePanel.
|
||||
|
||||
Supports:
|
||||
- Drag-drop samples from tell_sample_panel
|
||||
- Internal reordering via drag
|
||||
- Delete key to remove items
|
||||
"""
|
||||
|
||||
samples_dropped = Signal(list) # list[SampleShortInfo]
|
||||
item_reordered = Signal(str, int) # item_id, new_index
|
||||
delete_requested = Signal(list) # list[item_ids]
|
||||
|
||||
def __init__(self, parent: QWidget | None = None):
|
||||
super().__init__(parent)
|
||||
|
||||
self.setAcceptDrops(True)
|
||||
self.setDragEnabled(True)
|
||||
self.setDragDropMode(QAbstractItemView.DragDropMode.DragDrop)
|
||||
self.setDefaultDropAction(Qt.DropAction.MoveAction)
|
||||
self.setSelectionMode(QAbstractItemView.SelectionMode.ExtendedSelection)
|
||||
self.setContextMenuPolicy(Qt.ContextMenuPolicy.CustomContextMenu)
|
||||
self.customContextMenuRequested.connect(self._show_context_menu)
|
||||
|
||||
# Delete shortcut
|
||||
self._delete_shortcut = QShortcut(QKeySequence.StandardKey.Delete, self)
|
||||
self._delete_shortcut.activated.connect(self._on_delete_pressed)
|
||||
|
||||
# Store item_id -> row mapping
|
||||
self._item_ids: list[str] = []
|
||||
|
||||
def set_item_ids(self, item_ids: list[str]) -> None:
|
||||
"""Track item IDs for reordering."""
|
||||
self._item_ids = item_ids
|
||||
|
||||
def dragEnterEvent(self, event: QDragEnterEvent) -> None:
|
||||
"""Accept drops from sample panels."""
|
||||
mime = event.mimeData()
|
||||
if mime.hasText():
|
||||
# Check if it's sample data (JSON)
|
||||
try:
|
||||
text = mime.text()
|
||||
if text.startswith("{") or text.startswith("["):
|
||||
event.acceptProposedAction()
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
# Accept internal moves
|
||||
if event.source() == self:
|
||||
event.acceptProposedAction()
|
||||
return
|
||||
event.ignore()
|
||||
|
||||
def dragMoveEvent(self, event) -> None:
|
||||
"""Show drop indicator."""
|
||||
if event.mimeData().hasText() or event.source() == self:
|
||||
event.acceptProposedAction()
|
||||
else:
|
||||
event.ignore()
|
||||
|
||||
def dropEvent(self, event: QDropEvent) -> None:
|
||||
"""Handle drop - either samples from panel or internal reorder."""
|
||||
mime = event.mimeData()
|
||||
|
||||
if mime.hasText():
|
||||
text = mime.text()
|
||||
try:
|
||||
# Try to parse as SampleShortInfoList
|
||||
sample_list = SampleShortInfoList.model_validate_json(text)
|
||||
if sample_list.s:
|
||||
self.samples_dropped.emit(sample_list.s)
|
||||
event.acceptProposedAction()
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
try:
|
||||
# Try single sample
|
||||
sample = SampleShortInfo.model_validate_json(text)
|
||||
self.samples_dropped.emit([sample])
|
||||
event.acceptProposedAction()
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Internal reorder
|
||||
if event.source() == self:
|
||||
# Get drop position
|
||||
drop_row = self.indexAt(event.position().toPoint()).row()
|
||||
if drop_row < 0:
|
||||
drop_row = self.count()
|
||||
|
||||
# Get selected items
|
||||
selected = self.selectedItems()
|
||||
if selected and self._item_ids:
|
||||
for item in selected:
|
||||
row = self.row(item)
|
||||
if 0 <= row < len(self._item_ids):
|
||||
item_id = self._item_ids[row]
|
||||
self.item_reordered.emit(item_id, drop_row)
|
||||
|
||||
event.acceptProposedAction()
|
||||
return
|
||||
|
||||
event.ignore()
|
||||
|
||||
def _on_delete_pressed(self) -> None:
|
||||
"""Handle delete key press."""
|
||||
selected = self.selectedItems()
|
||||
if not selected:
|
||||
return
|
||||
|
||||
item_ids = []
|
||||
for item in selected:
|
||||
row = self.row(item)
|
||||
if 0 <= row < len(self._item_ids):
|
||||
item_ids.append(self._item_ids[row])
|
||||
|
||||
if item_ids:
|
||||
self.delete_requested.emit(item_ids)
|
||||
|
||||
def _show_context_menu(self, pos) -> None:
|
||||
"""Show context menu for queue items."""
|
||||
item = self.itemAt(pos)
|
||||
if not item:
|
||||
return
|
||||
|
||||
row = self.row(item)
|
||||
if row < 0 or row >= len(self._item_ids):
|
||||
return
|
||||
|
||||
menu = QMenu(self)
|
||||
|
||||
delete_action = menu.addAction("🗑 Remove from queue")
|
||||
start_action = menu.addAction("▶ Start this item")
|
||||
|
||||
menu.addSeparator()
|
||||
move_top_action = menu.addAction("⬆ Move to top")
|
||||
move_bottom_action = menu.addAction("⬇ Move to bottom")
|
||||
|
||||
action = menu.exec_(self.mapToGlobal(pos))
|
||||
|
||||
if action == delete_action:
|
||||
self.delete_requested.emit([self._item_ids[row]])
|
||||
elif action == start_action:
|
||||
# Will need to emit a signal for this
|
||||
pass
|
||||
elif action == move_top_action:
|
||||
self.item_reordered.emit(self._item_ids[row], 0)
|
||||
elif action == move_bottom_action:
|
||||
self.item_reordered.emit(self._item_ids[row], 999999)
|
||||
|
||||
|
||||
class WorkflowPanel(QWidget):
|
||||
"""
|
||||
Main workflow panel combining queue display and controls.
|
||||
|
||||
Supports:
|
||||
- Drag-drop samples from TELL sample panel
|
||||
- Queue management (delete, reorder, clear, add all)
|
||||
- Step-by-step guided mode
|
||||
- Full automation mode
|
||||
"""
|
||||
|
||||
# Signals for DAQ worker
|
||||
request_queue_refresh = Signal()
|
||||
request_add_sample = Signal(object) # SampleShortInfo
|
||||
request_add_samples = Signal(list) # list[SampleShortInfo]
|
||||
request_delete_item = Signal(str) # item_id
|
||||
request_move_item = Signal(str, int) # item_id, new_order_index
|
||||
request_clear_queue = Signal()
|
||||
request_start_item = Signal(str) # item_id
|
||||
request_next_step = Signal()
|
||||
request_pause = Signal()
|
||||
request_resume = Signal()
|
||||
request_abort = Signal()
|
||||
request_skip = Signal()
|
||||
request_start_automation = Signal()
|
||||
request_stop_automation = Signal()
|
||||
|
||||
def __init__(self, parent: QWidget | None = None):
|
||||
super().__init__(parent)
|
||||
self._queue_items: list[QueueItem] = []
|
||||
self._all_samples: list[SampleShortInfo] = [] # Cache for "Add All"
|
||||
self._runtime: RuntimeState | None = None
|
||||
self._control: ControlState | None = None
|
||||
self._automation_enabled = False
|
||||
self._setup_ui()
|
||||
|
||||
def _setup_ui(self) -> None:
|
||||
layout = QVBoxLayout(self)
|
||||
layout.setContentsMargins(8, 8, 8, 8)
|
||||
layout.setSpacing(8)
|
||||
|
||||
# === Status Section ===
|
||||
status_group = QGroupBox("Current Status")
|
||||
status_layout = QVBoxLayout(status_group)
|
||||
|
||||
self._status_label = QLabel("⏹️ Idle")
|
||||
self._status_label.setStyleSheet("font-size: 14px; font-weight: bold;")
|
||||
status_layout.addWidget(self._status_label)
|
||||
|
||||
self._current_item_label = QLabel("No item running")
|
||||
status_layout.addWidget(self._current_item_label)
|
||||
|
||||
self._step_progress = StepProgressWidget()
|
||||
status_layout.addWidget(self._step_progress)
|
||||
|
||||
layout.addWidget(status_group)
|
||||
|
||||
# === Control Buttons ===
|
||||
controls_group = QGroupBox("Controls")
|
||||
controls_layout = QVBoxLayout(controls_group)
|
||||
|
||||
# Mode toggle
|
||||
mode_layout = QHBoxLayout()
|
||||
self._mode_label = QLabel("Mode:")
|
||||
self._guided_btn = QPushButton("Guided Manual")
|
||||
self._guided_btn.setCheckable(True)
|
||||
self._guided_btn.setChecked(True)
|
||||
self._auto_btn = QPushButton("Automation")
|
||||
self._auto_btn.setCheckable(True)
|
||||
|
||||
self._guided_btn.clicked.connect(self._on_guided_mode)
|
||||
self._auto_btn.clicked.connect(self._on_automation_mode)
|
||||
|
||||
mode_layout.addWidget(self._mode_label)
|
||||
mode_layout.addWidget(self._guided_btn)
|
||||
mode_layout.addWidget(self._auto_btn)
|
||||
mode_layout.addStretch()
|
||||
controls_layout.addLayout(mode_layout)
|
||||
|
||||
# Step controls
|
||||
step_layout = QHBoxLayout()
|
||||
|
||||
self._next_btn = QPushButton("Next Step")
|
||||
self._next_btn.setStyleSheet("background-color: #4CAF50; color: white;")
|
||||
self._next_btn.clicked.connect(self.request_next_step.emit)
|
||||
|
||||
self._skip_btn = QPushButton("Skip")
|
||||
self._skip_btn.setStyleSheet("background-color: #FF9800; color: white;")
|
||||
self._skip_btn.clicked.connect(self.request_skip.emit)
|
||||
|
||||
self._pause_btn = QPushButton("Pause")
|
||||
self._pause_btn.setStyleSheet("background-color: #2196F3; color: white;")
|
||||
self._pause_btn.clicked.connect(self._on_pause_resume)
|
||||
|
||||
self._abort_btn = QPushButton("Abort")
|
||||
self._abort_btn.setStyleSheet("background-color: #F44336; color: white;")
|
||||
self._abort_btn.clicked.connect(self.request_abort.emit)
|
||||
|
||||
step_layout.addWidget(self._next_btn)
|
||||
step_layout.addWidget(self._skip_btn)
|
||||
step_layout.addWidget(self._pause_btn)
|
||||
step_layout.addWidget(self._abort_btn)
|
||||
|
||||
controls_layout.addLayout(step_layout)
|
||||
layout.addWidget(controls_group)
|
||||
|
||||
# === Queue Section ===
|
||||
queue_group = QGroupBox("Queue (drag samples here)")
|
||||
queue_layout = QVBoxLayout(queue_group)
|
||||
|
||||
# Info label
|
||||
self._drop_hint = QLabel("💡 Drag samples from the Sample List to add them")
|
||||
self._drop_hint.setStyleSheet("color: #666; font-style: italic;")
|
||||
queue_layout.addWidget(self._drop_hint)
|
||||
|
||||
# Draggable queue list
|
||||
self._queue_list = DraggableQueueListWidget()
|
||||
self._queue_list.setMinimumHeight(150)
|
||||
self._queue_list.itemDoubleClicked.connect(self._on_item_double_clicked)
|
||||
self._queue_list.samples_dropped.connect(self._on_samples_dropped)
|
||||
self._queue_list.delete_requested.connect(self._on_delete_requested)
|
||||
self._queue_list.item_reordered.connect(self._on_item_reordered)
|
||||
queue_layout.addWidget(self._queue_list)
|
||||
|
||||
# Queue management buttons
|
||||
queue_btn_layout = QHBoxLayout()
|
||||
|
||||
self._add_all_btn = QPushButton("➕ Add All Samples")
|
||||
self._add_all_btn.clicked.connect(self._on_add_all_clicked)
|
||||
self._add_all_btn.setToolTip("Add all samples from the Sample List to the queue")
|
||||
|
||||
self._remove_selected_btn = QPushButton("🗑 Remove Selected")
|
||||
self._remove_selected_btn.clicked.connect(self._on_remove_selected)
|
||||
|
||||
self._clear_btn = QPushButton("✖ Clear Queue")
|
||||
self._clear_btn.clicked.connect(self._on_clear_queue)
|
||||
|
||||
self._refresh_btn = QPushButton("🔄")
|
||||
self._refresh_btn.setFixedWidth(40)
|
||||
self._refresh_btn.setToolTip("Refresh queue")
|
||||
self._refresh_btn.clicked.connect(self.request_queue_refresh.emit)
|
||||
|
||||
queue_btn_layout.addWidget(self._add_all_btn)
|
||||
queue_btn_layout.addWidget(self._remove_selected_btn)
|
||||
queue_btn_layout.addWidget(self._clear_btn)
|
||||
queue_btn_layout.addStretch()
|
||||
queue_btn_layout.addWidget(self._refresh_btn)
|
||||
queue_layout.addLayout(queue_btn_layout)
|
||||
|
||||
layout.addWidget(queue_group)
|
||||
|
||||
# Initial state
|
||||
self._update_button_states()
|
||||
|
||||
def _on_guided_mode(self) -> None:
|
||||
"""Switch to guided manual mode."""
|
||||
self._guided_btn.setChecked(True)
|
||||
self._auto_btn.setChecked(False)
|
||||
if self._automation_enabled:
|
||||
self.request_stop_automation.emit()
|
||||
self._update_button_states()
|
||||
|
||||
def _on_automation_mode(self) -> None:
|
||||
"""Switch to automation mode."""
|
||||
self._auto_btn.setChecked(True)
|
||||
self._guided_btn.setChecked(False)
|
||||
if not self._automation_enabled:
|
||||
self.request_start_automation.emit()
|
||||
self._update_button_states()
|
||||
|
||||
def _on_pause_resume(self) -> None:
|
||||
"""Toggle pause/resume."""
|
||||
if self._runtime and self._runtime.paused:
|
||||
self.request_resume.emit()
|
||||
else:
|
||||
self.request_pause.emit()
|
||||
|
||||
def _on_item_double_clicked(self, item: QListWidgetItem) -> None:
|
||||
"""Start processing double-clicked item."""
|
||||
idx = self._queue_list.row(item)
|
||||
if 0 <= idx < len(self._queue_items):
|
||||
queue_item = self._queue_items[idx]
|
||||
if queue_item.status == QueueItemStatus.PENDING:
|
||||
self.request_start_item.emit(queue_item.item_id)
|
||||
|
||||
def _on_samples_dropped(self, samples: list[SampleShortInfo]) -> None:
|
||||
"""Handle samples dropped onto the queue."""
|
||||
logger.info(f"Adding {len(samples)} samples to workflow queue")
|
||||
for sample in samples:
|
||||
self.request_add_sample.emit(sample)
|
||||
# Refresh after a short delay to let server process
|
||||
QTimer.singleShot(500, self.request_queue_refresh.emit)
|
||||
|
||||
def _on_delete_requested(self, item_ids: list[str]) -> None:
|
||||
"""Handle delete request from queue list."""
|
||||
for item_id in item_ids:
|
||||
self.request_delete_item.emit(item_id)
|
||||
QTimer.singleShot(300, self.request_queue_refresh.emit)
|
||||
|
||||
def _on_item_reordered(self, item_id: str, new_index: int) -> None:
|
||||
"""Handle item reorder."""
|
||||
self.request_move_item.emit(item_id, new_index)
|
||||
QTimer.singleShot(300, self.request_queue_refresh.emit)
|
||||
|
||||
def _on_remove_selected(self) -> None:
|
||||
"""Remove selected items from queue."""
|
||||
selected = self._queue_list.selectedItems()
|
||||
if not selected:
|
||||
return
|
||||
|
||||
item_ids = self._queue_list._item_ids
|
||||
for item in selected:
|
||||
row = self._queue_list.row(item)
|
||||
if 0 <= row < len(item_ids):
|
||||
self.request_delete_item.emit(item_ids[row])
|
||||
|
||||
QTimer.singleShot(300, self.request_queue_refresh.emit)
|
||||
|
||||
def _on_clear_queue(self) -> None:
|
||||
"""Clear all items from queue."""
|
||||
if not self._queue_items:
|
||||
return
|
||||
|
||||
reply = QMessageBox.question(
|
||||
self,
|
||||
"Clear Queue",
|
||||
f"Remove all {len(self._queue_items)} items from the queue?",
|
||||
QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No,
|
||||
)
|
||||
|
||||
if reply == QMessageBox.StandardButton.Yes:
|
||||
self.request_clear_queue.emit()
|
||||
QTimer.singleShot(500, self.request_queue_refresh.emit)
|
||||
|
||||
def _on_add_all_clicked(self) -> None:
|
||||
"""Add all available samples to the queue."""
|
||||
if not self._all_samples:
|
||||
QMessageBox.information(
|
||||
self,
|
||||
"No Samples",
|
||||
"No samples available to add. Load samples in the Sample List first.",
|
||||
)
|
||||
return
|
||||
|
||||
reply = QMessageBox.question(
|
||||
self,
|
||||
"Add All Samples",
|
||||
f"Add all {len(self._all_samples)} samples to the workflow queue?",
|
||||
QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No,
|
||||
)
|
||||
|
||||
if reply == QMessageBox.StandardButton.Yes:
|
||||
self.request_add_samples.emit(self._all_samples)
|
||||
QTimer.singleShot(500, self.request_queue_refresh.emit)
|
||||
|
||||
def _update_button_states(self) -> None:
|
||||
"""Update button enabled/disabled states based on current state."""
|
||||
is_running = self._runtime is not None and self._runtime.running
|
||||
is_paused = self._runtime is not None and self._runtime.paused
|
||||
is_automation = self._auto_btn.isChecked()
|
||||
|
||||
# In automation mode, hide the Next button
|
||||
self._next_btn.setVisible(not is_automation)
|
||||
self._next_btn.setEnabled(is_running and not is_paused)
|
||||
|
||||
self._skip_btn.setEnabled(is_running)
|
||||
self._abort_btn.setEnabled(is_running)
|
||||
|
||||
if is_paused:
|
||||
self._pause_btn.setText("Resume")
|
||||
self._pause_btn.setStyleSheet("background-color: #4CAF50; color: white;")
|
||||
else:
|
||||
self._pause_btn.setText("Pause")
|
||||
self._pause_btn.setStyleSheet("background-color: #2196F3; color: white;")
|
||||
|
||||
self._pause_btn.setEnabled(is_running)
|
||||
|
||||
# Update drop hint visibility
|
||||
has_items = len(self._queue_items) > 0
|
||||
self._drop_hint.setVisible(not has_items)
|
||||
|
||||
# === Public Update Methods ===
|
||||
|
||||
@Slot(list)
|
||||
def update_queue(self, items: list[QueueItem]) -> None:
|
||||
"""Update queue display."""
|
||||
self._queue_items = items
|
||||
self._queue_list.clear()
|
||||
|
||||
item_ids = []
|
||||
for item in items:
|
||||
status_emoji = {
|
||||
QueueItemStatus.PENDING: "⏳",
|
||||
QueueItemStatus.RUNNING: "▶️",
|
||||
QueueItemStatus.COMPLETED: "✅",
|
||||
QueueItemStatus.FAILED: "❌",
|
||||
QueueItemStatus.ABORTED: "🛑",
|
||||
QueueItemStatus.SKIPPED: "⏭️",
|
||||
}.get(item.status, "❓")
|
||||
|
||||
display_text = f"{status_emoji} {item.sample_name or item.item_id}"
|
||||
list_item = QListWidgetItem(display_text)
|
||||
|
||||
if item.status == QueueItemStatus.RUNNING:
|
||||
list_item.setBackground(QColor("#E6F3FF"))
|
||||
elif item.status == QueueItemStatus.COMPLETED:
|
||||
list_item.setBackground(QColor("#E6FFE6"))
|
||||
elif item.status == QueueItemStatus.FAILED:
|
||||
list_item.setBackground(QColor("#FFE6E6"))
|
||||
|
||||
self._queue_list.addItem(list_item)
|
||||
item_ids.append(item.item_id)
|
||||
|
||||
self._queue_list.set_item_ids(item_ids)
|
||||
self._update_button_states()
|
||||
|
||||
@Slot(object)
|
||||
def update_runtime(self, runtime: RuntimeState) -> None:
|
||||
"""Update from runtime state."""
|
||||
self._runtime = runtime
|
||||
|
||||
if runtime.running:
|
||||
if runtime.paused:
|
||||
self._status_label.setText("⏸️ Paused")
|
||||
self._status_label.setStyleSheet(
|
||||
"font-size: 14px; font-weight: bold; color: #FF9800;"
|
||||
)
|
||||
else:
|
||||
self._status_label.setText("▶️ Running")
|
||||
self._status_label.setStyleSheet(
|
||||
"font-size: 14px; font-weight: bold; color: #4CAF50;"
|
||||
)
|
||||
|
||||
self._current_item_label.setText(
|
||||
f"Item: {runtime.current_item_id or 'Unknown'}"
|
||||
)
|
||||
else:
|
||||
self._status_label.setText("⏹️ Idle")
|
||||
self._status_label.setStyleSheet(
|
||||
"font-size: 14px; font-weight: bold; color: #757575;"
|
||||
)
|
||||
self._current_item_label.setText("No item running")
|
||||
|
||||
self._update_button_states()
|
||||
|
||||
@Slot(object)
|
||||
def update_control(self, control: ControlState) -> None:
|
||||
"""Update from control state."""
|
||||
self._control = control
|
||||
self._update_button_states()
|
||||
|
||||
@Slot(object)
|
||||
def update_current_item(self, item: QueueItem) -> None:
|
||||
"""Update step progress for current item."""
|
||||
self._step_progress.set_steps(item.steps, item.current_step_index)
|
||||
|
||||
for i, step in enumerate(item.steps):
|
||||
self._step_progress.update_step(i, step.status)
|
||||
|
||||
@Slot(bool)
|
||||
def update_automation_enabled(self, enabled: bool) -> None:
|
||||
"""Update automation mode state."""
|
||||
self._automation_enabled = enabled
|
||||
self._auto_btn.setChecked(enabled)
|
||||
self._guided_btn.setChecked(not enabled)
|
||||
self._update_button_states()
|
||||
|
||||
@Slot(list)
|
||||
def update_available_samples(self, samples: list[SampleShortInfo]) -> None:
|
||||
"""Update the cached sample list for 'Add All' functionality."""
|
||||
self._all_samples = samples
|
||||
|
||||
@Slot(str, str)
|
||||
def on_workflow_event(self, event_type: str, message: str) -> None:
|
||||
"""Handle workflow events from SSE stream."""
|
||||
logger.debug(f"Workflow event: {event_type} - {message}")
|
||||
|
||||
if event_type in ("item_started", "item_completed", "step_completed", "step_skipped"):
|
||||
self.request_queue_refresh.emit()
|
||||
@@ -51,6 +51,14 @@ class DAQWorker(QObject):
|
||||
last_error_payload_changed = Signal(dict)
|
||||
last_error_payloads_changed = Signal(list)
|
||||
|
||||
# Workflow signals
|
||||
workflow_queue_loaded = Signal(list) # list[QueueItem]
|
||||
workflow_runtime_changed = Signal(object) # RuntimeState
|
||||
workflow_control_changed = Signal(object) # ControlState
|
||||
workflow_item_updated = Signal(object) # QueueItem
|
||||
workflow_automation_status = Signal(bool) # enabled
|
||||
workflow_event = Signal(object) # WorkflowEvent
|
||||
|
||||
def __init__(self, base_url: str | None, token: str, parent=None):
|
||||
super().__init__(parent)
|
||||
self._active_status_error_key = None
|
||||
@@ -1252,4 +1260,162 @@ class DAQWorker(QObject):
|
||||
query.append(f"message={quote(message)}")
|
||||
|
||||
suffix = f"?{'&'.join(query)}" if query else ""
|
||||
self.generic_post(f"samcam/send_screenshot_db{suffix}")
|
||||
self.generic_post(f"samcam/send_screenshot_db{suffix}")
|
||||
|
||||
# ─────────────────────────────────────────────
|
||||
# Workflow API methods
|
||||
# ─────────────────────────────────────────────
|
||||
|
||||
@Slot()
|
||||
def workflow_load_queue(self):
|
||||
"""Load workflow queue."""
|
||||
if self.__base_url is None:
|
||||
return
|
||||
|
||||
request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/queue"))
|
||||
request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode())
|
||||
reply = self.__net_manager.get(request)
|
||||
reply.finished.connect(lambda: self._handle_workflow_queue_response(reply))
|
||||
|
||||
def _handle_workflow_queue_response(self, reply: QNetworkReply):
|
||||
try:
|
||||
response_data = self.handle_response(reply)
|
||||
data = json.loads(response_data)
|
||||
from aare.common.automation_models import QueueItem
|
||||
items = [QueueItem.model_validate(i) for i in data.get("items", [])]
|
||||
self.workflow_queue_loaded.emit(items)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to load workflow queue: {e}")
|
||||
|
||||
@Slot(object) # SampleShortInfo
|
||||
def workflow_add_sample(self, sample: "SampleShortInfo"):
|
||||
"""Add a sample to the workflow queue."""
|
||||
if self.__base_url is None:
|
||||
logger.info(f"POST /workflow/queue: {sample.sample_name}")
|
||||
return
|
||||
|
||||
from aare.common.automation_models import CreateQueueItemRequest
|
||||
request_data = CreateQueueItemRequest(
|
||||
sample_id=sample.db_id,
|
||||
sample_name=sample.sample_name,
|
||||
priority=int(sample.priority or 100),
|
||||
)
|
||||
|
||||
request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/queue"))
|
||||
request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode())
|
||||
request.setRawHeader(b"Content-Type", b"application/json")
|
||||
reply = self.__net_manager.post(request, QByteArray(request_data.model_dump_json().encode()))
|
||||
reply.finished.connect(lambda: self.handle_req_response(reply))
|
||||
|
||||
@Slot(list) # list[SampleShortInfo]
|
||||
def workflow_add_samples(self, samples: list):
|
||||
"""Add multiple samples to the workflow queue."""
|
||||
for sample in samples:
|
||||
self.workflow_add_sample(sample)
|
||||
|
||||
@Slot(str)
|
||||
def workflow_delete_item(self, item_id: str):
|
||||
"""Delete an item from the workflow queue."""
|
||||
if self.__base_url is None:
|
||||
logger.info(f"DELETE /workflow/queue/{item_id}")
|
||||
return
|
||||
|
||||
self.generic_delete(f"workflow/queue/{item_id}")
|
||||
|
||||
@Slot(str, int)
|
||||
def workflow_move_item(self, item_id: str, new_order_index: int):
|
||||
"""Move/reorder an item in the workflow queue."""
|
||||
if self.__base_url is None:
|
||||
logger.info(f"POST /workflow/queue/{item_id}/move order={new_order_index}")
|
||||
return
|
||||
|
||||
body = json.dumps({"new_order_index": new_order_index})
|
||||
self.generic_post(f"workflow/queue/{item_id}/move", body)
|
||||
|
||||
@Slot()
|
||||
def workflow_clear_queue(self):
|
||||
"""Clear all pending items from the queue."""
|
||||
# Load queue first, then delete all pending items
|
||||
if self.__base_url is None:
|
||||
return
|
||||
|
||||
request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/queue"))
|
||||
request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode())
|
||||
reply = self.__net_manager.get(request)
|
||||
reply.finished.connect(lambda: self._handle_clear_queue_response(reply))
|
||||
|
||||
def _handle_clear_queue_response(self, reply: QNetworkReply):
|
||||
try:
|
||||
response_data = self.handle_response(reply)
|
||||
data = json.loads(response_data)
|
||||
from aare.common.automation_models import QueueItem, QueueItemStatus
|
||||
items = [QueueItem.model_validate(i) for i in data.get("items", [])]
|
||||
|
||||
# Delete all non-running items
|
||||
for item in items:
|
||||
if item.status != QueueItemStatus.RUNNING:
|
||||
self.workflow_delete_item(item.item_id)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to clear workflow queue: {e}")
|
||||
|
||||
@Slot(str)
|
||||
def workflow_start_item(self, item_id: str):
|
||||
"""Start processing a queue item."""
|
||||
self.generic_post(f"workflow/start/{item_id}")
|
||||
QTimer.singleShot(500, self.workflow_load_queue)
|
||||
|
||||
@Slot()
|
||||
def workflow_next_step(self):
|
||||
"""Run next step in guided mode."""
|
||||
if self.__base_url is None:
|
||||
return
|
||||
|
||||
request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/next"))
|
||||
request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode())
|
||||
request.setRawHeader(b"Content-Type", b"application/json")
|
||||
reply = self.__net_manager.post(request, QByteArray(b""))
|
||||
reply.finished.connect(lambda: self._handle_workflow_step_response(reply))
|
||||
|
||||
def _handle_workflow_step_response(self, reply: QNetworkReply):
|
||||
try:
|
||||
response_data = self.handle_response(reply)
|
||||
data = json.loads(response_data)
|
||||
if data.get("item"):
|
||||
from aare.common.automation_models import QueueItem
|
||||
item = QueueItem.model_validate(data["item"])
|
||||
self.workflow_item_updated.emit(item)
|
||||
# Also refresh queue
|
||||
self.workflow_load_queue()
|
||||
except Exception as e:
|
||||
logger.error(f"Workflow step error: {e}")
|
||||
self.http_error.emit(str(e))
|
||||
|
||||
@Slot()
|
||||
def workflow_pause(self):
|
||||
"""Request pause."""
|
||||
self.generic_post("workflow/control/pause")
|
||||
|
||||
@Slot()
|
||||
def workflow_resume(self):
|
||||
"""Request resume."""
|
||||
self.generic_post("workflow/control/resume")
|
||||
|
||||
@Slot()
|
||||
def workflow_abort(self):
|
||||
"""Request abort."""
|
||||
self.generic_post("workflow/control/abort")
|
||||
|
||||
@Slot()
|
||||
def workflow_skip(self):
|
||||
"""Request skip."""
|
||||
self.generic_post("workflow/control/skip")
|
||||
|
||||
@Slot()
|
||||
def workflow_start_automation(self):
|
||||
"""Start automation mode."""
|
||||
self.generic_post("workflow/automation/start")
|
||||
|
||||
@Slot()
|
||||
def workflow_stop_automation(self):
|
||||
"""Stop automation mode."""
|
||||
self.generic_post("workflow/automation/stop")
|
||||
@@ -0,0 +1,158 @@
|
||||
"""
|
||||
SSE client for workflow events.
|
||||
|
||||
Connects to the /workflow/sse endpoint and emits signals for:
|
||||
- Runtime state changes
|
||||
- Control state changes
|
||||
- Individual workflow events
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
from PySide6.QtCore import QObject, Signal, Slot, QTimer
|
||||
from PySide6.QtNetwork import QNetworkAccessManager, QNetworkRequest, QNetworkReply
|
||||
from PySide6.QtCore import QUrl, QByteArray
|
||||
|
||||
from aare.common.automation_models import RuntimeState, ControlState, WorkflowEvent
|
||||
from aare.common.logger_config import setup_logger
|
||||
|
||||
logger = setup_logger("aareGUI")
|
||||
|
||||
|
||||
class WorkflowSSEClient(QObject):
|
||||
"""
|
||||
SSE client that subscribes to workflow events.
|
||||
|
||||
Emits signals when state changes are received.
|
||||
"""
|
||||
|
||||
# Signals
|
||||
runtime_changed = Signal(object) # RuntimeState
|
||||
control_changed = Signal(object) # ControlState
|
||||
workflow_event = Signal(object) # WorkflowEvent
|
||||
connected = Signal()
|
||||
disconnected = Signal()
|
||||
error = Signal(str)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
token: str,
|
||||
parent: QObject | None = None,
|
||||
):
|
||||
super().__init__(parent)
|
||||
|
||||
self._base_url = base_url
|
||||
self._token = token
|
||||
self._manager = QNetworkAccessManager(self)
|
||||
self._reply: QNetworkReply | None = None
|
||||
self._buffer = ""
|
||||
|
||||
# Reconnection
|
||||
self._reconnect_timer = QTimer(self)
|
||||
self._reconnect_timer.setInterval(5000) # 5 seconds
|
||||
self._reconnect_timer.timeout.connect(self.connect)
|
||||
self._should_reconnect = False
|
||||
|
||||
def connect(self) -> None:
|
||||
"""Start SSE connection."""
|
||||
if self._reply is not None:
|
||||
return # Already connected
|
||||
|
||||
self._should_reconnect = True
|
||||
|
||||
url = QUrl(f"{self._base_url}/workflow/sse")
|
||||
# Add token as query param for SSE (can't use headers easily)
|
||||
url.setQuery(f"token={self._token}")
|
||||
|
||||
request = QNetworkRequest(url)
|
||||
request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode())
|
||||
request.setRawHeader(b"Accept", b"text/event-stream")
|
||||
request.setRawHeader(b"Cache-Control", b"no-cache")
|
||||
|
||||
self._reply = self._manager.get(request)
|
||||
self._reply.readyRead.connect(self._on_data_ready)
|
||||
self._reply.finished.connect(self._on_finished)
|
||||
self._reply.errorOccurred.connect(self._on_error)
|
||||
|
||||
self._reconnect_timer.stop()
|
||||
logger.debug("Workflow SSE: connecting...")
|
||||
|
||||
def disconnect(self) -> None:
|
||||
"""Stop SSE connection."""
|
||||
self._should_reconnect = False
|
||||
self._reconnect_timer.stop()
|
||||
|
||||
if self._reply is not None:
|
||||
self._reply.abort()
|
||||
self._reply.deleteLater()
|
||||
self._reply = None
|
||||
|
||||
self.disconnected.emit()
|
||||
|
||||
@Slot()
|
||||
def _on_data_ready(self) -> None:
|
||||
"""Handle incoming SSE data."""
|
||||
if self._reply is None:
|
||||
return
|
||||
|
||||
data = self._reply.readAll().data().decode("utf-8")
|
||||
self._buffer += data
|
||||
|
||||
# Process complete events (separated by double newlines)
|
||||
while "\n\n" in self._buffer:
|
||||
event_data, self._buffer = self._buffer.split("\n\n", 1)
|
||||
self._parse_event(event_data)
|
||||
|
||||
def _parse_event(self, event_data: str) -> None:
|
||||
"""Parse a single SSE event."""
|
||||
event_type = "message"
|
||||
data_lines = []
|
||||
|
||||
for line in event_data.split("\n"):
|
||||
if line.startswith("event:"):
|
||||
event_type = line[6:].strip()
|
||||
elif line.startswith("data:"):
|
||||
data_lines.append(line[5:].strip())
|
||||
|
||||
if not data_lines:
|
||||
return
|
||||
|
||||
data_str = "\n".join(data_lines)
|
||||
|
||||
try:
|
||||
if event_type == "runtime":
|
||||
runtime = RuntimeState.model_validate_json(data_str)
|
||||
self.runtime_changed.emit(runtime)
|
||||
elif event_type == "control":
|
||||
control = ControlState.model_validate_json(data_str)
|
||||
self.control_changed.emit(control)
|
||||
elif event_type == "workflow_event":
|
||||
event = WorkflowEvent.model_validate_json(data_str)
|
||||
self.workflow_event.emit(event)
|
||||
except Exception as e:
|
||||
logger.warning(f"Workflow SSE: failed to parse {event_type}: {e}")
|
||||
|
||||
@Slot()
|
||||
def _on_finished(self) -> None:
|
||||
"""Handle connection finished."""
|
||||
if self._reply is not None:
|
||||
self._reply.deleteLater()
|
||||
self._reply = None
|
||||
|
||||
self._buffer = ""
|
||||
self.disconnected.emit()
|
||||
|
||||
# Reconnect if desired
|
||||
if self._should_reconnect:
|
||||
logger.debug("Workflow SSE: disconnected, will reconnect...")
|
||||
self._reconnect_timer.start()
|
||||
|
||||
@Slot(QNetworkReply.NetworkError)
|
||||
def _on_error(self, error: QNetworkReply.NetworkError) -> None:
|
||||
"""Handle connection error."""
|
||||
error_msg = self._reply.errorString() if self._reply else str(error)
|
||||
logger.warning(f"Workflow SSE error: {error_msg}")
|
||||
self.error.emit(error_msg)
|
||||
Reference in New Issue
Block a user