Major update from x10sa #48

Merged
appleb_m merged 58 commits from x10sa into master 2026-03-27 12:00:23 +01:00
39 changed files with 6892 additions and 1018 deletions
+1 -1
View File
@@ -18,7 +18,7 @@ dependencies = [
"python-redis-lock==4.0.0",
"fastapi==0.115.13",
"uvicorn==0.34.2",
"aaredb==0.1.1a42",
"aaredb==0.1.1a43",
"python_multipart==0.0.20",
"websocket-client==1.8.0",
"sseclient-py==1.8.0",
+155
View File
@@ -0,0 +1,155 @@
from enum import Enum
from typing import Optional
from pydantic import BaseModel, ConfigDict, Field
from aare.common.rotation_scan import RotationScanRequest
class TaskEnum(Enum):
TASK_0 = 0
TASK_1 = 1
TASK_2 = 2
TASK_3 = 3
TASK_4 = 4
TASK_5 = 5
TASK_6 = 6
TASK_7 = 7
TASK_8 = 8
TASK_9 = 9
class AxisEnum(Enum):
X = "x"
Y = "y"
Z = "z"
OMEGA = "u"
class AerotechRunEnum(Enum):
STOP = 0
START = 1
RUN = 2
LOAD = 3
PAUSE = 4
RESET = 5
class VariableTypeEnum(Enum):
INT = 0
REAL = 1
STRING = 2
class AerotechAxisStatus(BaseModel):
enabled: bool
fault: int
homed: bool
is_fault: bool
moving: bool
position: float
status: int
velocity: float
class AerotechStatus(BaseModel):
state: str
x: Optional[AerotechAxisStatus] = None
y: Optional[AerotechAxisStatus] = None
z: Optional[AerotechAxisStatus] = None
u: Optional[AerotechAxisStatus] = None
model_config = ConfigDict(extra="allow")
def __str__(self) -> str:
return self.to_pretty_string()
def to_pretty_string(self) -> str:
def axis_line(name: str, axis: AerotechAxisStatus | None) -> str:
if axis is None:
return f"{name.upper()}: unavailable"
return (
f"{name.upper():>2} | pos={axis.position:>12.6f} | "
f"homed={axis.homed!s:<5} | moving={axis.moving!s:<5} | "
f"enabled={axis.enabled!s:<5} | fault={axis.fault} | "
f"faulted={axis.is_fault!s:<5} | vel={axis.velocity:>10.6f}"
)
return "\n".join(
[
f"STATE: {self.state}",
axis_line("x", self.x),
axis_line("y", self.y),
axis_line("z", self.z),
axis_line("u", self.u),
]
)
def to_compact_string(self) -> str:
axes = []
for name in ("x", "y", "z", "u"):
axis = getattr(self, name)
if axis is not None:
axes.append(
f"{name}={axis.position:.4f} "
f"({'H' if axis.homed else 'NH'}, {'M' if axis.moving else '-'})"
)
return f"state={self.state} | " + " | ".join(axes)
def to_colored_string(self) -> str:
# ANSI colors for terminal use
RESET = "\033[0m"
BOLD = "\033[1m"
CYAN = "\033[36m"
GREEN = "\033[32m"
YELLOW = "\033[33m"
RED = "\033[31m"
def color_bool(value: bool) -> str:
return f"{GREEN}True{RESET}" if value else f"{RED}False{RESET}"
def axis_line(name: str, axis: AerotechAxisStatus | None) -> str:
if axis is None:
return f"{YELLOW}{name.upper()}: unavailable{RESET}"
return (
f"{BOLD}{name.upper()}{RESET} | "
f"pos={CYAN}{axis.position:>12.6f}{RESET} | "
f"homed={color_bool(axis.homed)} | "
f"moving={color_bool(axis.moving)} | "
f"enabled={color_bool(axis.enabled)} | "
f"fault={axis.fault} | "
f"faulted={color_bool(axis.is_fault)} | "
f"vel={axis.velocity:>10.6f}"
)
return "\n".join(
[
f"{BOLD}STATE:{RESET} {CYAN}{self.state}{RESET}",
axis_line("x", self.x),
axis_line("y", self.y),
axis_line("z", self.z),
axis_line("u", self.u),
]
)
class AerotechTarget(BaseModel):
x: Optional[float] = None
y: Optional[float] = None
z: Optional[float] = None
u: Optional[float] = None
def to_payload(self) -> dict:
return self.model_dump(exclude_none=True)
class AerotechRotationScanRequest(RotationScanRequest):
rotation_deg: float
time_sec: float
start_pos_deg: float
async_move: bool = Field(default=False, alias="async")
model_config = ConfigDict(populate_by_name=True)
def to_payload(self) -> dict:
return self.model_dump(by_alias=True, exclude_none=True)
+52
View File
@@ -0,0 +1,52 @@
from enum import Enum
from pydantic import BaseModel
from datetime import datetime
class BatonRequestStatus(Enum):
PENDING = "pending"
ACCEPTED = "accepted"
REFUSED = "refused"
TIMEOUT = "timeout"
CANCELLED = "cancelled"
class BatonHolderInfo(BaseModel):
"""Information about the current baton holder."""
username: str
session: int
is_staff: bool
pgroup: str | None = None
class BatonRequest(BaseModel):
"""A request from one user to take the baton from another."""
request_id: str
requester_username: str
requester_session: int
requester_is_staff: bool
holder_username: str | None = None
holder_session: int | None = None
created_at: float # Unix timestamp
timeout_seconds: int = 30
status: BatonRequestStatus = BatonRequestStatus.PENDING
class BatonTransferQueue(BaseModel):
"""Queued baton transfer waiting for beamline to be available."""
target_session: int
target_username: str
target_is_staff: bool
target_pgroup: str | None = None
queued_at: float # Unix timestamp
reason: str = "beamline_busy"
class BatonStatus(BaseModel):
"""Full baton status for GUI display."""
holder: BatonHolderInfo | None = None
pending_request: BatonRequest | None = None
queued_transfer: BatonTransferQueue | None = None
you_are_holder: bool = False
you_have_pending_request: bool = False
incoming_request: bool = False
allow_non_staff_request: bool = False
+182
View File
@@ -0,0 +1,182 @@
from __future__ import annotations
from dataclasses import dataclass, field
from enum import Enum
from typing import Any
import time
import uuid
from pydantic import BaseModel, Field
class WorkflowMode(str, Enum):
FLEXIBLE_MANUAL = "flexible_manual"
GUIDED_MANUAL = "guided_manual"
AUTOMATION = "automation"
class WorkflowStateKind(str, Enum):
MOUNT = "mount"
LOOP_CENTRE = "loop_centre"
RASTER = "raster"
DATA_COLLECTION = "data_collection"
class StepStatus(str, Enum):
PENDING = "pending"
RUNNING = "running"
SUCCESS = "success"
FAILED = "failed"
SKIPPED = "skipped"
PAUSED = "paused"
@dataclass(frozen=True)
class TransitionRule:
to_state: WorkflowStateKind
allowed_modes: frozenset[WorkflowMode] = frozenset()
optional: bool = False
condition_name: str | None = None
@dataclass(frozen=True)
class StateDefinition:
kind: WorkflowStateKind
transitions: tuple[TransitionRule, ...]
description: str = ""
@dataclass
class StateResult:
state: WorkflowStateKind
status: StepStatus
message: str = ""
payload: dict[str, Any] = field(default_factory=dict)
@dataclass
class WorkflowContext:
mode: WorkflowMode
queue_id: str
item_id: str
sample_id: int | None = None
current_state: WorkflowStateKind | None = None
current_step_index: int = 0
paused: bool = False
abort_requested: bool = False
last_message: str = ""
metadata: dict[str, Any] = field(default_factory=dict)
class QueueItemStatus(str, Enum):
PENDING = "pending"
RUNNING = "running"
COMPLETED = "completed"
FAILED = "failed"
SKIPPED = "skipped"
ABORTED = "aborted"
class WorkflowStepRecord(BaseModel):
kind: str
status: str = "pending"
message: str = ""
started_at: float | None = None
completed_at: float | None = None
error_detail: str | None = None
class QueueItem(BaseModel):
item_id: str
beamline: str
sample_id: int | None = None
sample_name: str = ""
owner_pgroup: str = ""
created_by: str = ""
created_at: float = Field(default_factory=time.time)
priority: int = 100
order_index: int = 0
status: QueueItemStatus = QueueItemStatus.PENDING
steps: list[WorkflowStepRecord] = Field(default_factory=list)
current_step_index: int = 0
recipe: dict[str, Any] = Field(default_factory=dict)
metadata: dict[str, Any] = Field(default_factory=dict)
class RuntimeState(BaseModel):
running: bool = False
paused: bool = False
current_queue_id: str = ""
current_item_id: str | None = None
current_state: str | None = None
current_step_index: int = 0
last_error: str | None = None
last_update: float = Field(default_factory=time.time)
class ControlState(BaseModel):
pause_requested: bool = False
resume_requested: bool = False
abort_requested: bool = False
skip_requested: bool = False
next_sample_requested: bool = False
requested_by: str | None = None
requested_at: float | None = None
class WorkflowEvent(BaseModel):
event_id: str = Field(default_factory=lambda: str(uuid.uuid4()))
beamline: str = ""
item_id: str | None = None
step: str | None = None
event_type: str = ""
timestamp: float = Field(default_factory=time.time)
actor: str = ""
message: str = ""
payload: dict[str, Any] = Field(default_factory=dict)
class CreateQueueItemRequest(BaseModel):
sample_id: int | None = None
sample_name: str = ""
priority: int = 100
recipe: dict = Field(default_factory=dict)
steps: list[str] | None = None # If None, use default steps
class MoveItemRequest(BaseModel):
new_order_index: int
class QueueListResponse(BaseModel):
items: list[QueueItem]
total: int
class RuntimeResponse(BaseModel):
runtime: RuntimeState
control: ControlState
class ControlActionResponse(BaseModel):
ok: bool
control: ControlState
message: str = ""
class StepActionResponse(BaseModel):
ok: bool
item: QueueItem | None = None
step: str | None = None
status: str = ""
message: str = ""
class EventListResponse(BaseModel):
events: list[WorkflowEvent]
class AutomationStatusResponse(BaseModel):
enabled: bool
running: bool
runtime: RuntimeState
control: ControlState
+245
View File
@@ -0,0 +1,245 @@
from __future__ import annotations
import time
from typing import Any
import redis
from aare.common.automation_models import (
QueueItem,
WorkflowEvent,
ControlState,
RuntimeState,
QueueItemStatus,
WorkflowStepRecord,
WorkflowStateKind,
)
CURRENT_QUEUE_LIMIT = 576
def build_default_steps() -> list[WorkflowStepRecord]:
return [
WorkflowStepRecord(kind=WorkflowStateKind.MOUNT.value),
WorkflowStepRecord(kind=WorkflowStateKind.LOOP_CENTRE.value),
WorkflowStepRecord(kind=WorkflowStateKind.RASTER.value),
WorkflowStepRecord(kind=WorkflowStateKind.DATA_COLLECTION.value),
]
class WorkflowRedisManager:
def __init__(self, client: redis.Redis, beamline: str):
self._client = client
self._bl = beamline.lower()
def _key(self, suffix: str) -> str:
return f"{self._bl}:workflow:{suffix}"
def _item_key(self, item_id: str) -> str:
return self._key(f"item:{item_id}")
def _queue_key(self) -> str:
return self._key("queue")
def _runtime_key(self) -> str:
return self._key("runtime")
def _control_key(self) -> str:
return self._key("control")
def _events_key(self) -> str:
return self._key("events")
def _next_item_id(self) -> str:
n = int(self._client.incr(self._key("item_seq")))
return f"wf_{int(time.time())}_{n:06d}"
# ─────────────────────────────────────────────
# Queue item operations
# ─────────────────────────────────────────────
def create_item(self, item: QueueItem) -> QueueItem:
if not item.item_id:
item.item_id = self._next_item_id()
current_count = self._client.zcard(self._queue_key())
if current_count >= CURRENT_QUEUE_LIMIT:
raise ValueError(f"Queue size limit reached ({CURRENT_QUEUE_LIMIT} items). Please clear the queue.")
pipe = self._client.pipeline(transaction=True)
pipe.set(self._item_key(item.item_id), item.model_dump_json())
pipe.zadd(self._queue_key(), {item.item_id: float(item.order_index)})
pipe.execute()
self.append_event(WorkflowEvent(
beamline=self._bl,
item_id=item.item_id,
event_type="item_created",
message="Queue item created",
payload={"status": item.status.value},
))
return item
def get_item(self, item_id: str) -> QueueItem | None:
raw = self._client.get(self._item_key(item_id))
if raw is None:
return None
return QueueItem.model_validate_json(raw)
def update_item(self, item_id: str, patch: dict[str, Any]) -> QueueItem:
item = self.get_item(item_id)
if item is None:
raise KeyError(f"Queue item not found: {item_id}")
updated = item.model_copy(update=patch)
self._client.set(self._item_key(item_id), updated.model_dump_json())
return updated
def delete_item(self, item_id: str) -> None:
pipe = self._client.pipeline(transaction=True)
pipe.delete(self._item_key(item_id))
pipe.zrem(self._queue_key(), item_id)
pipe.execute()
def list_queue_order(self) -> list[str]:
return [str(x) for x in self._client.zrange(self._queue_key(), 0, -1)]
def list_items(self, *, include_finished: bool = True) -> list[QueueItem]:
out: list[QueueItem] = []
for item_id in self.list_queue_order():
item = self.get_item(item_id)
if item is None:
continue
if not include_finished and item.status in {
QueueItemStatus.COMPLETED,
QueueItemStatus.FAILED,
QueueItemStatus.SKIPPED,
QueueItemStatus.ABORTED,
}:
continue
out.append(item)
return out
def get_next_pending_item(self) -> QueueItem | None:
for item_id in self.list_queue_order():
item = self.get_item(item_id)
if item is not None and item.status == QueueItemStatus.PENDING:
return item
return None
def move_item(self, item_id: str, new_order_index: int) -> None:
if self.get_item(item_id) is None:
raise KeyError(f"Queue item not found: {item_id}")
self._client.zadd(self._queue_key(), {item_id: float(new_order_index)})
self.update_item(item_id, {"order_index": new_order_index})
# ─────────────────────────────────────────────
# Runtime state
# ─────────────────────────────────────────────
def get_runtime(self) -> RuntimeState:
raw = self._client.get(self._runtime_key())
if raw is None:
return RuntimeState()
return RuntimeState.model_validate_json(raw)
def set_runtime(self, runtime: RuntimeState) -> RuntimeState:
runtime.last_update = time.time()
self._client.set(self._runtime_key(), runtime.model_dump_json())
return runtime
def patch_runtime(self, patch: dict[str, Any]) -> RuntimeState:
runtime = self.get_runtime()
updated = runtime.model_copy(update=patch)
return self.set_runtime(updated)
# ─────────────────────────────────────────────
# Control state
# ─────────────────────────────────────────────
def get_control(self) -> ControlState:
raw = self._client.get(self._control_key())
if raw is None:
return ControlState()
return ControlState.model_validate_json(raw)
def request_control(self, patch: dict[str, Any], *, requested_by: str) -> ControlState:
control = self.get_control()
updated = control.model_copy(update={
**patch,
"requested_by": requested_by,
"requested_at": time.time(),
})
self._client.set(self._control_key(), updated.model_dump_json())
self.append_event(WorkflowEvent(
beamline=self._bl,
item_id=self.get_runtime().current_item_id,
event_type="control_requested",
actor=requested_by,
payload=patch,
))
return updated
def clear_control(self) -> ControlState:
control = ControlState()
self._client.set(self._control_key(), control.model_dump_json())
return control
# ─────────────────────────────────────────────
# Event log
# ─────────────────────────────────────────────
def append_event(self, event: WorkflowEvent) -> None:
payload = event.model_dump_json()
self._client.xadd(self._events_key(), {"json": payload}, maxlen=5000, approximate=True)
def read_events(self, *, limit: int = 200) -> list[WorkflowEvent]:
rows = self._client.xrevrange(self._events_key(), count=limit)
out: list[WorkflowEvent] = []
for _, fields in reversed(rows):
raw = fields.get("json")
if raw:
out.append(WorkflowEvent.model_validate_json(raw))
return out
# ─────────────────────────────────────────────
# Step status updates
# ─────────────────────────────────────────────
def update_step(
self,
item_id: str,
step_index: int,
*,
status: str | None = None,
message: str | None = None,
error_detail: str | None = None,
) -> QueueItem:
item = self.get_item(item_id)
if item is None:
raise KeyError(f"Queue item not found: {item_id}")
if not (0 <= step_index < len(item.steps)):
raise IndexError(f"Step index out of range: {step_index}")
step = item.steps[step_index]
if status is not None:
step.status = status
if status == "running" and step.started_at is None:
step.started_at = time.time()
elif status in ("completed", "failed", "skipped"):
step.completed_at = time.time()
if message is not None:
step.message = message
if error_detail is not None:
step.error_detail = error_detail
item.steps[step_index] = step
self._client.set(self._item_key(item_id), item.model_dump_json())
return item
+351
View File
@@ -0,0 +1,351 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from aare.common.automation_models import (
WorkflowStateKind,
WorkflowMode,
StepStatus,
StateResult,
WorkflowContext,
TransitionRule,
StateDefinition,
)
STATE_REGISTRY: dict[WorkflowStateKind, StateDefinition] = {
WorkflowStateKind.MOUNT: StateDefinition(
kind=WorkflowStateKind.MOUNT,
description="Mount the sample",
transitions=(
TransitionRule(
to_state=WorkflowStateKind.LOOP_CENTRE,
allowed_modes=frozenset({
WorkflowMode.FLEXIBLE_MANUAL,
WorkflowMode.GUIDED_MANUAL,
WorkflowMode.AUTOMATION,
}),
),
),
),
WorkflowStateKind.LOOP_CENTRE: StateDefinition(
kind=WorkflowStateKind.LOOP_CENTRE,
description="Centre the loop",
transitions=(
TransitionRule(
to_state=WorkflowStateKind.RASTER,
allowed_modes=frozenset({
WorkflowMode.FLEXIBLE_MANUAL,
WorkflowMode.GUIDED_MANUAL,
WorkflowMode.AUTOMATION,
}),
),
TransitionRule(
to_state=WorkflowStateKind.DATA_COLLECTION,
allowed_modes=frozenset({
WorkflowMode.FLEXIBLE_MANUAL,
WorkflowMode.GUIDED_MANUAL,
}),
optional=True,
),
),
),
WorkflowStateKind.RASTER: StateDefinition(
kind=WorkflowStateKind.RASTER,
description="Run raster scan",
transitions=(
TransitionRule(
to_state=WorkflowStateKind.DATA_COLLECTION,
allowed_modes=frozenset({
WorkflowMode.FLEXIBLE_MANUAL,
WorkflowMode.GUIDED_MANUAL,
WorkflowMode.AUTOMATION,
}),
),
),
),
WorkflowStateKind.DATA_COLLECTION: StateDefinition(
kind=WorkflowStateKind.DATA_COLLECTION,
description="Collect diffraction data",
transitions=(),
),
}
def get_state_definition(kind: WorkflowStateKind) -> StateDefinition:
try:
return STATE_REGISTRY[kind]
except KeyError as exc:
raise KeyError(f"Unknown workflow state: {kind}") from exc
def get_allowed_next_states(
kind: WorkflowStateKind,
mode: WorkflowMode | None = None,
) -> list[WorkflowStateKind]:
definition = get_state_definition(kind)
out: list[WorkflowStateKind] = []
for transition in definition.transitions:
if mode is None:
out.append(transition.to_state)
continue
if not transition.allowed_modes or mode in transition.allowed_modes:
out.append(transition.to_state)
return out
def can_transition(
from_state: WorkflowStateKind,
to_state: WorkflowStateKind,
mode: WorkflowMode | None = None,
) -> bool:
return to_state in get_allowed_next_states(from_state, mode=mode)
class StateHandler(ABC):
state_kind: WorkflowStateKind
def __init__(self, registry: dict[WorkflowStateKind, StateDefinition] | None = None):
self._registry = registry or STATE_REGISTRY
def definition(self) -> StateDefinition:
return get_state_definition(self.state_kind)
def can_run(self, context: WorkflowContext) -> bool:
return context.current_state in (None, self.state_kind)
@abstractmethod
def validate(self, context: WorkflowContext) -> None:
pass
@abstractmethod
def execute(self, context: WorkflowContext) -> StateResult:
pass
class MountHandler(StateHandler):
state_kind = WorkflowStateKind.MOUNT
def validate(self, context: WorkflowContext) -> None:
if context.abort_requested:
raise RuntimeError("Abort requested; cannot mount.")
def execute(self, context: WorkflowContext) -> StateResult:
self.validate(context)
context.current_state = WorkflowStateKind.MOUNT
context.last_message = "Sample mounted"
return StateResult(
state=self.state_kind,
status=StepStatus.SUCCESS,
message="Sample mounted successfully.",
payload={"mounted": True},
)
class LoopCentreHandler(StateHandler):
state_kind = WorkflowStateKind.LOOP_CENTRE
def validate(self, context: WorkflowContext) -> None:
if context.abort_requested:
raise RuntimeError("Abort requested; cannot loop-centre.")
def execute(self, context: WorkflowContext) -> StateResult:
self.validate(context)
context.current_state = WorkflowStateKind.LOOP_CENTRE
context.last_message = "Loop centred"
return StateResult(
state=self.state_kind,
status=StepStatus.SUCCESS,
message="Loop centring completed.",
payload={"centred": True},
)
class RasterHandler(StateHandler):
state_kind = WorkflowStateKind.RASTER
def validate(self, context: WorkflowContext) -> None:
if context.abort_requested:
raise RuntimeError("Abort requested; cannot raster.")
def execute(self, context: WorkflowContext) -> StateResult:
self.validate(context)
context.current_state = WorkflowStateKind.RASTER
context.last_message = "Raster completed"
return StateResult(
state=self.state_kind,
status=StepStatus.SUCCESS,
message="Raster scan completed.",
payload={"best_spot_found": True},
)
class DataCollectionHandler(StateHandler):
state_kind = WorkflowStateKind.DATA_COLLECTION
def validate(self, context: WorkflowContext) -> None:
if context.abort_requested:
raise RuntimeError("Abort requested; cannot collect data.")
def execute(self, context: WorkflowContext) -> StateResult:
self.validate(context)
context.current_state = WorkflowStateKind.DATA_COLLECTION
context.last_message = "Data collected"
return StateResult(
state=self.state_kind,
status=StepStatus.SUCCESS,
message="Data collection completed.",
payload={"frames_collected": 1},
)
class WorkflowRunner:
"""Simple in-memory runner (no persistence)."""
def __init__(
self,
registry: dict[WorkflowStateKind, StateDefinition] | None = None,
handlers: dict[WorkflowStateKind, StateHandler] | None = None,
):
self._registry = registry or STATE_REGISTRY
self._handlers = handlers or HANDLER_REGISTRY
def get_handler(self, state: WorkflowStateKind) -> StateHandler:
try:
return self._handlers[state]
except KeyError as exc:
raise KeyError(f"No handler registered for state: {state}") from exc
def can_move_to(
self,
current: WorkflowStateKind,
next_state: WorkflowStateKind,
mode: WorkflowMode,
) -> bool:
return can_transition(current, next_state, mode)
def run_state(
self,
context: WorkflowContext,
state: WorkflowStateKind,
) -> StateResult:
if context.current_state is not None:
if not self.can_move_to(context.current_state, state, context.mode):
raise RuntimeError(
f"Transition not allowed: {context.current_state} -> {state}"
)
handler = self.get_handler(state)
result = handler.execute(context)
context.current_state = state
context.current_step_index += 1
context.last_message = result.message
return result
class SimulatedMountHandler(StateHandler):
"""Simulated mount handler for testing - doesn't actually mount."""
state_kind = WorkflowStateKind.MOUNT
def validate(self, context: WorkflowContext) -> None:
if context.abort_requested:
raise RuntimeError("Abort requested; cannot mount.")
def execute(self, context: WorkflowContext) -> StateResult:
self.validate(context)
import time
time.sleep(10) # Simulate some work
context.current_state = WorkflowStateKind.MOUNT
context.last_message = "[SIMULATION] Sample would be mounted"
return StateResult(
state=self.state_kind,
status=StepStatus.SUCCESS,
message=f"[SIMULATION] Would mount sample_id={context.sample_id}",
payload={"mounted": True, "simulated": True},
)
class SimulatedLoopCentreHandler(StateHandler):
"""Simulated loop centre handler for testing."""
state_kind = WorkflowStateKind.LOOP_CENTRE
def validate(self, context: WorkflowContext) -> None:
if context.abort_requested:
raise RuntimeError("Abort requested; cannot loop-centre.")
def execute(self, context: WorkflowContext) -> StateResult:
self.validate(context)
import time
time.sleep(10)
context.current_state = WorkflowStateKind.LOOP_CENTRE
context.last_message = "[SIMULATION] Loop would be centred"
return StateResult(
state=self.state_kind,
status=StepStatus.SUCCESS,
message="[SIMULATION] Would run loop centering algorithm",
payload={"centred": True, "simulated": True},
)
class SimulatedRasterHandler(StateHandler):
"""Simulated raster handler for testing."""
state_kind = WorkflowStateKind.RASTER
def validate(self, context: WorkflowContext) -> None:
if context.abort_requested:
raise RuntimeError("Abort requested; cannot raster.")
def execute(self, context: WorkflowContext) -> StateResult:
self.validate(context)
import time
time.sleep(10)
context.current_state = WorkflowStateKind.RASTER
context.last_message = "[SIMULATION] Raster scan would be performed"
return StateResult(
state=self.state_kind,
status=StepStatus.SUCCESS,
message="[SIMULATION] Would run raster scan, find best diffraction spot",
payload={"best_spot_found": True, "simulated": True},
)
class SimulatedDataCollectionHandler(StateHandler):
"""Simulated data collection handler for testing."""
state_kind = WorkflowStateKind.DATA_COLLECTION
def validate(self, context: WorkflowContext) -> None:
if context.abort_requested:
raise RuntimeError("Abort requested; cannot collect data.")
def execute(self, context: WorkflowContext) -> StateResult:
self.validate(context)
import time
time.sleep(10)
context.current_state = WorkflowStateKind.DATA_COLLECTION
context.last_message = "[SIMULATION] Data collection would be performed"
return StateResult(
state=self.state_kind,
status=StepStatus.SUCCESS,
message="[SIMULATION] Would collect 1800 frames at 0.2° oscillation",
payload={"frames_collected": 1800, "simulated": True},
)
# Simulated handler registry for testing
SIMULATED_HANDLER_REGISTRY: dict[WorkflowStateKind, StateHandler] = {
WorkflowStateKind.MOUNT: SimulatedMountHandler(STATE_REGISTRY),
WorkflowStateKind.LOOP_CENTRE: SimulatedLoopCentreHandler(STATE_REGISTRY),
WorkflowStateKind.RASTER: SimulatedRasterHandler(STATE_REGISTRY),
WorkflowStateKind.DATA_COLLECTION: SimulatedDataCollectionHandler(STATE_REGISTRY),
}
HANDLER_REGISTRY: dict[WorkflowStateKind, StateHandler] = {
WorkflowStateKind.MOUNT: MountHandler(STATE_REGISTRY),
WorkflowStateKind.LOOP_CENTRE: LoopCentreHandler(STATE_REGISTRY),
WorkflowStateKind.RASTER: RasterHandler(STATE_REGISTRY),
WorkflowStateKind.DATA_COLLECTION: DataCollectionHandler(STATE_REGISTRY),
}
+111 -1
View File
@@ -1,5 +1,7 @@
from __future__ import annotations
import time
from aare.common.logger_config import setup_logger
from aare.common.error_codes import AuthErrorCode, DAQErrorCode
@@ -104,6 +106,8 @@ class AuthenticationException(Exception):
class UserRightsException(Exception):
_last_log_ts_by_message: dict[str, float] = {}
_throttle_window_s = 30.0
def __init__(self,
message: str = "User does not have rights to perform this action.",
*,
@@ -115,7 +119,12 @@ class UserRightsException(Exception):
self.status_code = status_code
self.headers = headers
self.code = code
logger.error(message, extra={"exception:": Exception})
now = time.monotonic()
last_ts = self._last_log_ts_by_message.get(message, 0.0)
if (now - last_ts) >= self._throttle_window_s:
self._last_log_ts_by_message[message] = now
logger.warning(message, extra={"exception:": Exception})
def __str__(self) -> str:
return self.message
@@ -186,5 +195,106 @@ class TellCommunicationError(Exception):
},
)
def __str__(self) -> str:
return self.message
class JFJochCommunicationError(Exception):
"""
Raised when JFJoch HTTP/API communication fails.
Intended for scan-time fallbacks and GUI-visible alerts.
"""
def __init__(
self,
message: str = "JFJoch communication error",
*,
operation: str | None = None,
endpoint: str | None = None,
base_url: str | None = None,
status_code: int | None = None,
):
super().__init__(message)
self.message = message
self.operation = operation
self.endpoint = endpoint
self.base_url = base_url
self.status_code = status_code
logger.error(
message,
extra={
"device": "jfjjoch",
"operation": operation,
"endpoint": endpoint,
"base_url": base_url,
"status_code": status_code,
},
)
def __str__(self) -> str:
return self.message
class AareDBCommunicationError(Exception):
"""Raised when AareDB HTTPS communication fails."""
def __init__(
self,
message: str = "AareDB communication error",
*,
operation: str | None = None,
endpoint: str | None = None,
base_url: str | None = None,
status_code: int | None = None,
):
super().__init__(message)
self.message = message
self.operation = operation
self.endpoint = endpoint
self.base_url = base_url
self.status_code = status_code
logger.error(
message,
extra={
"device": "aaredb",
"operation": operation,
"endpoint": endpoint,
"base_url": base_url,
"status_code": status_code,
},
)
def __str__(self) -> str:
return self.message
class AerotechCommunicationError(Exception):
"""
Raised when Aerotech HTTP/API communication fails (connection refused, timeout, bad HTTP status, etc).
Keep the original exception in `__cause__` by using `raise ... from e`.
"""
def __init__(
self,
message: str = "Aerotech communication error",
*,
endpoint: str | None = None,
base_url: str | None = None,
operation: str | None = None,
status_code: int | None = None,
):
super().__init__(message)
self.message = message
self.endpoint = endpoint
self.base_url = base_url
self.operation = operation
self.status_code = status_code
logger.error(
message,
extra={
"device": "aerotech",
"operation": operation,
"endpoint": endpoint,
"base_url": base_url,
"status_code": status_code,
},
)
def __str__(self) -> str:
return self.message
+3
View File
@@ -725,6 +725,9 @@ class DAQStatusModel(BaseModel):
smargon_connected: bool = True
smargon_error: str | None = None
aerotech_connected: bool = True
aerotech_error: str | None = None
class BeamlineSettingsModel(BaseModel):
dtz_max: float | None = 1600.0
dtz_min: float | None = 120.0
+31 -20
View File
@@ -42,17 +42,28 @@ class AareWrapper:
def __init__(
self,
bl: MXBeamline,
host: str = "https://mx-db-01.psi.ch/dispatcher",
host: str = "https://mx-aaredb-dmz-01.psi.ch/dispatcher",
):
configuration = aareDB.Configuration(host=host)
configuration.verify_ssl = False # Disable SSL verification
# --- mTLS & SSL CONFIGURATION ---
# 1. Trust the Server (CA that signed mx-aaredb-dmz-01)
configuration.verify_ssl = True
configuration.ssl_ca_cert = "/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem"
# 2. Present Machine Identity (The certs that worked in curl)
configuration.cert_file = "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_dmz-01.crt"
configuration.key_file = "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_dmz-01.key"
# 3. Initialize the Client with this config
self.client = aareDB.ApiClient(configuration)
# Identity Forwarding (Optional now that mTLS is active, but safe to keep)
self.client.default_headers["X-Shared-Password"] = os.getenv("AAREDB_SHARED_PASSWORD")
self.__host = host
self.__tell_api = aareDB.TellsRunnerApi(self.client)
self.__sample_api = aareDB.SamplesRunnerApi(self.client)
self.__proc_api = aareDB.ProcessingsRunnerApi(self.client)
self.__proc_api = aareDB.ProcessingsRunnerApi(self.client)
self.__raster_api = aareDB.GridscanRunnerApi(self.client)
self.__bl = bl
@@ -70,7 +81,7 @@ class AareWrapper:
ret = self.__tell_api.set_tell_positions(
set_tell_position_request=payload,
)
print(ret)
logger.debug(ret)
def create_manual_sample(self, s: SampleShortInfo):
from aareDB.models import ManualSampleCreate
@@ -84,7 +95,7 @@ class AareWrapper:
try:
s.db_id = self.__sample_api.insert_sample(manual_sample).id
except Exception as e:
print(f"Error inserting sample: {e}")
logger.error(f"Error inserting sample: {e}")
def sample_mounted(self, s: Optional[SampleShortInfo]):
if s is not None:
@@ -94,7 +105,7 @@ class AareWrapper:
sample_event_create=SampleEventCreate(event_type=SampleEventType("Mounted")),
)
except Exception as e:
print(e)
logger.error(e)
def sample_unmounted(self, s: Optional[SampleShortInfo]):
if s is not None:
@@ -104,7 +115,7 @@ class AareWrapper:
sample_event_create=SampleEventCreate(event_type=SampleEventType("Unmounted")),
)
except Exception as e:
print(e)
logger.error(e)
def sample_centered(self, s: Optional[SampleShortInfo]):
if s is not None:
@@ -114,7 +125,7 @@ class AareWrapper:
sample_event_create=SampleEventCreate(event_type=SampleEventType("Centered")),
)
except Exception as e:
print(e)
logger.error(e)
def sample_collected(self, s: Optional[SampleShortInfo]):
if s is None:
@@ -125,7 +136,7 @@ class AareWrapper:
sample_event_create=SampleEventCreate(event_type=SampleEventType("Collected")),
)
except Exception as e:
print(e)
logger.error(e)
def sample_failed(self, s: Optional[SampleShortInfo], failed_comment: Optional[str] = None):
if s is None:
@@ -136,7 +147,7 @@ class AareWrapper:
sample_event_create=SampleEventCreate(event_type=SampleEventType("Failed"), comment=failed_comment),
)
except Exception as e:
print(e)
logger.error(e)
def axc_failed(self, s: Optional[SampleShortInfo]):
if s is None:
@@ -147,7 +158,7 @@ class AareWrapper:
sample_event_create=SampleEventCreate(event_type=SampleEventType("AXCFailed")),
)
except Exception as e:
print(e)
logger.error(e)
def alc_failed(self, s: Optional[SampleShortInfo], alc_comment: Optional[str] = None):
if s is None:
@@ -158,7 +169,7 @@ class AareWrapper:
sample_event_create=SampleEventCreate(event_type=SampleEventType("ALCFailed"), comment=alc_comment),
)
except Exception as e:
print(e)
logger.error(e)
def sample_lost(self, s: Optional[SampleShortInfo]):
if s is None:
@@ -288,9 +299,9 @@ class AareWrapper:
sample_id=s.db_id,
experiment_parameters_create=experiment_params_payload
)
print("Experiment parameters created:", response)
logger.debug("Experiment parameters created:", response)
except Exception as e:
print(e)
logger.error(e)
def create_gridscan_run(self, s: Optional[SampleShortInfo], r:RasterGridRequest, d:DAQStatusModel):
if s is None:
@@ -352,9 +363,9 @@ class AareWrapper:
sample_id=s.db_id,
experiment_parameters_create=experiment_params_payload
)
print("Experiment parameters created:", response)
logger.info("Experiment parameters created:", response)
except Exception as e:
print(e)
logger.debug(e)
def ingest_gridscan(self, sample: Optional[SampleShortInfo], raster_result: ScanResult,
raster_request: RasterGridRequest, geom: SampleGeometryModel,
@@ -378,7 +389,7 @@ class AareWrapper:
headers=headers, data=json.dumps(payload), timeout=30, verify=False)
response.raise_for_status()
print(f"Response status code: {response.status_code}")
logger.info(f"Response status code: {response.status_code}")
def format_gridscan_payload(self, sample: Optional[SampleShortInfo], raster_result:ScanResult,
@@ -421,7 +432,7 @@ class AareWrapper:
return payload
except Exception as e:
print(e)
logger.error(e)
raise e
def ingest_scan(self, sample: Optional[SampleShortInfo], result: ScanResult,
@@ -445,7 +456,7 @@ class AareWrapper:
headers=headers, data=json.dumps(payload), timeout=30, verify=False)
response.raise_for_status()
print(f"Response status code: {response.status_code}")
logger.info(f"Response status code: {response.status_code}")
def format_scan_payload(self, sample: Optional[SampleShortInfo], result:ScanResult,
geom:SampleGeometryModel,
@@ -463,5 +474,5 @@ class AareWrapper:
return payload
except Exception as e:
print(e)
logger.error(e)
raise e
+274 -4
View File
@@ -1,14 +1,18 @@
import grp
import os
import pwd
import uuid
from datetime import datetime, timedelta, UTC
from typing import List
import time
import jwt
from fastapi import HTTPException, status
from fastapi.security import OAuth2PasswordRequestForm
from fastapi import Depends, HTTPException, status
from fastapi.security import OAuth2PasswordRequestForm, OAuth2PasswordBearer
from pydantic import BaseModel
from aare.common.auth_models import BatonRequestStatus, BatonTransferQueue, BatonRequest, BatonStatus
from aare.common.models import SessionsStateEnum
from aare.daq.config import BeamlineConfig
from aare.common.exception_handler import AuthenticationException, UserRightsException, AuthErrorCode
@@ -21,10 +25,13 @@ SECRET_KEY = os.environ.get("JWT_AAREDAQ_KEY")
ALGORITHM = "HS256"
ACCESS_TOKEN_EXPIRE_MINUTES = 24 * 60 * 7 # 1 week
SESSION_EXPIRE_SECONDS = 60 * 10
BATON_REQUEST_TIMEOUT_SECONDS = 30
STAFF_GROUP = "unx-MXgroup"
SUPER_USERS = ["e10019", "e11206", "e18147"]
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token")
class TokenData(BaseModel):
sub: str # Username
pgroups: List[str]
@@ -54,7 +61,7 @@ def authenticate_user(cfg: BeamlineConfig, form_data: OAuth2PasswordRequestForm)
return create_access_token(token)
def parse_token(token: str) -> TokenData:
def parse_token(token: str = Depends(oauth2_scheme)) -> TokenData:
try:
payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
token = TokenData(**payload)
@@ -112,4 +119,267 @@ def check_jwt_staff(cfg: BeamlineConfig, data: TokenData) -> None:
) from e
def force_current_sesion(cfg: BeamlineConfig, data: TokenData) -> None:
cfg.force_set_active_session(data.session, SESSION_EXPIRE_SECONDS)
cfg.execute_baton_transfer(
to_session=data.session,
to_username=data.sub,
to_is_staff=data.staff,
to_pgroup=cfg.pgroup,
expiry_sec=SESSION_EXPIRE_SECONDS
)
def _finalize_expired_baton_request(cfg: BeamlineConfig, pending: BatonRequest) -> BatonStatus:
"""
Resolve an expired baton request in one place.
Rules:
- staff can override immediately if beamline is free
- otherwise the transfer is queued
- pending request is cleared once it is no longer pending
"""
requester_is_staff = bool(pending.requester_is_staff)
if cfg.can_transfer_baton_now():
cfg.execute_baton_transfer(
to_session=pending.requester_session,
to_username=pending.requester_username,
to_is_staff=requester_is_staff,
to_pgroup=cfg.pgroup,
expiry_sec=SESSION_EXPIRE_SECONDS
)
else:
cfg.queued_baton_transfer = BatonTransferQueue(
target_session=pending.requester_session,
target_username=pending.requester_username,
target_is_staff=requester_is_staff,
target_pgroup=cfg.pgroup,
queued_at=time.time(),
reason="timeout_beamline_busy"
)
cfg.clear_pending_baton_request()
return get_baton_status(cfg, TokenData(
sub=pending.requester_username,
pgroups=[],
session=pending.requester_session,
staff=requester_is_staff
))
def resolve_baton_timeout_if_needed(cfg: BeamlineConfig) -> BatonStatus | None:
"""
Check the current pending request and resolve it if expired.
Returns the updated BatonStatus when a timeout was processed, else None.
"""
pending = cfg.pending_baton_request
if pending is None or pending.status != BatonRequestStatus.PENDING:
return None
elapsed = time.time() - pending.created_at
if elapsed < pending.timeout_seconds:
return None
return _finalize_expired_baton_request(cfg, pending)
def get_baton_status(cfg: BeamlineConfig, data: TokenData) -> BatonStatus:
"""
Build baton status scoped to the requesting session.
Important:
- requester sees you_have_pending_request
- holder sees incoming_request
- nobody else sees the request as actionable
"""
holder = cfg.baton_holder
pending = cfg.pending_baton_request
is_requester = bool(
pending
and pending.status in (BatonRequestStatus.PENDING, BatonRequestStatus.REFUSED)
and pending.requester_session == data.session
)
is_holder = bool(
pending
and pending.status == BatonRequestStatus.PENDING
and pending.holder_session == data.session
)
# Only expose the pending request object to the two relevant sessions
scoped_pending = pending if (is_requester or is_holder) else None
return BatonStatus(
holder=holder,
pending_request=scoped_pending,
queued_transfer=cfg.queued_baton_transfer,
you_are_holder=bool(holder and holder.session == data.session),
you_have_pending_request=is_requester,
incoming_request=is_holder,
allow_non_staff_request=cfg.allow_non_staff_request_from_staff,
)
def request_baton(cfg: BeamlineConfig, data: TokenData) -> dict:
resolve_baton_timeout_if_needed(cfg)
session_state = cfg.session_state(data.session)
if session_state == SessionsStateEnum.Vacant:
cfg.execute_baton_transfer(
to_session=data.session,
to_username=data.sub,
to_is_staff=data.staff,
to_pgroup=cfg.pgroup,
expiry_sec=SESSION_EXPIRE_SECONDS,
)
return {"granted": True, "message": "Baton acquired (beamline was vacant)"}
if session_state == SessionsStateEnum.OwnedByYou:
cfg.try_set_active_session(data.session, SESSION_EXPIRE_SECONDS)
return {"already_holder": True, "message": "You already hold the baton"}
holder = cfg.baton_holder
print(cfg.allow_non_staff_request_from_staff)
if holder and holder.is_staff and not data.staff and not cfg.allow_non_staff_request_from_staff:
return {
"error": True,
"message": "Requesting baton from staff is disabled by backend policy.",
}
if data.staff:
if not cfg.can_transfer_baton_now():
cfg.queued_baton_transfer = BatonTransferQueue(
target_session=data.session,
target_username=data.sub,
target_is_staff=data.staff,
target_pgroup=cfg.pgroup,
queued_at=time.time(),
reason="beamline_busy_staff_override",
)
return {
"queued": True,
"message": "Staff override queued - will transfer when beamline is available",
}
cfg.execute_baton_transfer(
to_session=data.session,
to_username=data.sub,
to_is_staff=data.staff,
to_pgroup=cfg.pgroup,
expiry_sec=SESSION_EXPIRE_SECONDS,
)
return {"granted": True, "override": True, "message": "Staff override - baton acquired"}
existing_request = cfg.pending_baton_request
if existing_request and existing_request.status == BatonRequestStatus.PENDING:
if existing_request.requester_session == data.session:
elapsed = time.time() - existing_request.created_at
if elapsed >= existing_request.timeout_seconds:
return {"timeout": True, "message": "Request timed out"}
remaining = existing_request.timeout_seconds - elapsed
return {
"pending": True,
"existing": True,
"remaining_seconds": max(0, remaining),
"message": f"Request already pending ({remaining:.0f}s remaining)",
}
return {
"error": True,
"message": "Another user already has a pending request",
}
request = BatonRequest(
request_id=str(uuid.uuid4()),
requester_username=data.sub,
requester_session=data.session,
requester_is_staff=data.staff,
holder_username=holder.username if holder else None,
holder_session=holder.session if holder else None,
created_at=time.time(),
timeout_seconds=BATON_REQUEST_TIMEOUT_SECONDS,
status=BatonRequestStatus.PENDING,
)
cfg.set_pending_baton_request(request, timeout_sec=BATON_REQUEST_TIMEOUT_SECONDS)
return {
"pending": True,
"request_id": request.request_id,
"timeout_seconds": BATON_REQUEST_TIMEOUT_SECONDS,
"message": f"Request sent to {holder.username if holder else 'current holder'}",
}
def respond_to_baton_request(cfg: BeamlineConfig, data: TokenData, accept: bool) -> dict:
"""
Current baton holder responds to a pending request.
"""
resolve_baton_timeout_if_needed(cfg)
holder = cfg.baton_holder
if holder is None or holder.session != data.session:
return {"error": True, "message": "You are not the current baton holder"}
pending = cfg.pending_baton_request
if pending is None or pending.status != BatonRequestStatus.PENDING:
return {"error": True, "message": "No pending request to respond to"}
if accept:
if cfg.can_transfer_baton_now():
cfg.execute_baton_transfer(
to_session=pending.requester_session,
to_username=pending.requester_username,
to_is_staff=pending.requester_is_staff,
to_pgroup=cfg.pgroup,
expiry_sec=SESSION_EXPIRE_SECONDS
)
return {"accepted": True, "transferred": True, "message": "Baton transferred"}
else:
cfg.queued_baton_transfer = BatonTransferQueue(
target_session=pending.requester_session,
target_username=pending.requester_username,
target_is_staff=pending.requester_is_staff,
target_pgroup=cfg.pgroup,
queued_at=time.time(),
reason="accepted_beamline_busy"
)
cfg.clear_pending_baton_request()
return {
"accepted": True,
"queued": True,
"message": "Request accepted - will transfer when beamline is available"
}
else:
pending.status = BatonRequestStatus.REFUSED
cfg.set_pending_baton_request(pending, timeout_sec=5)
return {"refused": True, "message": "Request refused"}
def release_baton(cfg: BeamlineConfig, data: TokenData) -> dict:
"""
Voluntarily release the baton (set session to free).
"""
resolve_baton_timeout_if_needed(cfg)
holder = cfg.baton_holder
if holder is None:
return {"released": True, "message": "Baton was already vacant"}
if holder.session != data.session:
return {"info": True, "message": "You don't hold the baton"}
cfg.end_active_session(data.session)
cfg.baton_holder = None
cfg.clear_pending_baton_request()
return {"released": True, "message": "Baton released - beamline is now vacant"}
def cancel_baton_request(cfg: BeamlineConfig, data: TokenData) -> dict:
"""
Cancel your own pending baton request.
"""
resolve_baton_timeout_if_needed(cfg)
pending = cfg.pending_baton_request
if pending is None:
return {"error": True, "message": "No pending request to cancel"}
if pending.requester_session != data.session:
return {"error": True, "message": "You can only cancel your own request"}
cfg.clear_pending_baton_request()
return {"cancelled": True, "message": "Request cancelled"}
+761
View File
@@ -0,0 +1,761 @@
from __future__ import annotations
import asyncio
import json
from typing import AsyncGenerator, TYPE_CHECKING
from fastapi import APIRouter, Depends, HTTPException, status
from pydantic import BaseModel, Field
from starlette.responses import StreamingResponse
from aare.common.automation_models import (
QueueItem,
QueueItemStatus,
WorkflowContext,
WorkflowEvent,
WorkflowMode,
WorkflowStepRecord,
RuntimeState,
ControlState,
WorkflowStateKind, EventListResponse, StepActionResponse, ControlActionResponse, RuntimeResponse,
CreateQueueItemRequest, QueueListResponse, MoveItemRequest, AutomationStatusResponse,
)
from aare.common.automation_queue_manager import (
WorkflowRedisManager,
build_default_steps,
)
from aare.common.automation_workflow import STATE_REGISTRY, HANDLER_REGISTRY
from aare.daq.automation_runner import PersistentWorkflowRunner
from aare.daq.auth import parse_token, check_jwt_rw, check_jwt_ro, oauth2_scheme
from aare.common.models import TokenData
from aare.daq.config import BeamlineConfig
from aare.daq.automation_runner import AutomationLoop
router = APIRouter(prefix="/workflow", tags=["workflow"])
# ─────────────────────────────────────────────
# Dependency: get managers
# ─────────────────────────────────────────────
# These will be set up when the router is included
_redis_manager: WorkflowRedisManager | None = None
_runner: PersistentWorkflowRunner | None = None
_cfg: BeamlineConfig | None = None
_automation_loop: AutomationLoop | None = None
def set_workflow_dependencies(
redis_manager: WorkflowRedisManager,
runner: PersistentWorkflowRunner,
cfg: BeamlineConfig,
) -> None:
global _redis_manager, _runner, _cfg, _automation_loop
_redis_manager = redis_manager
_runner = runner
_cfg = cfg
_automation_loop = AutomationLoop(runner, redis_manager)
def get_automation_loop() -> AutomationLoop:
if _automation_loop is None:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail="Automation loop not initialized",
)
return _automation_loop
def get_cfg() -> BeamlineConfig:
if _cfg is None:
raise HTTPException(
status_code=503,
detail="Workflow system not initialized",
)
return _cfg
def get_redis_manager() -> WorkflowRedisManager:
if _redis_manager is None:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail="Workflow system not initialized",
)
return _redis_manager
def get_runner() -> PersistentWorkflowRunner:
if _runner is None:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail="Workflow runner not initialized",
)
return _runner
# ─────────────────────────────────────────────
# Queue management endpoints
# ─────────────────────────────────────────────
@router.get("/queue", response_model=QueueListResponse)
async def list_queue(
include_finished: bool = False,
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
):
"""List all items in the workflow queue."""
data = parse_token(token)
check_jwt_ro(get_cfg(), data)
items = redis_mgr.list_items(include_finished=include_finished)
# Filter by pgroup if not staff
if not data.staff:
items = [i for i in items if i.owner_pgroup in data.pgroups]
return QueueListResponse(items=items, total=len(items))
@router.post("/queue", response_model=QueueItem)
async def create_queue_item(
request: CreateQueueItemRequest,
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
):
"""Add a new item to the workflow queue."""
data = parse_token(token)
check_jwt_rw(get_cfg(), data)
# Build steps
if request.steps:
steps = [
WorkflowStepRecord(kind=s)
for s in request.steps
if s in [sk.value for sk in WorkflowStateKind]
]
else:
steps = build_default_steps()
item = QueueItem(
item_id="",
beamline=redis_mgr._bl,
sample_id=request.sample_id,
sample_name=request.sample_name,
owner_pgroup=data.pgroups[0] if data.pgroups else "",
created_by=data.sub,
priority=request.priority,
order_index=int(asyncio.get_event_loop().time() * 1000),
steps=steps,
recipe=request.recipe,
)
try:
item = redis_mgr.create_item(item)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
return item
@router.get("/queue/{item_id}", response_model=QueueItem)
async def get_queue_item(
item_id: str,
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
):
"""Get a single queue item by ID."""
data = parse_token(token)
check_jwt_ro(get_cfg(), data)
item = redis_mgr.get_item(item_id)
if item is None:
raise HTTPException(status_code=404, detail="Item not found")
# Check access
if not data.staff and item.owner_pgroup not in data.pgroups:
raise HTTPException(status_code=403, detail="Access denied")
return item
@router.delete("/queue/{item_id}")
async def delete_queue_item(
item_id: str,
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
):
"""Remove an item from the queue."""
data = parse_token(token)
check_jwt_rw(get_cfg(), data)
item = redis_mgr.get_item(item_id)
if item is None:
raise HTTPException(status_code=404, detail="Item not found")
# Check access
if not data.staff and item.owner_pgroup not in data.pgroups:
raise HTTPException(status_code=403, detail="Access denied")
# Don't allow deleting running items
if item.status == QueueItemStatus.RUNNING:
raise HTTPException(status_code=409, detail="Cannot delete running item")
redis_mgr.delete_item(item_id)
return {"ok": True, "message": f"Item {item_id} deleted"}
@router.delete("/queue/clear")
async def clear_queue(
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
):
"""Clear all non-running items from the queue."""
data = parse_token(token)
check_jwt_rw(get_cfg(), data)
items = redis_mgr.list_items(include_finished=True)
deleted_count = 0
for item in items:
# Skip running items
if item.status == QueueItemStatus.RUNNING:
continue
# Check access
if not data.staff and item.owner_pgroup not in data.pgroups:
continue
redis_mgr.delete_item(item.item_id)
deleted_count += 1
return {"ok": True, "message": f"Cleared {deleted_count} items from queue"}
@router.post("/control/skip_sample", response_model=ControlActionResponse)
async def request_skip_sample(
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
runner: PersistentWorkflowRunner = Depends(get_runner),
):
"""Skip the current sample entirely (abort current sample and move to next)."""
data = parse_token(token)
check_jwt_rw(get_cfg(), data)
runtime = redis_mgr.get_runtime()
if not runtime.running or not runtime.current_item_id:
return ControlActionResponse(
ok=False,
control=redis_mgr.get_control(),
message="No sample currently running",
)
# Mark current item as skipped and complete it
runner.complete_item(runtime.current_item_id, QueueItemStatus.SKIPPED)
redis_mgr.append_event(WorkflowEvent(
beamline=runner.beamline,
item_id=runtime.current_item_id,
event_type="sample_skipped",
actor=data.sub,
message="Sample skipped by user",
))
return ControlActionResponse(
ok=True,
control=redis_mgr.get_control(),
message="Sample skipped",
)
@router.post("/queue/{item_id}/move")
async def move_queue_item(
item_id: str,
request: MoveItemRequest,
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
):
"""Reorder an item in the queue."""
data = parse_token(token)
check_jwt_rw(get_cfg(), data)
item = redis_mgr.get_item(item_id)
if item is None:
raise HTTPException(status_code=404, detail="Item not found")
if not data.staff and item.owner_pgroup not in data.pgroups:
raise HTTPException(status_code=403, detail="Access denied")
redis_mgr.move_item(item_id, request.new_order_index)
return {"ok": True, "message": f"Item {item_id} moved"}
# ─────────────────────────────────────────────
# Runtime and control endpoints
# ─────────────────────────────────────────────
@router.get("/runtime", response_model=RuntimeResponse)
async def get_runtime(
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
):
"""Get current runtime and control state."""
data = parse_token(token)
check_jwt_ro(get_cfg(), data)
return RuntimeResponse(
runtime=redis_mgr.get_runtime(),
control=redis_mgr.get_control(),
)
@router.get("/control", response_model=ControlState)
async def get_control(
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
):
"""Get current control state."""
data = parse_token(token)
check_jwt_ro(get_cfg(), data)
return redis_mgr.get_control()
@router.post("/control/pause", response_model=ControlActionResponse)
async def request_pause(
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
):
"""Request workflow pause after current step."""
data = parse_token(token)
check_jwt_rw(get_cfg(), data)
control = redis_mgr.request_control(
{"pause_requested": True, "resume_requested": False},
requested_by=data.sub,
)
return ControlActionResponse(
ok=True,
control=control,
message="Pause requested",
)
@router.post("/control/resume", response_model=ControlActionResponse)
async def request_resume(
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
):
"""Resume paused workflow."""
data = parse_token(token)
check_jwt_rw(get_cfg(), data)
control = redis_mgr.request_control(
{"pause_requested": False, "resume_requested": True},
requested_by=data.sub,
)
return ControlActionResponse(
ok=True,
control=control,
message="Resume requested",
)
@router.post("/control/abort", response_model=ControlActionResponse)
async def request_abort(
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
):
"""Abort current workflow execution."""
data = parse_token(token)
check_jwt_rw(get_cfg(), data)
# Clear pause_requested when aborting
control = redis_mgr.request_control(
{"abort_requested": True, "pause_requested": False, "resume_requested": False},
requested_by=data.sub,
)
return ControlActionResponse(
ok=True,
control=control,
message="Abort requested",
)
@router.post("/control/skip", response_model=ControlActionResponse)
async def request_skip(
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
):
"""Skip current step."""
data = parse_token(token)
check_jwt_rw(get_cfg(), data)
control = redis_mgr.request_control(
{"skip_requested": True},
requested_by=data.sub,
)
return ControlActionResponse(
ok=True,
control=control,
message="Skip requested",
)
@router.post("/control/clear", response_model=ControlActionResponse)
async def clear_control(
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
):
"""Clear all control requests."""
data = parse_token(token)
check_jwt_rw(get_cfg(), data)
control = redis_mgr.clear_control()
return ControlActionResponse(
ok=True,
control=control,
message="Control state cleared",
)
# ─────────────────────────────────────────────
# Execution endpoints
# ─────────────────────────────────────────────
@router.post("/start/{item_id}", response_model=StepActionResponse)
async def start_item(
item_id: str,
mode: WorkflowMode = WorkflowMode.GUIDED_MANUAL,
token: str = Depends(oauth2_scheme),
runner: PersistentWorkflowRunner = Depends(get_runner),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
):
"""Start processing a queue item."""
data = parse_token(token)
check_jwt_rw(get_cfg(), data)
item = redis_mgr.get_item(item_id)
if item is None:
raise HTTPException(status_code=404, detail="Item not found")
if not data.staff and item.owner_pgroup not in data.pgroups:
raise HTTPException(status_code=403, detail="Access denied")
if item.status == QueueItemStatus.RUNNING:
raise HTTPException(status_code=409, detail="Item already running")
# Check if another item is running
runtime = redis_mgr.get_runtime()
if runtime.running:
raise HTTPException(
status_code=409,
detail=f"Another item is running: {runtime.current_item_id}",
)
# Clear any stale control requests
redis_mgr.clear_control()
# Start the item
item = runner.start_item(item_id)
return StepActionResponse(
ok=True,
item=item,
step=item.steps[0].kind if item.steps else None,
status="started",
message=f"Started processing {item.sample_name}",
)
@router.post("/next", response_model=StepActionResponse)
async def run_next_step(
token: str = Depends(oauth2_scheme),
runner: PersistentWorkflowRunner = Depends(get_runner),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
):
"""
Run the next step in guided mode.
This is the main endpoint for guided manual operation.
User clicks "Next" and this executes one step.
"""
data = parse_token(token)
check_jwt_rw(get_cfg(), data)
runtime = redis_mgr.get_runtime()
if not runtime.running or not runtime.current_item_id:
raise HTTPException(status_code=409, detail="No item is currently running")
item = redis_mgr.get_item(runtime.current_item_id)
if item is None:
raise HTTPException(status_code=404, detail="Running item not found")
if not data.staff and item.owner_pgroup not in data.pgroups:
raise HTTPException(status_code=403, detail="Access denied")
# Check if all steps completed
if item.current_step_index >= len(item.steps):
item = runner.complete_item(item.item_id, QueueItemStatus.COMPLETED)
return StepActionResponse(
ok=True,
item=item,
status="completed",
message="All steps completed",
)
# Check for abort
if runner.should_abort():
item = runner.complete_item(item.item_id, QueueItemStatus.ABORTED)
redis_mgr.clear_control()
return StepActionResponse(
ok=True,
item=item,
status="aborted",
message="Workflow aborted by user",
)
# Check for skip
control = runner.check_control()
if control.skip_requested:
step_index = item.current_step_index
step_kind = item.steps[step_index].kind
redis_mgr.update_step(item.item_id, step_index, status="skipped")
redis_mgr.update_item(item.item_id, {"current_step_index": step_index + 1})
redis_mgr.request_control({"skip_requested": False}, requested_by=data.sub)
redis_mgr.append_event(WorkflowEvent(
beamline=runner.beamline,
item_id=item.item_id,
step=step_kind,
event_type="step_skipped",
actor=data.sub,
message=f"Step {step_kind} skipped by user",
))
item = redis_mgr.get_item(item.item_id)
return StepActionResponse(
ok=True,
item=item,
step=step_kind,
status="skipped",
message=f"Skipped {step_kind}",
)
# Build context
context = WorkflowContext(
mode=WorkflowMode.GUIDED_MANUAL,
queue_id="default",
item_id=item.item_id,
sample_id=item.sample_id,
current_state=WorkflowStateKind(runtime.current_state) if runtime.current_state else None,
current_step_index=item.current_step_index,
)
# Run the step
try:
result = runner.run_current_step(context)
except Exception as e:
return StepActionResponse(
ok=False,
item=redis_mgr.get_item(item.item_id),
step=item.steps[item.current_step_index].kind,
status="failed",
message=str(e),
)
# Get updated item
item = redis_mgr.get_item(item.item_id)
# Check if completed
if item.current_step_index >= len(item.steps):
item = runner.complete_item(item.item_id, QueueItemStatus.COMPLETED)
return StepActionResponse(
ok=True,
item=item,
step=result.state.value,
status="completed",
message="All steps completed",
)
return StepActionResponse(
ok=True,
item=item,
step=result.state.value,
status=result.status.value,
message=result.message,
)
# ─────────────────────────────────────────────
# Events endpoints
# ─────────────────────────────────────────────
@router.get("/events", response_model=EventListResponse)
async def get_events(
limit: int = 100,
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
):
"""Get recent workflow events."""
data = parse_token(token)
check_jwt_ro(get_cfg(), data)
events = redis_mgr.read_events(limit=limit)
return EventListResponse(events=events)
# ─────────────────────────────────────────────
# SSE stream
# ─────────────────────────────────────────────
async def workflow_event_stream(
redis_mgr: WorkflowRedisManager,
) -> AsyncGenerator[str, None]:
"""
Server-Sent Events stream for live workflow updates.
Polls runtime state and streams changes.
"""
last_runtime_json = ""
last_control_json = ""
last_event_id = "0-0"
try:
while True:
# Check runtime state
runtime = redis_mgr.get_runtime()
runtime_json = runtime.model_dump_json()
if runtime_json != last_runtime_json:
last_runtime_json = runtime_json
yield f"event: runtime\ndata: {runtime_json}\n\n"
# Check control state
control = redis_mgr.get_control()
control_json = control.model_dump_json()
if control_json != last_control_json:
last_control_json = control_json
yield f"event: control\ndata: {control_json}\n\n"
# Check for new events (using Redis streams)
try:
events_key = redis_mgr._events_key()
new_events = redis_mgr._client.xread(
{events_key: last_event_id},
count=10,
block=0,
)
if new_events:
for _, messages in new_events:
for msg_id, fields in messages:
last_event_id = msg_id
raw = fields.get("json")
if raw:
yield f"event: workflow_event\ndata: {raw}\n\n"
except Exception:
pass # Redis stream read failed, continue polling
await asyncio.sleep(0.2)
except asyncio.CancelledError:
return
# ─────────────────────────────────────────────
# Automation mode endpoints
# ─────────────────────────────────────────────
@router.get("/automation/status", response_model=AutomationStatusResponse)
async def get_automation_status(
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
loop: "AutomationLoop" = Depends(get_automation_loop),
):
"""Get automation loop status."""
data = parse_token(token)
check_jwt_ro(get_cfg(), data)
return AutomationStatusResponse(
enabled=loop.is_enabled,
running=loop.is_running,
runtime=redis_mgr.get_runtime(),
control=redis_mgr.get_control(),
)
@router.post("/automation/start", response_model=AutomationStatusResponse)
async def start_automation(
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
loop: "AutomationLoop" = Depends(get_automation_loop),
):
"""Start automation mode - processes queue automatically."""
data = parse_token(token)
check_jwt_rw(get_cfg(), data)
loop.start()
return AutomationStatusResponse(
enabled=loop.is_enabled,
running=loop.is_running,
runtime=redis_mgr.get_runtime(),
control=redis_mgr.get_control(),
)
@router.post("/automation/stop", response_model=AutomationStatusResponse)
async def stop_automation(
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
loop: "AutomationLoop" = Depends(get_automation_loop),
):
"""Stop automation mode - completes current step then stops."""
data = parse_token(token)
check_jwt_rw(get_cfg(), data)
loop.stop()
return AutomationStatusResponse(
enabled=loop.is_enabled,
running=loop.is_running,
runtime=redis_mgr.get_runtime(),
control=redis_mgr.get_control(),
)
@router.get("/sse")
async def workflow_sse(
token: str = Depends(oauth2_scheme),
redis_mgr: WorkflowRedisManager = Depends(get_redis_manager),
):
"""
SSE endpoint for live workflow updates.
Events:
- runtime: RuntimeState changes
- control: ControlState changes
- workflow_event: Individual workflow events
"""
data = parse_token(token)
check_jwt_ro(get_cfg(), data)
return StreamingResponse(
workflow_event_stream(redis_mgr),
media_type="text/event-stream",
headers={
"Cache-Control": "no-cache",
"Connection": "keep-alive",
"Access-Control-Allow-Origin": "*",
"Access-Control-Allow-Headers": "Cache-Control",
},
)
+461
View File
@@ -0,0 +1,461 @@
import asyncio
from typing import Callable
import redis
from aare.common.automation_models import (
WorkflowEvent,
QueueItemStatus,
QueueItem,
ControlState,
WorkflowStateKind,
WorkflowContext,
StateResult,
StateDefinition,
WorkflowMode,
)
from aare.common.automation_queue_manager import (
WorkflowRedisManager,
build_default_steps
)
from aare.common.automation_workflow import (
can_transition,
StateHandler,
STATE_REGISTRY,
HANDLER_REGISTRY
)
class PersistentWorkflowRunner:
"""
Orchestrates workflow execution with Redis persistence.
Combines:
- State handlers (from automation_workflow)
- Redis persistence (from automation_queue_manager)
"""
def __init__(
self,
redis_manager: WorkflowRedisManager,
registry: dict[WorkflowStateKind, StateDefinition] | None = None,
handlers: dict[WorkflowStateKind, StateHandler] | None = None,
):
self._redis = redis_manager
self._registry = registry or STATE_REGISTRY
self._handlers = handlers or HANDLER_REGISTRY
@property
def beamline(self) -> str:
return self._redis._bl
def get_handler(self, state: WorkflowStateKind) -> StateHandler:
try:
return self._handlers[state]
except KeyError as exc:
raise KeyError(f"No handler registered for state: {state}") from exc
def start_item(self, item_id: str) -> QueueItem:
item = self._redis.get_item(item_id)
if item is None:
raise KeyError(f"Item not found: {item_id}")
item = self._redis.update_item(item_id, {
"status": QueueItemStatus.RUNNING,
"current_step_index": 0,
})
self._redis.patch_runtime({
"running": True,
"paused": False,
"current_item_id": item_id,
"current_step_index": 0,
"current_state": None,
"last_error": None,
})
self._redis.append_event(WorkflowEvent(
beamline=self.beamline,
item_id=item_id,
event_type="item_started",
message=f"Started processing {item.sample_name}",
))
return item
def run_current_step(self, context: WorkflowContext) -> StateResult:
item = self._redis.get_item(context.item_id)
if item is None:
raise KeyError(f"Item not found: {context.item_id}")
step_index = item.current_step_index
if step_index >= len(item.steps):
raise RuntimeError("No more steps to run")
step_record = item.steps[step_index]
state_kind = WorkflowStateKind(step_record.kind)
# Check transition is allowed
if context.current_state is not None:
if not can_transition(context.current_state, state_kind, context.mode):
raise RuntimeError(
f"Transition not allowed: {context.current_state} -> {state_kind}"
)
# Mark step as running
self._redis.update_step(context.item_id, step_index, status="running")
self._redis.patch_runtime({
"current_state": state_kind.value,
"current_step_index": step_index,
})
self._redis.append_event(WorkflowEvent(
beamline=self.beamline,
item_id=context.item_id,
step=state_kind.value,
event_type="step_started",
message=f"Starting {state_kind.value}",
))
# Execute handler
handler = self.get_handler(state_kind)
try:
result = handler.execute(context)
except Exception as e:
self._redis.update_step(
context.item_id,
step_index,
status="failed",
error_detail=str(e),
)
self._redis.patch_runtime({"last_error": str(e)})
self._redis.append_event(WorkflowEvent(
beamline=self.beamline,
item_id=context.item_id,
step=state_kind.value,
event_type="step_failed",
message=str(e),
))
raise
# Mark step as completed
self._redis.update_step(
context.item_id,
step_index,
status=result.status.value,
message=result.message,
)
# Advance step index
self._redis.update_item(context.item_id, {
"current_step_index": step_index + 1,
})
context.current_state = state_kind
context.current_step_index = step_index + 1
self._redis.append_event(WorkflowEvent(
beamline=self.beamline,
item_id=context.item_id,
step=state_kind.value,
event_type="step_completed",
message=result.message,
payload=result.payload,
))
return result
def complete_item(self, item_id: str, status: QueueItemStatus) -> QueueItem:
item = self._redis.update_item(item_id, {"status": status})
self._redis.patch_runtime({
"running": False,
"current_item_id": None,
"current_state": None,
})
self._redis.append_event(WorkflowEvent(
beamline=self.beamline,
item_id=item_id,
event_type="item_completed",
message=f"Item finished with status {status.value}",
))
if status in (QueueItemStatus.COMPLETED, QueueItemStatus.ABORTED, QueueItemStatus.SKIPPED):
self._redis.delete_item(item_id)
return item
return item
def check_control(self) -> ControlState:
return self._redis.get_control()
def should_pause(self) -> bool:
return self.check_control().pause_requested
def should_abort(self) -> bool:
return self.check_control().abort_requested
def should_skip(self) -> bool:
return self.check_control().skip_requested
class AutomationLoop:
"""
Background task that drives fully automated workflow execution.
Polls control state and processes the queue automatically.
IMPORTANT: This loop does NOT auto-start. It must be explicitly started
via the start() method (triggered by the "Start Automation" button).
"""
def __init__(
self,
runner: PersistentWorkflowRunner,
redis_manager: WorkflowRedisManager,
poll_interval: float = 0.3,
step_delay: float = 0.5, # Delay between steps for control checks
):
self._runner = runner
self._redis = redis_manager
self._poll_interval = poll_interval
self._step_delay = step_delay
self._task: asyncio.Task | None = None
self._enabled = False
self._on_step_complete: Callable[[str, str, StateResult], None] | None = None
@property
def is_running(self) -> bool:
return self._task is not None and not self._task.done()
@property
def is_enabled(self) -> bool:
return self._enabled
def set_step_callback(self, callback: Callable[[str, str, StateResult], None]) -> None:
"""Set callback for step completion: callback(item_id, step_name, result)"""
self._on_step_complete = callback
def start(self) -> None:
"""Start the automation loop. Must be explicitly called."""
if self._task is not None and not self._task.done():
return # Already running
self._enabled = True
self._task = asyncio.create_task(self._run_loop())
self._redis.append_event(WorkflowEvent(
beamline=self._runner.beamline,
event_type="automation_started",
message="Automation mode enabled",
))
def stop(self) -> None:
"""Stop the automation loop gracefully."""
self._enabled = False
# Cancel the task if it exists
if self._task is not None and not self._task.done():
self._task.cancel()
self._redis.append_event(WorkflowEvent(
beamline=self._runner.beamline,
event_type="automation_stopped",
message="Automation mode disabled",
))
async def _run_loop(self) -> None:
"""Main automation loop."""
while self._enabled:
try:
# Check control state FIRST before doing anything
control = self._runner.check_control()
runtime = self._redis.get_runtime()
# Handle abort immediately
if control.abort_requested:
if runtime.current_item_id:
self._runner.complete_item(runtime.current_item_id, QueueItemStatus.ABORTED)
# Clear all control flags including pause
self._redis.clear_control()
self.stop()
await asyncio.sleep(self._poll_interval)
continue
# Handle pause
if control.pause_requested:
if runtime.running and not runtime.paused:
self._redis.patch_runtime({"paused": True})
self._redis.append_event(WorkflowEvent(
beamline=self._runner.beamline,
item_id=runtime.current_item_id,
event_type="workflow_paused",
message="Workflow paused by user request",
))
await asyncio.sleep(self._poll_interval)
continue
# Handle resume
if control.resume_requested and runtime.paused:
self._redis.patch_runtime({"paused": False})
self._redis.request_control(
{"resume_requested": False},
requested_by="automation_loop",
)
self._redis.append_event(WorkflowEvent(
beamline=self._runner.beamline,
item_id=runtime.current_item_id,
event_type="workflow_resumed",
message="Workflow resumed",
))
# Don't process if paused
if runtime.paused:
await asyncio.sleep(self._poll_interval)
continue
await self._tick()
except asyncio.CancelledError:
# Loop was cancelled (stop() was called)
break
except Exception as e:
self._redis.patch_runtime({"last_error": str(e)})
self._redis.append_event(WorkflowEvent(
beamline=self._runner.beamline,
event_type="automation_error",
message=f"Automation error: {e}",
))
await asyncio.sleep(2.0)
await asyncio.sleep(self._poll_interval)
async def _tick(self) -> None:
"""Single iteration of the automation loop - process one step."""
runtime = self._redis.get_runtime()
control = self._runner.check_control()
# If nothing running, try to start next item
if not runtime.running or not runtime.current_item_id:
next_item = self._redis.get_next_pending_item()
if next_item is None:
return # Queue empty, nothing to do
self._runner.start_item(next_item.item_id)
# Add delay after starting to allow UI to catch up
await asyncio.sleep(self._step_delay)
runtime = self._redis.get_runtime()
# Re-check control after potential start
control = self._runner.check_control()
if control.abort_requested or control.pause_requested:
return # Let main loop handle it
# Process current item
item = self._redis.get_item(runtime.current_item_id)
if item is None:
self._redis.patch_runtime({"running": False, "current_item_id": None})
return
# Check if item is complete
if item.current_step_index >= len(item.steps):
self._runner.complete_item(item.item_id, QueueItemStatus.COMPLETED)
# Clear next_sample_requested if set
if control.next_sample_requested:
self._redis.request_control(
{"next_sample_requested": False},
requested_by="automation_loop",
)
return
# Handle skip request (skip current step)
if control.skip_requested:
step_kind = item.steps[item.current_step_index].kind
self._redis.update_step(item.item_id, item.current_step_index, status="skipped")
self._redis.update_item(item.item_id, {"current_step_index": item.current_step_index + 1})
self._redis.request_control({"skip_requested": False}, requested_by="automation_loop")
self._redis.append_event(WorkflowEvent(
beamline=self._runner.beamline,
item_id=item.item_id,
step=step_kind,
event_type="step_skipped",
message=f"Step {step_kind} skipped",
))
return
# Build context and run step
context = WorkflowContext(
mode=WorkflowMode.AUTOMATION,
queue_id="default",
item_id=item.item_id,
sample_id=item.sample_id,
current_state=WorkflowStateKind(runtime.current_state) if runtime.current_state else None,
current_step_index=item.current_step_index,
)
# Run the step (this is blocking in the async context)
step_name = item.steps[item.current_step_index].kind
result = await asyncio.to_thread(self._runner.run_current_step, context)
# Add delay after step to allow control checks
await asyncio.sleep(self._step_delay)
# Notify callback if set
if self._on_step_complete:
self._on_step_complete(item.item_id, step_name, result)
if __name__ == "__main__":
client = redis.Redis(host="localhost", port=6379, db=0, decode_responses=True)
redis_mgr = WorkflowRedisManager(client, beamline="x10sa")
# Create a queue item
item = QueueItem(
item_id="",
beamline="x10sa",
sample_id=123,
sample_name="lysozyme_01",
owner_pgroup="p12345",
created_by="user@example.com",
steps=build_default_steps(),
)
item = redis_mgr.create_item(item)
# Create runner
runner = PersistentWorkflowRunner(
redis_manager=redis_mgr,
registry=STATE_REGISTRY,
handlers=HANDLER_REGISTRY,
)
# Start processing
runner.start_item(item.item_id)
# Build context
context = WorkflowContext(
mode=WorkflowMode.GUIDED_MANUAL,
queue_id="default",
item_id=item.item_id,
sample_id=item.sample_id,
)
# Run each step (in guided mode, user triggers each one)
while context.current_step_index < len(item.steps):
if runner.should_pause():
print("Paused by user")
break
if runner.should_abort():
runner.complete_item(item.item_id, QueueItemStatus.ABORTED)
break
result = runner.run_current_step(context)
print(f"Step {result.state.value}: {result.status.value}")
# Mark complete if all steps done
if context.current_step_index >= len(item.steps):
runner.complete_item(item.item_id, QueueItemStatus.COMPLETED)
+158 -6
View File
@@ -6,7 +6,7 @@ from typing import Tuple, List
import numpy as np
import redis
import redis_lock
from aare.common.coordinate import Coordinate
from aare.common.coordinate import Coordinate, AerotechCoordinate
from aare.common.models import (
BeamlineSettingsModel,
BeamMarkCoeffModel,
@@ -18,13 +18,25 @@ from aare.common.models import (
FluorescenceSpectrumOutputModel, CrystalSize, SimpleStrategyInputModel, SimpleScanParameters
)
from aare.common.auth_models import (
BatonStatus,
BatonRequest,
BatonHolderInfo,
BatonRequestStatus,
BatonTransferQueue,
)
from aare.common.beamline import MXBeamline
from aare.common.logger_config import setup_logger
from aare.common.exception_handler import BeamlineBusyException
ABR_POS_ALIGN_DEF = Coordinate(x=-18, y=-0.266, z=0)
ABR_POS_MOUNT = Coordinate(x=0, y=0, z=0)#Coordinate(x=-18, y=0, z=0)
#TODO WHAT SHOULD THIS BE?
ABR_POS_ALIGN_DEF = AerotechCoordinate(at_mm=Coordinate(x=-18, y=-0.266, z=0))
ABR_POS_MOUNT = AerotechCoordinate(
at_mm=Coordinate(x=0, y=0, z=0),
omega_deg=0
)
#Coordinate(x=-18, y=0, z=0)
ABR_OMEGA_MOUNT = 0.0
logger = setup_logger("aareDAQ")
@@ -74,9 +86,24 @@ class BeamlineConfig:
else:
host = f"{self.__bl}-redis.psi.ch"
self.__client = redis.Redis(host=host, port=6379, db=0, decode_responses=True)
self.simulated_detector = bl is MXBeamline.SIMULATED
# Session and authentication management
@property
def allow_non_staff_request_from_staff(self) -> bool:
raw = self.__client.get(f"{self.__bl}:allow_non_staff_request_from_staff")
if raw is None:
return False
return str(raw).strip().lower() in {"1", "true", "yes", "on"}
@allow_non_staff_request_from_staff.setter
def allow_non_staff_request_from_staff(self, enabled: bool) -> None:
if enabled:
self.__client.set(f"{self.__bl}:allow_non_staff_request_from_staff", "1")
else:
self.__client.delete(f"{self.__bl}:allow_non_staff_request_from_staff")
def generate_session(self) -> int:
return int(self.__client.incr(f"{self.__bl}:session"))
@@ -152,6 +179,7 @@ class BeamlineConfig:
return
if active == session:
self.__client.delete(f"{self.__bl}:active_session")
self.__client.delete(f"{self.__bl}:baton_holder")
def force_set_active_session(self, session: int, expiry_sec: int) -> None:
# Ensure that there is no active try-set for active session
@@ -161,6 +189,125 @@ class BeamlineConfig:
self.__client.set(f"{self.__bl}:active_session", session)
self.__client.expire(f"{self.__bl}:active_session", expiry_sec)
# ========== BATON SYSTEM ==========
@property
def baton_holder(self) -> BatonHolderInfo | None:
"""Get information about the current baton holder."""
tmp = self.__client.get(f"{self.__bl}:baton_holder")
if tmp is None:
return None
try:
return BatonHolderInfo(**json.loads(tmp))
except Exception:
return None
@baton_holder.setter
def baton_holder(self, info: BatonHolderInfo | None) -> None:
if info is None:
self.__client.delete(f"{self.__bl}:baton_holder")
else:
self.__client.set(f"{self.__bl}:baton_holder", info.model_dump_json())
@property
def pending_baton_request(self) -> BatonRequest | None:
"""Get the current pending baton request, if any."""
tmp = self.__client.get(f"{self.__bl}:baton_request")
if tmp is None:
return None
try:
return BatonRequest(**json.loads(tmp))
except Exception:
return None
def set_pending_baton_request(self, request: BatonRequest | None, timeout_sec: int = 30) -> None:
"""Set a pending baton request with auto-expiry for timeout."""
if request is None:
self.__client.delete(f"{self.__bl}:baton_request")
else:
self.__client.set(f"{self.__bl}:baton_request", request.model_dump_json())
# Add a few seconds buffer so we can detect timeout vs expiry
self.__client.expire(f"{self.__bl}:baton_request", timeout_sec + 5)
def clear_pending_baton_request(self) -> None:
self.__client.delete(f"{self.__bl}:baton_request")
@property
def queued_baton_transfer(self) -> BatonTransferQueue | None:
"""Get queued transfer waiting for beamline to be available."""
tmp = self.__client.get(f"{self.__bl}:baton_transfer_queue")
if tmp is None:
return None
try:
return BatonTransferQueue(**json.loads(tmp))
except Exception:
return None
@queued_baton_transfer.setter
def queued_baton_transfer(self, transfer: BatonTransferQueue | None) -> None:
if transfer is None:
self.__client.delete(f"{self.__bl}:baton_transfer_queue")
else:
self.__client.set(f"{self.__bl}:baton_transfer_queue", transfer.model_dump_json())
def can_transfer_baton_now(self) -> bool:
"""Check if baton can be transferred (beamline not mid-operation)."""
# Can't transfer while beamline is busy
if self.state_busy:
return False
# Add automation queue check here when you implement it
# if self.automation_queue_running:
# return False
return True
def execute_baton_transfer(
self,
to_session: int,
to_username: str,
to_is_staff: bool,
to_pgroup: str | None,
expiry_sec: int
) -> None:
"""
Atomically transfer the baton to a new holder.
Use existing active_session_lock for consistency.
"""
with redis_lock.Lock(
self.__client, f"{self.__bl}:active_session_lock", expire=10
):
self.__client.set(f"{self.__bl}:active_session", to_session)
self.__client.expire(f"{self.__bl}:active_session", expiry_sec)
self.baton_holder = BatonHolderInfo(
username=to_username,
session=to_session,
is_staff=to_is_staff,
pgroup=to_pgroup
)
# Clear any pending request or queued transfer
self.clear_pending_baton_request()
self.queued_baton_transfer = None
def process_queued_transfer_if_ready(self, expiry_sec: int) -> bool:
"""
Check if there's a queued transfer and beamline is now available.
Returns True if transfer was executed.
"""
queued = self.queued_baton_transfer
if queued is None:
return False
if not self.can_transfer_baton_now():
return False
self.execute_baton_transfer(
to_session=queued.target_session,
to_username=queued.target_username,
to_is_staff=queued.target_is_staff,
to_pgroup=queued.target_pgroup,
expiry_sec=expiry_sec
)
return True
@property
def pgroup(self) -> str | None:
tmp = self.__client.get(f"{self.__bl}:pgroup")
@@ -458,16 +605,16 @@ class BeamlineConfig:
self.__client.set(f"{self.__bl}:{self.zoom_setting_string(mode)}", data.model_dump_json())
@property
def abr_meas_pos(self) -> Coordinate:
def abr_meas_pos(self) -> AerotechCoordinate:
tmp = self.__client.get(f"{self.__bl}:abr_meas_pos")
if tmp is None:
return ABR_POS_ALIGN_DEF
data_dict = json.loads(tmp)
return Coordinate(**data_dict)
return AerotechCoordinate(**data_dict)
@abr_meas_pos.setter
def abr_meas_pos(self, data: Coordinate):
def abr_meas_pos(self, data: AerotechCoordinate):
self.__client.set(f"{self.__bl}:abr_meas_pos", data.model_dump_json())
@property
@@ -625,3 +772,8 @@ class BeamlineConfig:
def increment_failed_mount_count(self) -> int:
return int(self.__client.incr(f"{self.__bl}:failed_mount_count"))
if __name__ == "__main__":
from aare.common.beamline import mx_beamline
cfg = BeamlineConfig(bl=mx_beamline())
cfg.allow_non_staff_request_from_staff = True
+233 -65
View File
@@ -1,4 +1,4 @@
import copy
import secrets
import time
import traceback
@@ -11,6 +11,7 @@ from typing import List, Tuple, Optional, Callable
import cv2
import numpy as np
from jfjoch_client import ScanResult, ScanResultImagesInner
import aare.common.face_detection as fd
from aare.daq import workflows
@@ -21,7 +22,7 @@ from aare.daq.config import BeamlineStateEnum
from aare.daq.devices import BeamlineDevices
from aare.daq.mlbox import MlBox
from aare.common.beamline import MXBeamline
from aare.common.coordinate import Coordinate, SmargonCoordinate
from aare.common.coordinate import Coordinate, SmargonCoordinate, AerotechCoordinate
from aare.common.diffraction_geometry import DiffractionGeometry
from aare.common.logger_config import setup_logger
from aare.common.models import (
@@ -34,6 +35,7 @@ from aare.common.models import (
from aare.common.raster_grid import RasterGridRequest, CompletedRasterGrid, CompletedRasterGridElem
from aare.common.rotation_scan import RotationScanRequest, CompletedRotationScan
from aare.common.sample_geometry import SampleGeometryModel
from aare.daq.spreadsheetupdater import beamline
from aare.devices.area_detector import AutoEnum
from aare.devices.jfjoch import JFJochWrapper
from aare.devices.mx_lib import clean_filename
@@ -44,7 +46,10 @@ from aare.common.exception_handler import (
MountingFailed,
WarningTellException,
CriticalTellException,
AXCFailed, SmargonCommunicationError, TellCommunicationError,
AXCFailed,
SmargonCommunicationError,
TellCommunicationError,
JFJochCommunicationError
)
logger = setup_logger("aareDAQ")
@@ -62,7 +67,7 @@ class AareDAQ:
self.__bl = bl.value.upper()
self.__aare = AareWrapper(bl)
self.__saved_box = None
self._smargon_trace_path = Path("logs") / "smargon_trace.csv"
self._smargon_trace_path = Path("/sls/mx/applications/logs") / "smargon_trace.csv"
self._face_detection_progress_cb: Callable[[dict], None] | None = None
self._last_sample_sync_ts = 0.0
self._sample_sync_min_interval_s = 2.0
@@ -358,10 +363,10 @@ class AareDAQ:
def samcam_settings(self, s: SampleCameraSettings):
self.__devs.samcam_settings = s
def tweak_abr_meas_pos(self, c: Coordinate):
def tweak_abr_meas_pos(self, c: AerotechCoordinate):
self.__cfg.set_busy(BeamlineStateEnum.SampleAlignment)
try:
new_meas_pos = self.__cfg.abr_meas_pos + c
new_meas_pos = AerotechCoordinate(at_mm=self.__cfg.abr_meas_pos.at_mm + c.at_mm)
self.__cfg.abr_meas_pos = new_meas_pos
self.__devs.aerotech_pos = new_meas_pos
self.__saved_box = None
@@ -405,13 +410,27 @@ class AareDAQ:
self.__cfg.state_busy = False
raise
def park_and_dry(self):
self.__cfg.try_set_busy(timeout=360)
try:
self.__devs.tell.check_enable_motion()
self.__devs.tell.wait_not_busy()
self.__devs.tell.set_in_mount_position(True)
self.__devs.tell.unmount(wait=True, timeout=60.0)
self.__devs.tell.dry(wait_cold=-1, wait=True, timeout=360.0)
self.__cfg.current_sample = None
self.__cfg.state_busy = False
except Exception as e:
self.__cfg.state_busy = False
logger.error(f"Failed to park and dry: {e}")
raise e
def __mount_failure_handler(self, mount_error):
pass
def __mount(self, target: SampleShortInfo | None):
self.__devs.smargon_move_home()
self.__devs.aerotech_pos = ABR_POS_MOUNT
self.__devs.aerotech_omega = ABR_OMEGA_MOUNT
#collimator should be down!!!
self.__devs.tell.check_enable_motion()
self.__devs.tell.wait_not_busy()
@@ -458,14 +477,14 @@ class AareDAQ:
except Exception as e:
self.__cfg.state_busy = False
logger.debug(f"Failed to mount sample: {e}")
self.__aare.sample_failed(target, f"Mount failed due to {e}")
# self.__aare.sample_failed(target, f"Mount failed due to {e}")
raise e
workflows.rse2sa(devs=self.__devs, cfg=self.__cfg)
self.__cfg.state_busy = False
if target is not None:
if target.db_id is not None:
self.__aare.sample_mounted(target)
self.save_screenshot_db(target.db_id, f"{target.db_id}_mounted")
# if target is not None:
# if target.db_id is not None:
# self.__aare.sample_mounted(target)
# self.save_screenshot_db(target.db_id, f"{target.db_id}_mounted")
@property
@@ -676,30 +695,120 @@ class AareDAQ:
else:
return None
def __setup_datacollection(self, request: RasterGridRequest | RasterGridRequest):
if request.dtz is not None:
logger.info(f'requesting dtz to move to {request.dtz}')
self.__cfg.dtz = request.dtz
if self.sample is not None and self.sample.db_id is not None:
self.save_screenshot_db(self.sample.db_id, f"{self.sample.db_id}_before_data_collection")
def __raster(self, r: RasterGridRequest):# -> CompletedRasterGridElem:
max_time = r.exp_time_s * r.n_x * r.n_y + 60
row_time = r.exp_time_s * r.n_x
row_width_mm = r.grid_size_mm.x * r.n_x
self.__set_state(BeamlineStateEnum.DataCollection)
save_smargon_position = self.__devs.smargon_pos
if r.smargon_top_left is not None:
delta_mm = self.sample_geometry.smargon_nudge(Coordinate(x=r.grid_size_mm.x / 2, y=r.grid_size_mm.y / 2))
self.__devs.set_smargon_pos(SmargonCoordinate(sh_mm=r.smargon_top_left.sh_mm + delta_mm,
phi_deg=r.smargon_top_left.phi_deg,
chi_deg=r.smargon_top_left.chi_deg))
self.__devs.aerotech_omega = r.omega_deg
if request.transmission is not None:
logger.info(f'requesting transmission to move to {request.transmission}')
self.__devs.transmission = request.transmission
if hasattr(request, 'start') and request.start is not None:
self.__devs.smargon_pos = request.start
elif hasattr(request, 'smargon_top_left') and request.smargon_top_left is not None:
self.__devs.set_smargon_pos(SmargonCoordinate(sh_mm=request.smargon_top_left.sh_mm,
phi_deg=request.smargon_top_left.phi_deg,
chi_deg=request.smargon_top_left.chi_deg))
#if request.transmission is not None:
# self.__devs.transmission.wait()
self.__devs.smargon_wait(timeout=180)
return
def _build_fake_scan_result(self, *, file_prefix: str | None, image_count: int) -> ScanResult:
images = [
ScanResultImagesInner(
number=i,
efficiency=1.0,
bkg=0.0,
spots=0,
spots_low_res=0,
spots_indexed=0,
index=0,
b=0.0,
)
for i in range(max(1, image_count))
]
return ScanResult(file_prefix=file_prefix, images=images)
def _build_fake_rotation_result(self, request: RotationScanRequest) -> CompletedRotationScan:
result = self._build_fake_scan_result(
file_prefix=request.file_prefix,
image_count=request.steps,
)
return CompletedRotationScan(
request=copy.deepcopy(request),
result=result,
)
def _build_fake_raster_result(self, request: RasterGridRequest) -> CompletedRasterGridElem:
result = self._build_fake_scan_result(
file_prefix=request.file_prefix,
image_count=request.n_x * request.n_y,
)
return CompletedRasterGridElem(
request=copy.deepcopy(request),
result=result,
centre_of_mass=None,
)
def __raster(self, request: RasterGridRequest) -> CompletedRasterGridElem:
self.__devs.aerotech_omega = request.omega_deg
self.__setup_datacollection(request=request)
status = self.status
print(f"raster status {status}")
print(f'raster grid request: {r}')
#self.__jfjoch.measure_raster(r, status)
self.__devs.aerotech.run_grid_scan(cell_height_mm=r.grid_size_mm.y, num_rows=r.n_x,
row_width_mm=row_width_mm, time_per_row_s=row_time, task_id=3)
self.__devs.aerotech_pos = self.__cfg.abr_meas_pos
#result = self.__jfjoch.wait_till_done(60)
return None
#return CompletedRasterGridElem(request=copy.deepcopy(r), result=, centre_of_mass=None)
logger.info(f"raster status {status}")
logger.info(f'raster grid request: {request}')
total_time = request.exp_time_s*request.n_x*request.n_y
if not self.__cfg.simulated_detector:
try:
logger.info('initialise detector')
self.__jfjoch.measure_raster(request, status)
logger.info('detector initialised')
except JFJochCommunicationError as e:
logger.warning(f"Failed to communicate with jfjoch: {e}")
else:
logger.info("Simulated detector mode enabled; returning fake zero raster result.")
try:
self.__devs.aerotech.grid_scan(grid_elem_size_y_um=request.grid_size_mm.y*1000,
grid_elem_size_x_um=request.grid_size_mm.x*1000,
grid_elem_count_x=request.n_x,
grid_elem_count_y=request.n_y,
time_sec=request.exp_time_s,
run_async=True)
self.__devs.aerotech.wait_till_done(timeout=int(round(total_time*2,0)))
self.__devs.aerotech_pos = self.__cfg.abr_meas_pos
try:
result = self.__jfjoch.wait_till_done(60)
except JFJochCommunicationError as e:
logger.warning(f"Failed to communicate with jfjoch: {e}")
return self._build_fake_raster_result(request)
if result is None:
logger.warning("JFJoch returned no ScanResult; using fake result for raster scan.")
return self._build_fake_raster_result(request)
return CompletedRasterGridElem(request=copy.deepcopy(request), result=result, centre_of_mass=None)
except JFJochCommunicationError:
raise
except Exception as e:
logger.error(f"Failed during raster: {e}")
raise Exception(f"Failed during raster: {e}") from e
def measure_raster(self, r: RasterGridRequest, auto: bool) -> CompletedRasterGrid:
self.__cfg.try_set_busy(timeout=ceil(360))
@@ -708,25 +817,96 @@ class AareDAQ:
result = self.__auto_center(r)
else:
raster_result = self.__raster(r)
#result = CompletedRasterGrid(r=[raster_result])
result = CompletedRasterGrid(r=[raster_result])
self.__set_state(BeamlineStateEnum.SampleAlignment)
self.__cfg.state_busy = False
return CompletedRasterGrid(r=[])
#return result
#return CompletedRasterGrid(r=[])
return result
except Exception as e:
try:
self.__aare.axc_failed(self.sample)
except Exception as axc_e:
logger.error(f"Exception while reporting AXC failure: {axc_e}")
self.__set_state(BeamlineStateEnum.SampleAlignment)
if not self.busy:
self.__cfg.try_set_busy(timeout=ceil(360))
if self.status.state != BeamlineStateEnum.SampleAlignment:
self.__set_state(BeamlineStateEnum.SampleAlignment)
self.__cfg.state_busy = False
raise e
def __rotation(self, request: RotationScanRequest) -> CompletedRotationScan:
return None
omega_start = self.omega
status = self.status
#self.__aare.create_rotation_run(self.sample, request, status)
total_time = request.exp_time_s * request.steps
try:
if self.__cfg.simulated_detector:
logger.info("Simulated detector mode enabled; skipping JFJoch start.")
else:
self.__jfjoch.measure_rotation(request, status, self.__cfg.xrf)
if request.screening:
self.__devs.aerotech.screening_scan(
rotation_deg=request.steps*request.incr_omega_deg,
wedge_deg=request.wedge_omega_deg,
time_sec=total_time,
steps=request.steps,
run_async=True,
)
else:
self.__devs.aerotech.rotation_scan(
rotation_deg=request.steps*request.incr_omega_deg,
time_sec=total_time,
start_pos_deg=request.start_omega_deg,
run_async=True,
)
#Is this for helical scans...? do we do smargon scans?
if request.start is not None and request.end is not None:
smargon_time_step = request.time_sec / float(request.steps)
pos_step = (request.end.sh_mm - request.start.sh_mm) * (1.0 / float(request.steps))
for i in range(request.steps):
self.__devs.smargon.target = SmargonCoordinate(
sh_mm=request.start.sh_mm + pos_step * i
)
time.sleep(smargon_time_step)
self.__devs.aerotech.wait_till_done(timeout=int(round(total_time + 60,0)))
self.__devs.aerotech_omega = omega_start
# self.__aare.sample_collected(self.sample)
if self.__cfg.simulated_detector:
logger.warning("Detector in simulation mode, returning fake zero rotation result.")
result = self._build_fake_rotation_result(request)
else:
# Let JFJochCommunicationError propagate
result = self.__jfjoch.wait_till_done(60)
# if self.sample is not None and self.sample.db_id is not None:
# self.save_screenshot_db(self.sample.db_id, "after_dc")
# try:
# self.__aare.ingest_scan(sample=self.sample, result=result,
# geom=self.sample_geometry, beam_mark_pxl=self.__cfg.get_beam_mark(self.zoom))
#
# except Exception as e:
# logger.error(f"Exception ingesting scan: {e}")
except JFJochCommunicationError:
raise
except Exception as e:
logger.error(f"Exception during rotation scan: {e}")
raise
return CompletedRotationScan(request=copy.deepcopy(request), result=result)
def measure_rotation(self, request: RotationScanRequest) -> CompletedRotationScan:
total_time = request.exp_time_s * request.steps
logger.info(f"received rotation scan request: {request}, total time: {total_time}s, steps: {request.steps}")
self.__cfg.try_set_busy(timeout=ceil(total_time + 360))
try:
@@ -739,31 +919,6 @@ class AareDAQ:
self.__cfg.state_busy = False
raise e
def get_background(self):
#if self.sample is not None:
# raise Exception("Background cannot be measured when sample is mounted")
self.__cfg.try_set_busy(timeout=360)
self.__set_state(BeamlineStateEnum.SampleAlignment)
try:
self.__devs.lamp_light = 2.5
self.__cfg.zoom_mode = ZoomModeEnum.LoopCenter
zoom_settings = self.__cfg.zoom_settings.z
for zoom_value in zoom_settings:
exposure = zoom_settings[zoom_value].exposure
gain = zoom_settings[zoom_value].gain
self.__devs.samcam_settings = SampleCameraSettings(exposure=exposure, gain=gain)
self.__devs.set_zoom(zoom_value, wait=True)
time.sleep(5.0) # Wait for settings to stabilize
image = self.__devs.samcam_get_image(gray=False)
self.__cfg.put_alc_bkg(zoom_value, exposure, gain, image)
bgr_array = cv2.cvtColor(image, cv2.COLOR_RGB2BGR)
cv2.imwrite(f"bkg{zoom_value:.0f}_{exposure*1000:.0f}_{gain:.0f}.jpg", bgr_array)
self.__cfg.state_busy = False
except Exception as e:
self.__cfg.state_busy = False
raise e
@property
def dtz(self) -> float:
tmp = self.__cfg.dtz
@@ -823,8 +978,8 @@ class AareDAQ:
@property
def sample_geometry(self) -> SampleGeometryModel:
zoom = self.__devs.zoom
aerotech_pos_ref = self.__cfg.abr_meas_pos
aerotech_pos = self.__devs.aerotech_pos
aerotech_pos_ref = self.__cfg.abr_meas_pos.at_mm
aerotech_pos = self.__devs.aerotech_pos.at_mm
sample_geom = SampleGeometryModel(
beam_location_pxl=self.__cfg.beam_mark_coeff.apply(zoom),
pixel_in_mm=self.__cfg.pixel_to_mm(zoom),
@@ -1632,6 +1787,16 @@ class AareDAQ:
# Keep status flowing even if Tell code throws something unexpected
return self.__cfg.current_sample, False, f"TELL unavailable: {e}"
def _aerotech_status(self) -> tuple[bool, str | None]:
aerotech_ok = True
aerotech_err: str | None = None
try:
_ = self.__devs.aerotech.status()
except Exception as e:
aerotech_ok = False
aerotech_err = str(e)
return aerotech_ok, aerotech_err
def _safe_geom(self) -> tuple[SampleGeometryModel, bool, str | None]:
"""
Return (geom, smargon_connected, smargon_error) without raising.
@@ -1714,6 +1879,7 @@ class AareDAQ:
def status(self) -> DAQStatusModel:
safe_sample, tell_ok, tell_err = self._safe_sample()
safe_geom, smargon_ok, smargon_err = self._safe_geom()
aerotech_ok, aerotech_err = self._aerotech_status()
return DAQStatusModel(
state=self.state,
@@ -1731,11 +1897,13 @@ class AareDAQ:
tell_error=tell_err,
smargon_connected=smargon_ok,
smargon_error=smargon_err,
aerotech_connected=aerotech_ok,
aerotech_error=aerotech_err,
)
def cancel(self):
if self.__cfg.state == BeamlineStateEnum.DataCollection:
self.__devs.aerotech_stop()
self.__devs.aerotech.cancel()
self.__jfjoch.cancel()
def anneal(self, time_s: float):
+38 -26
View File
@@ -5,15 +5,11 @@
# - setter with option to do sync/async move
# - property setter, which assumes that sync move is done (excl. zoom, which is async by default)
import time
from enum import Enum
import numpy as np
from epics import PV
from fontTools.feaLib.ast import deviceToString
from aare.common.beamline import MXBeamline
from aare.common.coordinate import Coordinate, SmargonCoordinate
from aare.common.coordinate import SmargonCoordinate, AerotechCoordinate
from aare.common.models import SampleCameraSettings, StagePositionEnum
from aare.devices import smargon, aerotech
from aare.devices.area_detector import epicsAD, AutoEnum
@@ -27,8 +23,7 @@ class BeamlineDevices:
def __init__(self, beamline: MXBeamline):
BEAMLINE = beamline.value.upper()
self.tell = make_tell_client(beamline)
self.__aerotech = aerotech.AerotechControllerEpics(beamline)
self.aerotech = aerotech.AerotechController(controller_ip="129.129.118.96")
self.aerotech = aerotech.AerotechController(beamline)
self.__smargon = smargon.Smargon(beamline)
self.__ring_current_pv = PV(f"ARS07-DPCT-0100:CURR")
@@ -91,19 +86,24 @@ class BeamlineDevices:
self.__cryojet_x = MyMotor(f"{BEAMLINE}-ES-CS:TRX") #currently in is 5 out is 15?
self.magnet_position_sensor = PV(f"{BEAMLINE}-ES-DF1:CBOX-CMP1")
self.magnet_position_sensor_readout = PV(f"{BEAMLINE}-ES-DF1:CBOX-USER1")
# Transmission
@property
def transmission(self) -> float:
return 1.0
def set_transmission(self, value: float, /, wait: bool = True):
pass
@transmission.setter
def transmission(self, value: float):
self.set_transmission(value, wait=False)
def set_transmission(self, value: float, /, wait: bool = True):
pass
# Lamp light
@property
def lamp_light(self) -> float:
@@ -262,8 +262,6 @@ class BeamlineDevices:
"""
return int(self.__sample_cam.uid.get())
# Detector Z
@property
def dtz(self) -> float:
@@ -284,32 +282,46 @@ class BeamlineDevices:
def dtz_high(self) -> float:
return self.__dtz.get("HLM")
# Aertoech Automation1
@property
def aerotech_pos(self) -> Coordinate:
return Coordinate(x=self.__aerotech.gmx.readback,
y=self.__aerotech.gmy.readback, z=self.__aerotech.gmz.readback)
def aerotech_pos(self) -> AerotechCoordinate:
#TODO add u/omega????
return self.aerotech.get_position()
@aerotech_pos.setter
def aerotech_pos(self, pos: Coordinate):
# self.__aerotech.gmx.move(pos.x, wait=False)
# self.__aerotech.gmy.move(pos.y, wait=False)
# self.__aerotech.gmz.move(pos.z, wait=False)
self.__aerotech.gmx.move(pos.x, wait=True)
self.__aerotech.gmy.move(pos.y, wait=True)
self.__aerotech.gmz.move(pos.z, wait=True)
#TODO add wait pos?
def aerotech_pos(
self,
coord: AerotechCoordinate,
/,
wait: bool = True,
incremental: bool = False,
):
self.aerotech.position(
coord,
wait=wait,
incremental=incremental,
)
@property
def aerotech_omega(self) -> float:
return self.__aerotech.omega.readback
return self.aerotech.status().u.pos
@aerotech_omega.setter
def aerotech_omega(self, val: float):
self.set_aerotech_omega(val, wait=True)
def set_aerotech_omega(self, val: float, /, wait: bool = True):
self.__aerotech.omega.move(val, wait=wait)
def set_aerotech_omega(
self,
val: float,
/,
wait: bool = True,
incremental: bool = False,
):
target = AerotechCoordinate(omega_deg=val)
return self.aerotech.position(
target,
wait=wait,
incremental=incremental,
)
@property
def aerotech_lock(self) -> bool:
+1 -1
View File
@@ -30,7 +30,7 @@ class MlBox:
elif bl == MXBeamline.X06DA:
self.__url = "http://mx-aare-test.psi.ch:8002/predict/?model=best_v8_20102025.pt"
elif bl == MXBeamline.X10SA:
self.__url = "http://x10sa-spark-01.psi.ch:8002/predict/?model=best_v12_22092025.engine"
self.__url = "http://x10sa-spark-01.psi.ch:8002/predict/?model=best_yolo26l-seg-overlap-false_2026-03-16.engine"#v12_22092025.engine"
elif bl == MXBeamline.X06SA:
self.__url = ""
raise NotImplemented(f"MLBox not implemente for {bl}")
+181 -11
View File
@@ -7,6 +7,9 @@ import json
import cv2
import urllib3
import uvicorn
from aare.common.coordinate import SmargonCoordinate, Coordinate, AerotechCoordinate
from aare.common.auth_models import BatonStatus, BatonRequestStatus
from aare.common.coordinate import SmargonCoordinate, Coordinate
from aare.common.error_codes import export_error_codes, export_error_codes_grouped
from aare.common.logger_config import setup_logger
@@ -34,6 +37,12 @@ from aare.common.exception_handler import (
SampleException,
UserRightsException,
)
from aare.common.automation_queue_manager import WorkflowRedisManager
from aare.common.automation_workflow import STATE_REGISTRY, HANDLER_REGISTRY, SIMULATED_HANDLER_REGISTRY
from aare.daq.automation_runner import PersistentWorkflowRunner
from aare.daq.automation_api_router import router as workflow_router, set_workflow_dependencies
logger = setup_logger("aareDAQ")
app = FastAPI()
register_exception_handlers(app)
@@ -45,6 +54,31 @@ bl = mx_beamline()
cfg = BeamlineConfig(bl)
daq = AareDAQ(cfg, bl)
# ─────────────────────────────────────────────
# Initialize workflow system
# ─────────────────────────────────────────────
USE_SIMULATED_WORKFLOW = os.getenv("WORKFLOW_SIMULATION", "0") == "1"
workflow_redis_manager = WorkflowRedisManager(
client=cfg._BeamlineConfig__client, # Reuse existing Redis connection
beamline=bl.value,
)
workflow_runner = PersistentWorkflowRunner(
redis_manager=workflow_redis_manager,
registry=STATE_REGISTRY,
handlers=SIMULATED_HANDLER_REGISTRY if USE_SIMULATED_WORKFLOW else HANDLER_REGISTRY,
)
if USE_SIMULATED_WORKFLOW:
logger.warning("⚠️ Workflow system running in SIMULATION mode - no actual DAQ operations")
set_workflow_dependencies(workflow_redis_manager, workflow_runner, cfg)
# Include the workflow router
app.include_router(workflow_router)
try:
daq.sync_current_sample_from_tell(force=True)
except Exception as e:
@@ -205,8 +239,7 @@ async def smargon(val: SmargonCoordinate, token: str = Depends(oauth2_scheme)):
@app.post("/beamline/tweak_abr_meas_pos")
async def tweak_abr_meas_pos(val: Coordinate, token: str = Depends(oauth2_scheme)):
logger.debug(f"Setting abr to {val}")
async def tweak_abr_meas_pos(val: AerotechCoordinate, token: str = Depends(oauth2_scheme)):
auth.check_jwt_staff(cfg, auth.parse_token(token))
daq.tweak_abr_meas_pos(val)
return "OK"
@@ -316,6 +349,15 @@ async def sample(token: str = Depends(oauth2_scheme)) -> SampleShortInfo:
location=s.location
)
async def park_and_dry(token: str = Depends(oauth2_scheme)):
token_data = auth.parse_token(token)
auth.check_jwt_ro(cfg, auth.parse_token(token))
daq.park_and_dry()
return {
"ok": True,
"message": "TELL has been dryed and parked",
}
@app.post("/sample/mount")
async def mount(dbid: int, token: str = Depends(oauth2_scheme), reference: bool = False):
@@ -602,13 +644,6 @@ async def cancel(token: str = Depends(oauth2_scheme)):
daq.cancel()
# ALC routines
@app.post("/alc/background")
async def alc_background(token: str = Depends(oauth2_scheme)) -> str:
auth.check_jwt_rw(cfg, auth.parse_token(token))
daq.get_background()
return "OK"
@app.post("/alc/center_loop")
async def alc_center_loop(token: str = Depends(oauth2_scheme)) -> str:
logger.debug(f"ALC")
@@ -659,11 +694,19 @@ async def pgroup(token: str = Depends(oauth2_scheme)) -> str:
@app.put("/access/pgroup")
async def set_pgroup(val: str, token: str = Depends(oauth2_scheme)) -> str:
auth.check_jwt_ro(cfg, auth.parse_token(token))
data = auth.parse_token(token)
holder = cfg.baton_holder
is_current_holder = holder is not None and holder.session == data.session
# Staff or current baton holder may change p-group even if it is not currently active.
# Everyone else must still belong to the active p-group.
if not (data.staff or is_current_holder):
auth.check_jwt_ro(cfg, data)
cfg.pgroup = val
return "OK"
@app.delete("/access/pgroup")
async def del_pgroup(token: str = Depends(oauth2_scheme)) -> str:
auth.check_jwt_ro(cfg, auth.parse_token(token))
@@ -696,6 +739,133 @@ async def force_current_session(token: str = Depends(oauth2_scheme)) -> str:
auth.force_current_sesion(cfg, data)
return "OK"
# ========== BATON CONTROL ENDPOINTS ==========
@app.get("/baton/status")
async def baton_status(token: str = Depends(oauth2_scheme)) -> BatonStatus:
"""Get the current baton status for the requesting user."""
data = auth.parse_token(token)
auth.resolve_baton_timeout_if_needed(cfg)
return auth.get_baton_status(cfg, data)
@app.post("/baton/request")
async def baton_request(token: str = Depends(oauth2_scheme)) -> dict:
"""
Request control (baton) of the beamline.
- If vacant: granted immediately
- If staff requesting: granted immediately (or queued if busy)
- If same level: creates pending request with timeout
- Non-staff cannot request from staff
"""
logger.debug(cfg.allow_non_staff_request_from_staff)
data = auth.parse_token(token)
return auth.request_baton(cfg, data)
@app.post("/baton/respond")
async def baton_respond(accept: bool, token: str = Depends(oauth2_scheme)) -> dict:
"""
Current baton holder responds to a pending request.
- accept=true: transfers baton (or queues if busy)
- accept=false: refuses the request
"""
data = auth.parse_token(token)
return auth.respond_to_baton_request(cfg, data, accept)
@app.post("/baton/release")
async def baton_release(token: str = Depends(oauth2_scheme)) -> dict:
"""Voluntarily release the baton, making the beamline vacant."""
data = auth.parse_token(token)
return auth.release_baton(cfg, data)
@app.post("/baton/cancel")
async def baton_cancel(token: str = Depends(oauth2_scheme)) -> dict:
"""Cancel your own pending baton request."""
data = auth.parse_token(token)
return auth.cancel_baton_request(cfg, data)
@app.get("/baton/check_timeout")
async def baton_check_timeout(token: str = Depends(oauth2_scheme)) -> dict:
"""
Check if a pending request has timed out and process it.
Called by GUI to poll for timeout completion.
"""
data = auth.parse_token(token)
auth.resolve_baton_timeout_if_needed(cfg)
pending = cfg.pending_baton_request
if pending is None:
if cfg.baton_holder and cfg.baton_holder.session == data.session:
return {"granted": True, "message": "Baton acquired!"}
queued = cfg.queued_baton_transfer
if queued and queued.target_session == data.session:
return {"queued": True, "message": "Transfer queued"}
return {"no_pending": True}
if pending.requester_session != data.session:
return {"not_your_request": True}
if pending.status == BatonRequestStatus.REFUSED:
cfg.clear_pending_baton_request()
return {"refused": True, "message": "Request refused"}
elapsed = time.time() - pending.created_at
if elapsed < pending.timeout_seconds:
return {
"pending": True,
"remaining_seconds": pending.timeout_seconds - elapsed
}
return auth.request_baton(cfg, data)
@app.put("/access/allow_non_staff_request_from_staff")
async def set_allow_non_staff_request_from_staff(val: bool, token: str = Depends(oauth2_scheme)) -> str:
data = auth.parse_token(token)
auth.check_jwt_staff_only(data)
cfg.allow_non_staff_request_from_staff = val
return "OK"
async def baton_status_event_stream(data: TokenData) -> AsyncGenerator[str, None]:
"""SSE stream for baton status updates."""
last_status = None
try:
while True:
auth.resolve_baton_timeout_if_needed(cfg)
new_baton_status = auth.get_baton_status(cfg, data)
status_json = new_baton_status.model_dump_json()
if status_json != last_status:
last_status = status_json
yield f"data: {status_json}\n\n"
cfg.process_queued_transfer_if_ready(auth.SESSION_EXPIRE_SECONDS)
await asyncio.sleep(0.5)
except asyncio.CancelledError:
return
@app.get("/sse/baton")
async def sse_baton(token: str = Depends(oauth2_scheme)):
"""SSE endpoint for real-time baton status updates."""
data = auth.parse_token(token)
return StreamingResponse(
baton_status_event_stream(data),
media_type="text/event-stream",
headers={
"Cache-Control": "no-cache",
"Connection": "keep-alive",
"Access-Control-Allow-Origin": "*",
"Access-Control-Allow-Headers": "Cache-Control"
}
)
@app.get("/beamline/settings")
async def get_settings(token: str = Depends(oauth2_scheme)) -> BeamlineSettingsModel:
+35 -2
View File
@@ -17,7 +17,10 @@ from aare.common.exception_handler import (
AuthenticationException,
SampleException,
UserRightsException,
SmargonCommunicationError, TellCommunicationError
SmargonCommunicationError,
TellCommunicationError,
JFJochCommunicationError,
AerotechCommunicationError,
)
logger = setup_logger("aareDAQ")
@@ -145,9 +148,39 @@ def register_exception_handlers(app) -> None:
),
)
@app.exception_handler(JFJochCommunicationError)
async def jfjoch_comm_handler(request: Request, exc: JFJochCommunicationError) -> JSONResponse:
return JSONResponse(
status_code=api_status.HTTP_503_SERVICE_UNAVAILABLE,
content=_error_payload(
code="JFJOCH_UNAVAILABLE",
message=str(exc) or "JFJoch detector is unavailable",
extra={
"operation": getattr(exc, "operation", None),
"endpoint": getattr(exc, "endpoint", None),
"base_url": getattr(exc, "base_url", None),
},
),
)
@app.exception_handler(AerotechCommunicationError)
async def aerotech_comm_handler(request: Request, exc: AerotechCommunicationError) -> JSONResponse:
return JSONResponse(
status_code=api_status.HTTP_503_SERVICE_UNAVAILABLE,
content=_error_payload(
code="AEROTECH_UNAVAILABLE",
message=str(exc) or "Aerotech is unavailable",
extra={
"operation": getattr(exc, "operation", None),
"endpoint": getattr(exc, "endpoint", None),
"base_url": getattr(exc, "base_url", None),
},
),
)
@app.exception_handler(Exception)
async def unhandled_exception_handler(request: Request, exc: Exception) -> JSONResponse:
logger.exception("Unhandled server exception")
logger.exception(f"Unhandled server exception: {exc}")
return JSONResponse(
status_code=api_status.HTTP_500_INTERNAL_SERVER_ERROR,
content=_error_payload(code="INTERNAL_SERVER_ERROR", message=str(exc) or "Internal server error"),
+28 -5
View File
@@ -9,7 +9,9 @@ from aare.common.beamline import MXBeamline, mx_beamline
beamline = mx_beamline()
SLOT_IDENTIFIER = beamline.value.upper()
WS_URL = f"wss://mx-db-01.psi.ch/dispatcher/protected_router/tell_runner/ws/samples-spreadsheet/{SLOT_IDENTIFIER}"
#WS_URL = f"wss://mx-db-01.psi.ch/dispatcher/protected_router/tell_runner/ws/samples-spreadsheet/{SLOT_IDENTIFIER}"
WS_URL = f"wss://mx-aaredb-dmz-01.psi.ch/dispatcher/protected_router/tell_runner/ws/samples-spreadsheet/{SLOT_IDENTIFIER}"
# Ensure the environment variable for the shared password is set
password = os.getenv("AAREDB_SHARED_PASSWORD")
@@ -134,6 +136,25 @@ def main():
"""
while True:
try:
import ssl
import websocket
# FORCE a clean context
context = ssl.create_default_context(ssl.Purpose.SERVER_AUTH)
context.load_verify_locations(cafile="/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem")
context.load_cert_chain(
certfile="/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_dmz-01.crt",
keyfile="/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_dmz-01.key"
)
# Explicitly set the SNI hostname to match NGINX server_name
# This is often what's missing when NGINX says "No cert sent"
ssl_opt = {
"context": context,
"server_hostname": "mx-aaredb-dmz-01.psi.ch",
"check_hostname": True
}
ws = websocket.WebSocketApp(
WS_URL,
header=WS_HEADERS,
@@ -142,11 +163,13 @@ def main():
on_close=on_close,
on_open=on_open,
)
ws.run_forever(sslopt={"cert_reqs": 0})
except Exception as e:
print(f"[MAIN][ERROR] ws.run_forever() crashed with: {e}")
print("[WS][INFO] WebSocket connection lost. Reconnecting in 5 seconds...")
print(f"[WS][INFO] Connecting to {WS_URL}...")
ws.run_forever(sslopt=ssl_opt)
except Exception as e:
print(f"[MAIN][ERROR] WebSocket connection failed: {e}")
time.sleep(5)
if __name__ == "__main__":
+64 -19
View File
@@ -4,7 +4,6 @@ import threading
import websocket
import sseclient
import requests
import time
from aareDB.models import PuckWithTellPosition
@@ -19,7 +18,8 @@ logger = setup_logger("aareDAQ")
# Configuration
beamline = mx_beamline()
SLOT_IDENTIFIER = beamline.value.upper()
WS_URL = f"wss://mx-db-01.psi.ch/dispatcher/protected_router/wstell/ws/slot/{SLOT_IDENTIFIER}"
#WS_URL = f"wss://mx-db-01.psi.ch/dispatcher/protected_router/wstell/ws/slot/{SLOT_IDENTIFIER}"
WS_URL = f"wss://mx-aaredb-dmz-01.psi.ch/dispatcher/protected_router/wstell/ws/slot/{SLOT_IDENTIFIER}"
#WS_URL = f"wss://localhost:8001/protected_router/wstell/ws/slot/{SLOT_IDENTIFIER}"
WS_HEADERS = [f"X-Shared-Password: {os.getenv('AAREDB_SHARED_PASSWORD')}"]
print(WS_HEADERS)
@@ -39,19 +39,40 @@ def listen_to_sse():
print(f"[SSE][WARN] No TELL URL configured SSE listener not started. (tell_client.url={tell_client.url})")
return
sse_url = tell_client.url + "/events"
try:
#response = requests.get(sse_url, stream=True)
client = sseclient.SSEClient(sse_url)
print("[SSE][listen_to_sse] Initial detected pucks fetch on connect")
handle_tell_change_event()
while True:
try:
print(f"[SSE][INFO] Attempting to connect to {sse_url}...")
if sse_url.startswith("https://mx-aaredb-dmz-01"): #"https://mx-db-01"
# mTLS path
import requests
#cert_pair = ("/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_db-01.crt", "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_db-01.key")
cert_pair = ("/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_dmz-01.crt",
"/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_dmz-01.key")
#ca_root = "/etc/ssl/certs/secrets/mx-db-01_Full_Chain_CA.pem"
ca_root = "/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem"
response = requests.get(sse_url, stream=True, cert=cert_pair, verify=ca_root)
response.raise_for_status()
client = sseclient.SSEClient(response)
else:
# Robot path (PC17488) - Pass URL STRING directly
# SSEClient will handle the simple HTTP GET itself
client = sseclient.SSEClient(sse_url)
for event in client.events():
print(f"event = {event.event} with data: {event.data}")
if event.event == "DewarContentUpdate":
on_sse_event(event)
except Exception as exc:
print(f"[SSE][listen_to_sse][ERROR] Failed to connect to {sse_url}: {exc}")
print("[SSE][listen_to_sse] Initial detected pucks fetch on connect")
handle_tell_change_event()
# Compatibility: some SSEClient versions are iterable, others expose .events().
events_iter = client.events() if hasattr(client, "events") else iter(client)
for event in events_iter:
# print(f"event = {event.event} with data: {event.data}")
if event.event == "DewarContentUpdate":
on_sse_event(event)
except Exception as exc:
print(f"[SSE][listen_to_sse][ERROR] Connection lost or failed: {exc}")
print("[SSE][INFO] Reconnecting to SSE in 5 seconds...")
time.sleep(5)
def compare_and_report_change(old, new, key_func):
"""
@@ -140,14 +161,34 @@ def on_close(ws, close_status_code, close_msg):
def on_open(ws):
print("[WS][OPEN] WebSocket opened.")
def main():
# Start SSE listener in a separate background thread
sse_thread = threading.Thread(target=listen_to_sse, daemon=True)
sse_thread.start()
# Main thread runs websocket client loop
"""
Main function to initiate WebSocket connection with mTLS.
"""
while True:
try:
import ssl
# 1. Create a modern SSL Context for a TLS Client
context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
# 2. Load the CA to verify the NGINX server's identity
context.load_verify_locations(cafile="/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem")
# 3. Load the Client Certificate and Key (mTLS)
# Using the 'dmz-01' paths that worked in your curl
context.load_cert_chain(
certfile="/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_dmz-01.crt",
keyfile="/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_dmz-01.key"
)
# Optional: Ensure hostname matching is active (recommended)
context.check_hostname = True
ws = websocket.WebSocketApp(
WS_URL,
header=WS_HEADERS,
@@ -156,12 +197,16 @@ def main():
on_close=on_close,
on_open=on_open,
)
ws.run_forever(sslopt={"cert_reqs": 0})
except Exception as e:
print(f"[MAIN][ERROR] ws.run_forever() crashed with: {e}")
print("[WS][INFO] WebSocket connection lost. Reconnecting in 5 seconds...")
# 4. Pass the context directly via sslopt
print(f"[WS][INFO] Connecting to {WS_URL} using mTLS...")
ws.run_forever(sslopt={"context": context})
except Exception as e:
print(f"[MAIN][ERROR] WebSocket connection failed: {e}")
print("[WS][INFO] Reconnecting in 5 seconds...")
time.sleep(5)
if __name__ == "__main__":
main()
main()
+2 -4
View File
@@ -34,8 +34,7 @@ def common_2rse(devs: BeamlineDevices, cfg: BeamlineConfig):
print('move smargon home')
devs.smargon_move_home()
print('try to move aerotech')
devs.aerotech_pos = Coordinate(x=0.0, y=0.0, z=0.0)
devs.aerotech_omega = 0.0
devs.aerotech_pos = ABR_POS_MOUNT
#print('move bl to park')
#devs.reflector_up = StagePositionEnum.PARK
print('move bs to park')
@@ -53,8 +52,7 @@ def sa2se(devs: BeamlineDevices, cfg: BeamlineConfig):
print('move smargon home')
devs.smargon_move_home()
print('try to move aerotech')
devs.aerotech_pos = Coordinate(x=0.0, y=0.0, z=0.0)
devs.aerotech_omega = 0.0
devs.aerotech_pos = ABR_POS_MOUNT
print('move bl to park')
devs.reflector_up = StagePositionEnum.PARK
print('move bs to park')
+221 -361
View File
@@ -1,386 +1,246 @@
import os
import sys
import threading
import time
from enum import Enum
from typing import Union
from typing import Optional, Union
import automation1 as a1
from epics import PV, Motor, caput, caget
from aarescan_client.models.grid_request import GridRequest
from aarescan_client.models.screen_request import ScreenRequest
from aare.common.beamline import MXBeamline, mx_beamline
from aare.devices.my_motor import MyMotor
from aarescan_client import ApiClient, Status, RotationRequest, Configuration, DefaultApi, AxisStatus, Target
class TaskEnum(Enum):
TASK_0 = 0
TASK_1 = 1
TASK_2 = 2
TASK_3 = 3
TASK_4 = 4
from aare.common.coordinate import AerotechCoordinate, Coordinate
from aare.common.exception_handler import AerotechCommunicationError
class AxisEnum(Enum):
X = "gmx"
Y = "gmy"
Z = "gmz"
OMEGA = "Omega"
AEROTECH_HOME = AerotechCoordinate(at_mm=Coordinate(x=0, y=0, z=0), omega_deg=0)
class AerotechRunEnum(Enum):
STOP = 0
START = 1
RUN = 2
LOAD = 3
PAUSE = 4
RESET = 5
class VariableTypeEnum(Enum):
INT = 0
REAL = 1
STRING = 2
class AerotechController(object):
class AerotechControllerEpics:
def __init__(self, beamline: MXBeamline):
BEAMLINE = beamline.value.upper()
self.aerotech_pv_prefix = f"{BEAMLINE}-ES-DF1"
trx_prefix = f"{self.aerotech_pv_prefix}:TRX1"
try_prefix = f"{self.aerotech_pv_prefix}:TRY1"
trz_prefix = f"{self.aerotech_pv_prefix}:TRZ1"
rotu_prefix = f"{self.aerotech_pv_prefix}:ROTU"
self.gmx = MyMotor(trx_prefix)
self.gmy = MyMotor(try_prefix)
self.gmz = MyMotor(trz_prefix)
self.omega = MyMotor(rotu_prefix)
self.__enable_all = PV(f"{self.aerotech_pv_prefix}:EnableAll")
self.__disable_all = PV(f"{self.aerotech_pv_prefix}:DisableAll")
self.__acknowledge_all = PV(f"{self.aerotech_pv_prefix}:AckAll")
self.__stop_all = PV(f"{self.aerotech_pv_prefix}:StopAll")
self.__task_filename = PV(f"{self.aerotech_pv_prefix}:TASK:FILENAME")
self.__task_id = PV(f"{self.aerotech_pv_prefix}:TASK:TASKIDX") # values can be 1 to 4 DO NOT USE 0!!!!
self.__task_run_enum = PV(f"{self.aerotech_pv_prefix}:TASK:SWITCH") # 0 Stop, 1 Start, 2 Run,
# 3 Load, 4 Pause, 5 Reset
self.enable_all()
def __init__(self, bl: MXBeamline):
if bl == MXBeamline.X06DA:
self.__simulated = False
self.__base = "http://x06da-smargopolo.psi.ch:5234"
elif bl == MXBeamline.X10SA:
self.__simulated = False
self.__base = "http://x10sa-smargopolo.psi.ch:5234"
elif bl == MXBeamline.SIMULATED:
self.__simulated = True
self.__pos = AEROTECH_HOME
self.__vel = 0
else:
raise Exception("unknown beamline")
def enable_all(self):
self.__enable_all.put(1)
self.__enable_all.put(0)
if not self.__simulated:
self.__client = ApiClient(Configuration(host=self.__base))
self.__api = DefaultApi(self.__client)
def acknowledge_all(self):
self.__acknowledge_all.put(1)
self.__acknowledge_all.put(0)
def __make_aerotech_target(
self,
coord: AerotechCoordinate,
wait: bool = False,
incremental: bool = False,
) -> Target:
at_mm = coord.at_mm
def _disable_all(self):
self.__disable_all.put(1)
self.__disable_all.put(0)
return Target(
x=at_mm.x if at_mm is not None else None,
y=at_mm.y if at_mm is not None else None,
z=at_mm.z if at_mm is not None else None,
u=coord.omega_deg,
run_async=not wait,
incremental=incremental,
)
def stop_all(self):
self.__stop_all.put(1)
self.__stop_all.put(0)
def __make_aerotech_coordinate(self, target: Target) -> AerotechCoordinate:
return AerotechCoordinate(at_mm=Coordinate(x=target.x,y=target.y,z=target.z), omega_deg=target.u)
def __task_stop(self):
self.__task_run_enum.put(0)
def __task_start(self):
self.__task_run_enum.put(1)
def __task_run(self):
self.__task_run_enum.put(2)
def __task_pause(self):
self.__task_run_enum.put(4)
def __task_load(self):
self.__task_run_enum.put(3)
def __task_reset(self):
self.__task_run_enum.put(5)
def __set_task_id(self, task_id: int):
if task_id not in range(1, 5):
raise ValueError(f"Invalid task id {task_id}")
self.__task_id.put(task_id)
def __put(self, axis: MyMotor, attr: str, value: float):
def cancel(self):
try:
axis.put(attr=attr, value=value)
self.__api.cancel_post()
except Exception as e:
raise ValueError(f"Error setting {axis.name} {attr} to {value}: {e}")
raise AerotechCommunicationError(
"Aerotech cancel failed",
endpoint="cancel_post",
base_url=self.__base,
operation="POST",
) from e
def set_offset(self, axis: MyMotor, offset: float):
self.__put(axis=axis, attr="OFF", value=offset)
def home_all(self, task_id: int = 3):
#self.__task_id.put(2)
self.__task_reset()
time.sleep(0.2)
self.__task_id.put(task_id)
print(self.__task_id.get())
time.sleep(0.2)
self.__task_filename.put("home_all.a1exe")
print(self.__task_filename.get())
# self.__task_load()
#time.sleep(0.2)
#print(self.__task_run_enum.get())
time.sleep(1.0)
self.__task_run()
print(self.__task_run_enum.get())
def get_tast_status(self, task_id: int = 3):
task_status = PV(f"{self.aerotech_pv_prefix}:TASK:T{task_id}:STATUS")
return task_status.get(as_string=True)
def __set_global(self, index:int, value: Union[int, float, str],
timeout:float=10.0):
"""Set a global variable in aerotech, it takes ~200 ms for the value to be set
:param index: select a variable to change. an integer between 0..256 for int and real variable for 0..31 for strings
:param value: the value to set can be int, float (REAL) or string
:param timeout: timeout for vairable change, default 10.0 seconds """
var_type = self.__get_var_type(value)
caput(f"{self.aerotech_pv_prefix}:VAR:{var_type.upper()}-ADDR", index)
if var_type == 'STRING':
var_type = 'STRING-SHORT'
caput(f"{self.aerotech_pv_prefix}:VAR:{var_type.upper()}", value)
start = time.perf_counter()
while time.perf_counter() - start < timeout:
rbv = self.__read_global_from_index(index, var_type)
if rbv == str(value):
return
time.sleep(0.1)
raise TimeoutError(f"Timeout setting global variable {index} to {value}")
def __read_global_feedback(self, index, var_type: str):
caput(f"{self.aerotech_pv_prefix}:VAR:{var_type.upper()}-RBV.PROC", 1)
time.sleep(0.5)
return caget(f"{self.aerotech_pv_prefix}:VAR:{var_type.upper()}-RBV", as_string=True)
def __read_global_from_index(self, index, var_type: str):
return caget(f"{self.aerotech_pv_prefix}:VAR:{var_type.upper()}{index}_RBV", as_string=True)
def __get_var_type(self, value):
if type(value) is int:
return VariableTypeEnum.INT.name
elif type(value) is float:
return VariableTypeEnum.REAL.name
elif type(value) is str:
return VariableTypeEnum.STRING.name
else:
raise ValueError(f"Invalid type {type(value)} for global variable")
def get_global_variable(self, index: int, var_type:VariableTypeEnum):
return self.__read_global_from_index(index, var_type.name)
def set_global_variable(self, index: int, value: Union[int, float, str]):
self.__set_global(index, value)
class AerotechController:
def __init__(self, controller_ip: str):
if controller_ip is None:
self.controller = None
else:
self.controller = a1.Controller.connect(controller_ip)
self.status_item_configuration = a1.StatusItemConfiguration()
self.start_controller()
def start_controller(self):
self.controller.start()
def disconnect(self):
self.controller.disconnect()
def enable_motion(self, axis: str):
self.controller.runtime.commands.motion.enable(axis.upper())
def home_motor(self, axis: str):
self.__configure_axis_status(axis, a1.AxisStatusItem.AxisEnabled)
result = self.get_status_via_status_items(name=axis, status_item=a1.AxisStatusItem.AxisEnabled)
if not result == a1.AxisStatusItem.AxisEnabled:
self.enable_motion(axis)
self.controller.runtime.commands.motion.home(axis.upper())
def __configure_axis_status(self, axis_name: str, axis_status_item: a1.AxisStatusItem):
self.status_item_configuration.axis.add(axis_status_item=axis_status_item, axis=axis_name)
def __configure_task_status(self, task_id:int, task_status_item: a1.TaskStatusItem):
self.status_item_configuration.task.add(task_status_item=task_status_item, task=f"Task {task_id}")
def get_status_via_status_items(self, name: str | int, status_item: a1.TaskStatusItem | a1.AxisStatusItem):
result = self.controller.runtime.status.get_status_items(self.status_item_configuration)
if int(name):
return result.task.get(status_item, f"Task {name}").value
elif str(name):
return result.axis.get(status_item, name).value
else:
raise ValueError(f"Invalid status item {status_item} for task {name}")
def set_global_variable(self, index:int, value: Union[int, float, str]):
if type(value) is int:
self.controller.runtime.variables.global_.set_integer(index, value)
elif type(value) is float:
self.controller.runtime.variables.global_.set_real(index, value)
elif type(value) is str:
self.controller.runtime.variables.global_.set_string(index, value)
else:
raise ValueError(f"Invalid type {type(value)} for global variable")
def get_axis_status_via_status_items(self, axis_name:str, status_item: a1.AxisStatusItem):
result = self.controller.runtime.status.get_status_items(self.status_item_configuration)
return result.task.get(status_item, axis_name).value
def wait_for_status_to_change(self, name: str | int, status_item:a1.TaskStatusItem | a1.AxisStatusItem, enum, timeout: float = 60.0):
start = time.perf_counter()
while self.get_status_via_status_items(name, status_item) == enum:
time.sleep(0.1)
if time.perf_counter() - start > timeout:
print(self.get_status_via_status_items(name, status_item))
raise TimeoutError(f"Timeout waiting for task {name} to finish")
return
def wait_program_finish(self, task_id:int=3, timeout:float=60.0):
start = time.perf_counter()
status_item = a1.TaskStatusItem.TaskState
while True:
status = self.get_status_via_status_items(task_id, status_item)
if status != a1.TaskState.ProgramRunning:
if status == a1.TaskState.Idle:
return
elif status == a1.TaskState.ProgramComplete:
print(f"Program completed on Task {task_id}")
return
elif status == a1.TaskState.Error:
raise RuntimeError(f"Task {task_id} failed to start")
elif status == a1.TaskState.ProgramPaused:
print(f"Program paused on Task {task_id}, not sure how")
else:
raise RuntimeError(f"Unknown status {status} for task {task_id}")
if time.perf_counter() - start > timeout:
print(self.get_status_via_status_items(task_id, status_item))
raise TimeoutError(f"Timeout waiting for task {task_id} to finish")
time.sleep(0.05)
print(f"Task {task_id} finished with status {status}")
def __run_program(self, script_name:str, task_id:int=3, timeout:float=60.0):
self.__configure_task_status(task_id, a1.TaskStatusItem.TaskState)
state = self.get_status_via_status_items(task_id, a1.TaskStatusItem.TaskState)
print(f"Task {task_id} is in state {state}: {a1.TaskState(state).name}")
if state != a1.TaskState.ProgramComplete and state != a1.TaskState.Idle:
if a1.TaskState.ProgramRunning == state:
self.wait_for_status_to_change(task_id, a1.TaskStatusItem.TaskState, state, timeout)
elif a1.TaskState.ProgramComplete == state:
print('can continue')
else:
print(f"Task {task_id} is not Idle: {a1.TaskState(state).name} , aborting")
return
def is_idle(self) -> bool:
try:
print(f"Running script {script_name} on task {task_id}")
self.controller.runtime.tasks[task_id].program.run(script_name)
self.wait_program_finish(task_id, timeout)
status = self.__api.status_get()
return status.state == 'Idle'
except Exception as e:
print(f"Error executing script {script_name}: {e}")
raise AerotechCommunicationError(
"Aerotech status check failed",
endpoint="status_get",
base_url=self.__base,
operation="GET",
) from e
def run_grid_scan(self, cell_height_mm:float, num_rows:int,
row_width_mm:float, time_per_row_s: float, task_id:int =3):
self.set_global_variable(0, cell_height_mm)
self.set_global_variable(1, row_width_mm)
self.set_global_variable(2, time_per_row_s)
self.set_global_variable(1, num_rows)
self.set_global_variable(0, 1)
#self.__run_program(script_name="grid_scan.a1exe", task_id=task_id, timeout=120.0)
def home_all(self, task_id:int = 3):
self.__run_program(script_name="home_all.a1exe", task_id=task_id, timeout=120.0)
def get_position(self) -> AerotechCoordinate:
status = self.status()
return AerotechCoordinate(
at_mm=Coordinate(
x=status.x.pos,
y=status.y.pos,
z=status.z.pos,
),
omega_deg=status.u.pos,
)
def rotation_scan(self, task_id = 3, start_angle:float = 0.0, end_angle:float = 360.0, step_size:float = 1.0, num_steps:int = 10):
self.__run_program(script_name="rotation_scan.a1exe", task_id=task_id, timeout=90.0)
def status(self) -> Status:
if self.__simulated:
return Status(state=Status.State.IDLE,
x=AxisStatus(pos=self.__pos.x, vel=self.__vel,
enabled=False, homed=False, moving=False,fault=False),
y=AxisStatus(pos=self.__pos.y, vel=self.__vel,
enabled=False, homed=False, moving=False,fault=False),
z=AxisStatus(pos=self.__pos.z, vel=self.__vel,
enabled=False, homed=False, moving=False,fault=False),
u=AxisStatus(pos=self.__pos.u, vel=self.__vel,
enabled=False, homed=False, moving=False,fault=False),
)
try:
return self.__api.status_get()
except Exception as e:
raise AerotechCommunicationError(
"Aerotech status request failed",
endpoint="status_get",
base_url=self.__base,
operation="GET",
) from e
def move_motor_absolute(self, axis:str, position:float, speed:float=1.0):
self.controller.runtime.commands.motion.moveabsolute(axis.upper(), [position], [speed])
def move_home(self, wait:bool=True, incremental:bool=False):
if self.__simulated:
self.__pos = AEROTECH_HOME
return self.__pos
return self.position(AEROTECH_HOME, wait=wait, incremental=incremental)
def home_aerotech(self):
try:
return self.__api.home_post()
except Exception as e:
raise AerotechCommunicationError(
"Aerotech home failed",
endpoint="home_post",
base_url=self.__base,
operation="POST",
) from e
def wait_till_done(self, timeout=60):
try:
return self.__api.wait_till_done_post(timeout=timeout)
except Exception as e:
raise AerotechCommunicationError(
"Aerotech wait_till_done failed",
endpoint="wait_till_done_post",
base_url=self.__base,
operation="POST",
) from e
def position(
self,
target: AerotechCoordinate,
/,
wait: bool = True,
incremental: bool = False,
):
if self.__simulated:
return self.__pos
payload = self.__make_aerotech_target(target, wait=wait, incremental=incremental)
try:
return self.__api.position_post(payload)
except Exception as e:
raise AerotechCommunicationError(
"Aerotech position move failed",
endpoint="position_post",
base_url=self.__base,
operation="POST",
) from e
def rotation_scan(self, rotation_deg: float | int,
time_sec: float | int,
start_pos_deg: float | int,
run_async: bool = False):
payload = RotationRequest(
rotation_deg=rotation_deg,
time_sec=time_sec,
start_pos_deg=start_pos_deg,
run_async=run_async
)
if self.__simulated:
return payload
try:
return self.__api.rotation_scan_post(payload)
except Exception as e:
raise AerotechCommunicationError(
"Aerotech rotation scan failed",
endpoint="rotation_scan_post",
base_url=self.__base,
operation="POST",
) from e
def grid_scan(self,
grid_elem_count_y: int,
grid_elem_size_y_um: int | float,
time_sec: int|float,
grid_elem_size_x_um: Optional[Union[float, int]] = None,
grid_elem_count_x: Optional[int] = None,
run_async: Optional[bool] = False
):
payload = GridRequest(
grid_elem_count_x=grid_elem_count_x,
grid_elem_count_y=grid_elem_count_y,
grid_elem_size_x_um=grid_elem_size_x_um,
grid_elem_size_y_um=grid_elem_size_y_um,
time_sec=time_sec,
run_async=run_async,
)
if self.__simulated:
return payload
try:
return self.__api.grid_scan_post(payload)
except Exception as e:
raise AerotechCommunicationError(
"Aerotech grid scan failed",
endpoint="grid_scan_post",
base_url=self.__base,
operation="POST",
) from e
def screening_scan(self,
rotation_deg: float | int,
wedge_deg: float | int,
time_sec: float | int,
steps: int,
run_async: bool = False
):
payload = ScreenRequest(
rotation_deg=rotation_deg,
wedge_deg=wedge_deg,
time_sec=time_sec,
steps=steps,
run_async=run_async
)
if self.__simulated:
return payload
try:
return self.__api.screening_post(payload)
except Exception as e:
raise AerotechCommunicationError(
"Aerotech screening scan failed",
endpoint="screening_post",
base_url=self.__base,
operation="POST",
) from e
def move_motor_linear(self, axis: str, position: float, speed: float = 1.0):
self.controller.runtime.commands.motion.movelinear(axis.upper(), [position], speed)
if __name__ == "__main__":
beamline = mx_beamline()
print(beamline)
### test on 10S
aerotech = AerotechController(controller_ip="129.129.118.96")
#aerotech.enable_motion("X")
#aerotech.home_all()
rw = 0.320
ch = 0.010
nr = 10
tpr_s = rw / ch * 0.02
#aerotech.run_grid_scan(cell_height_mm=ch, num_rows=nr,
# row_width_mm=rw, time_per_row_s=tpr_s)
st = time.perf_counter()
aerotech.move_motor_absolute("Z", 0, 10000)
print(f"time to move: {time.perf_counter() - st}")
aerotech.disconnect()
#aerotech_epics = AerotechControllerEpics(beamline)
#aerotech_epics.omega.speed = 80.0
#print(aerotech_epics.get_global_variable(0, VariableTypeEnum.REAL))
#aerotech_epics.set_global_variable(0, 100.0)
#print(aerotech_epics.get_global_variable(0, VariableTypeEnum.REAL))
# print('enable all motors')
# aerotech_epics.enable_all()
# print('home_all')
# aerotech_epics.home_all()
# print('wait for home to finish')
# start = time.perf_counter()
# status = aerotech_epics.get_tast_status(2)
# print(f'current status: {status}')
# if status == 'Idle':
# time.sleep(0.2)
# test_counter = 0
# while status != 'Ready':
# time.sleep(0.1)
# if time.perf_counter() - start > 360.0:
# raise TimeoutError("Timeout waiting for home all task to finish")
# elif status == 'Idle':
# time.sleep(0.5)
# if test_counter == 1:
# raise RuntimeError("Home all task failed to start")
# time.sleep(1.0)
# print('restarting home all')
# aerotech_epics.home_all()
# time.sleep(1.0)
# test_counter = 1
#
# status = arotech_epics.get_tast_status(2)
# for i in range(10):
# print(f'moving to {(i+1)*90}')
# aerotech_epics.omega.move(90.0, relative=True, wait=True)
# time.sleep(1.0)
# if i == 5:
# print('simulating disable')
# aerotech_epics._disable_all()
# time.sleep(10.0)
# print('re-enabling')
# aerotech_epics.enable_all()
# time.sleep(1.0)
# print('home all')
# aerotech_epics.home_all()
# print('wait for home to finish')
# start = time.perf_counter()
# status = aerotech_epics.get_tast_status(2)
# print(f'current status: {status}')
# test_counter = 0
# while status != 'Ready':
# time.sleep(0.1)
# if time.perf_counter() - start > 360.0:
# raise TimeoutError("Timeout waiting for home all task to finish")
# status = aerotech_epics.get_tast_status(2)
# time.sleep(0.5)
# if test_counter == 1:
# raise RuntimeError("Home all task failed to start")
# time.sleep(1.0)
# print('restarting home all')
# aerotech_epics.home_all()
# time.sleep(1.0)
# test_counter = 1
#aerotech_epics.home_all()
#print(beamline)
controller = AerotechController(beamline)
#controller.print_status(colored=True, compact=True)
print(controller.get_position())
controller.cancel()
+187
View File
@@ -0,0 +1,187 @@
import os
import sys
import threading
import time
from enum import Enum
from typing import Union
import automation1 as a1
from aare.common.beamline import MXBeamline, mx_beamline
class AerotechController:
def __init__(self, controller_ip: str):
# if controller_ip is None:
# self.controller = None
# else:
self.controller = a1.Controller.connect(controller_ip)
self.status_item_configuration = a1.StatusItemConfiguration()
self.start_controller()
def start_controller(self):
self.controller.start()
def disconnect(self):
self.controller.disconnect()
def enable_motion(self, axis: str):
self.controller.runtime.commands.motion.enable(axis.upper())
def home_motor(self, axis: str):
self.__configure_axis_status(axis, a1.AxisStatusItem.AxisEnabled)
result = self.get_status_via_status_items(name=axis, status_item=a1.AxisStatusItem.AxisEnabled)
if not result == a1.AxisStatusItem.AxisEnabled:
self.enable_motion(axis)
self.controller.runtime.commands.motion.home(axis.upper())
def __configure_axis_status(self, axis_name: str, axis_status_item: a1.AxisStatusItem):
self.status_item_configuration.axis.add(axis_status_item=axis_status_item, axis=axis_name)
def __configure_task_status(self, task_id:int, task_status_item: a1.TaskStatusItem):
self.status_item_configuration.task.add(task_status_item=task_status_item, task=f"Task {task_id}")
def get_status_via_status_items(self, name: str | int, status_item: a1.TaskStatusItem | a1.AxisStatusItem):
result = self.controller.runtime.status.get_status_items(self.status_item_configuration)
if int(name):
return result.task.get(status_item, f"Task {name}").value
elif str(name):
return result.axis.get(status_item, name).value
else:
raise ValueError(f"Invalid status item {status_item} for task {name}")
def set_global_variable(self, index:int, value: Union[int, float, str]):
if type(value) is int:
self.controller.runtime.variables.global_.set_integer(index, value)
elif type(value) is float:
self.controller.runtime.variables.global_.set_real(index, value)
elif type(value) is str:
self.controller.runtime.variables.global_.set_string(index, value)
else:
raise ValueError(f"Invalid type {type(value)} for global variable")
def get_axis_status_via_status_items(self, axis_name:str, status_item: a1.AxisStatusItem):
result = self.controller.runtime.status.get_status_items(self.status_item_configuration)
return result.task.get(status_item, axis_name).value
def wait_for_status_to_change(self, name: str | int, status_item:a1.TaskStatusItem | a1.AxisStatusItem, enum, timeout: float = 60.0):
start = time.perf_counter()
while self.get_status_via_status_items(name, status_item) == enum:
time.sleep(0.1)
if time.perf_counter() - start > timeout:
print(self.get_status_via_status_items(name, status_item))
raise TimeoutError(f"Timeout waiting for task {name} to finish")
return
def wait_program_finish(self, task_id:int=3, timeout:float=60.0):
start = time.perf_counter()
status_item = a1.TaskStatusItem.TaskState
while True:
status = self.get_status_via_status_items(task_id, status_item)
if status != a1.TaskState.ProgramRunning:
if status == a1.TaskState.Idle:
return
elif status == a1.TaskState.ProgramComplete:
print(f"Program completed on Task {task_id}")
return
elif status == a1.TaskState.Error:
raise RuntimeError(f"Task {task_id} failed to start")
elif status == a1.TaskState.ProgramPaused:
print(f"Program paused on Task {task_id}, not sure how")
else:
raise RuntimeError(f"Unknown status {status} for task {task_id}")
if time.perf_counter() - start > timeout:
print(self.get_status_via_status_items(task_id, status_item))
raise TimeoutError(f"Timeout waiting for task {task_id} to finish")
time.sleep(0.05)
print(f"Task {task_id} finished with status {status}")
def __run_program(self, script_name:str, task_id:int=3, timeout:float=60.0):
self.__configure_task_status(task_id, a1.TaskStatusItem.TaskState)
state = self.get_status_via_status_items(task_id, a1.TaskStatusItem.TaskState)
print(f"Task {task_id} is in state {state}: {a1.TaskState(state).name}")
if state != a1.TaskState.ProgramComplete and state != a1.TaskState.Idle:
if a1.TaskState.ProgramRunning == state:
self.wait_for_status_to_change(task_id, a1.TaskStatusItem.TaskState, state, timeout)
elif a1.TaskState.ProgramComplete == state:
print('can continue')
else:
print(f"Task {task_id} is not Idle: {a1.TaskState(state).name} , aborting")
return
try:
print(f"Running script {script_name} on task {task_id}")
self.controller.runtime.tasks[task_id].program.run(script_name)
self.wait_program_finish(task_id, timeout)
except Exception as e:
print(f"Error executing script {script_name}: {e}")
def run_grid_scan(self, cell_height_mm:float, num_rows:int,
row_width_mm:float, time_per_row_s: float, task_id:int =3):
self.set_global_variable(0, cell_height_mm)
self.set_global_variable(1, row_width_mm)
self.set_global_variable(2, time_per_row_s)
self.set_global_variable(1, num_rows)
self.set_global_variable(0, 1)
#self.__run_program(script_name="grid_scan.a1exe", task_id=task_id, timeout=120.0)
def home_all(self, task_id:int = 3):
self.__run_program(script_name="home_all.a1exe", task_id=task_id, timeout=120.0)
def rotation_scan(self, task_id = 3, start_angle:float = 0.0, end_angle:float = 360.0, step_size:float = 1.0, num_steps:int = 10):
self.__run_program(script_name="rotation_scan.a1exe", task_id=task_id, timeout=90.0)
def move_motor_absolute(self, axis:str, position:float, speed:float=1.0):
self.controller.runtime.commands.motion.moveabsolute(axis.upper(), [position], [speed])
def move_motor_linear(self, axis: str, position: float, speed: float = 1.0):
self.controller.runtime.commands.motion.movelinear(axis.upper(), [position], speed)
if __name__ == "__main__":
beamline = mx_beamline()
print(beamline)
### test on 10S
aerotech = AerotechController(controller_ip="129.129.118.96")
#aerotech.enable_motion("X")
#aerotech.home_all()
rw = 0.320
ch = 0.010
nr = 10
tpr_s = rw / ch * 0.02
#aerotech.run_grid_scan(cell_height_mm=ch, num_rows=nr,
# row_width_mm=rw, time_per_row_s=tpr_s)
# st = time.perf_counter()
# aerotech.move_motor_absolute("Z", 0, 10000)
# print(f"time to move: {time.perf_counter() - st}")
# aerotech.disconnect()
controller = a1.Controller.connect("129.129.118.96")
status_item_configuration = a1.StatusItemConfiguration()
controller.start()
axis = "X"
pso_input = a1.PsoWindowInput.iXC4ePrimaryFeedback
window_number = 0
reverse_direction = False
execution_task_index = 3
min_x_mm = 0
max_x_mm = 1
counts_per_unit = controller.runtime.parameters.axes[axis].units.countsperunit.value
print(f"Counts per unit for {axis}: {counts_per_unit}")
units_to_counts = controller.runtime.commands.utility_and_conversion.unitstocounts(axis,5,execution_task_index)
print(f"Units to counts for {axis}: {units_to_counts}")
print(dir(controller.runtime.commands))
primary_emulated_quadrature_divider = controller.runtime.parameters.axes[axis].feedback.primaryemulatedquadraturedivider.value
pso_lower_bound = round(controller.runtime.commands.utility_and_conversion.unitstocounts(axis,min_x_mm,execution_task_index)/primary_emulated_quadrature_divider)
pso_upper_bound = round(controller.runtime.commands.utility_and_conversion.unitstocounts(axis,max_x_mm,execution_task_index)/primary_emulated_quadrature_divider)
controller.runtime.commands.pso.psoreset(axis)
controller.runtime.commands.pso.psowindowconfigureinput(axis,0, pso_input, True, execution_task_index)
controller.runtime.commands.pso.psowindowconfigurefixedrange(axis,window_number,pso_lower_bound,pso_upper_bound,execution_task_index)
controller.runtime.commands.pso.psowindowoutputon(axis,window_number, execution_task_index)
controller.runtime.commands.motion.movelinear(axis,[1],0.1,execution_task_index)
controller.runtime.commands.motion.waitformotiondone(axis,execution_task_index)
controller.runtime.commands.motion.movelinear(axis,[-1],0.1,execution_task_index)
controller.runtime.commands.motion.waitformotiondone(axis,execution_task_index)
controller.runtime.commands.pso.psowindowoutputoff(axis,window_number, execution_task_index)
controller.runtime.commands.pso.psoreset(axis,execution_task_index)
+40 -5
View File
@@ -3,6 +3,7 @@ import math
import jfjoch_client
from aare.common.beamline import MXBeamline
from aare.common.exception_handler import JFJochCommunicationError
from aare.common.models import DAQStatusModel, FluorescenceSpectrumOutputModel
from aare.common.raster_grid import RasterGridRequest
from aare.common.rotation_scan import RotationScanRequest
@@ -10,6 +11,7 @@ from aare.common.rotation_scan import RotationScanRequest
class JFJochWrapper:
def __init__(self, bl: MXBeamline):
self.__simulated = False
match bl:
case MXBeamline.X06DA:
self.__url = "http://sls-gpu-001:8080"
@@ -17,6 +19,10 @@ class JFJochWrapper:
self.__url = "http://sls-gpu-002:8080"
case MXBeamline.SIMULATED:
self.__url = "http://localhost:8080"
self.__client = None
self.__api = None
self.__simulated = True
return
case _:
raise Exception("unknown beamline")
@@ -93,7 +99,18 @@ class JFJochWrapper:
detect_ice_rings=True,
xray_fluorescence_spectrum=xrf
)
self.__api.start_post(dataset_settings=dataset_settings)
try:
self.__api.start_post(dataset_settings=dataset_settings)
except Exception as e:
scan_type = "rotation"
if r.screening:
scan_type = "screening"
raise JFJochCommunicationError(
f"JFJoch data collection failed to initialize for {scan_type} scan",
operation="POST",
endpoint="start_post",
base_url=self.__url,
) from e
def measure_raster(self,
r: RasterGridRequest,
@@ -141,11 +158,29 @@ class JFJochWrapper:
max_spot_count = 1000,
detect_ice_rings = True
)
self.__api.start_post(dataset_settings=dataset_settings)
try:
self.__api.start_post(dataset_settings=dataset_settings)
except Exception as e:
raise JFJochCommunicationError(
"JFJoch data collection failed to initialize for raster scan",
operation="POST",
endpoint="start_post",
base_url=self.__url,
) from e
def wait_till_done(self, timeout : int | float) -> jfjoch_client.models.ScanResult:
self.__api.wait_till_done_post_with_http_info(timeout=math.ceil(timeout))
return self.__api.result_scan_get()
def wait_till_done(self, timeout : int | float) -> jfjoch_client.models.ScanResult | None:
if self.__simulated:
return None
try:
self.__api.wait_till_done_post_with_http_info(timeout=math.ceil(timeout))
return self.__api.result_scan_get()
except Exception as e:
raise JFJochCommunicationError(
"JFJoch wait/result retrieval failed",
operation="GET",
endpoint="wait_till_done_post / result_scan_get",
base_url=self.__url,
) from e
def detector(self) -> jfjoch_client.models.DetectorListElement:
l = self.__api.config_select_detector_get()
+410
View File
@@ -0,0 +1,410 @@
import json
import re
import time
from typing import Any, Callable, Protocol
from urllib.parse import urlparse
import requests
from aare.common.exception_handler import TellCommunicationError
from aare.common.logger_config import setup_logger
from aare.common.beamline import MXBeamline # noqa: F401
from pshell import PShellClient
logger = setup_logger("aareDAQ")
VALID_DEWAR_POSITIONS = [f"{p}{n}" for n in "12345" for p in "ABCDEFX"]
def is_valid_dewar_position(position):
"""check if argument is a valid dewar position"""
return position in VALID_DEWAR_POSITIONS
POSITION_PARK = "pPark"
POSITION_COLD = "pCold"
POSITION_AUX = "pAux"
POSITION_DEWAR = "pDewar"
POSITION_HOME = "pHome"
POSITION_HEATER = "pHeatB"
class ManualMountException(Exception):
"""Custom exception for manual mounting"""
pass
class SmartMagnetFaultException(Exception):
"""Custom exception for smart magnet fault"""
pass
class TellMountFailedException(Exception):
"""Custom exception for mount failure"""
pass
class TellCommandWhileBusyException(Exception):
"""Custom exception for trying to move Tell when it is busy"""
pass
class TellConnectionException(Exception):
"""Custom exception for connection problems"""
pass
class TellBackend(Protocol):
@property
def url(self) -> str | None:
...
def get_state(self) -> str:
...
def get_result(self, command_id: int = -1):
...
def wait_state(self, state: str, timeout: float) -> None:
...
def wait_state_not(self, state: str, timeout: float) -> None:
...
def wait_events(self, events: dict[str, Any], timeout: float):
...
def eval(self, expr: str):
...
def start_eval(self, expr: str) -> int:
...
def run(self, path: str, pars: list[str] | None = None, background: bool = False) -> None:
...
def abort(self) -> None:
...
class PShellTellBackend:
def __init__(self, bl: MXBeamline):
self._url = self._resolve_url(bl)
print(f"Connecting TELL p-shell service at {self._url} ...", end="")
hostname = urlparse(self._url).hostname
try:
requests.get(f"{self._url}/history/0", timeout=1.0)
except requests.exceptions.RequestException as e:
print(f"...connection to {hostname} failed")
raise TellCommunicationError(
f"TELL connection failed ({hostname})",
base_url=self._url,
endpoint="history/0",
operation="GET",
) from e
except requests.ReadTimeout as e:
print(f"...PShell service {hostname} is down")
raise TellCommunicationError(
f"TELL connection timedout ({hostname})",
base_url=self._url,
endpoint="history/0",
operation="GET",
) from e
self._pshell = PShellClient(self._url)
@staticmethod
def _resolve_url(bl: MXBeamline) -> str:
beamline = bl.value.lower()
if bl == MXBeamline.X06DA:
return f"http://{beamline}-tell.psi.ch:22222"
if bl == MXBeamline.X10SA:
return "http://PC17488:22222"
if bl == MXBeamline.X06SA:
raise NotImplementedError(f"TellClient not implemented for {beamline}")
if bl == MXBeamline.SIMULATED:
raise NotImplementedError("Use SimTellBackend for MXBeamline.SIMULATED")
raise ValueError(f"Unknown beamline {beamline}")
@property
def url(self) -> str | None:
return self._url
def get_state(self) -> str:
return self._pshell.get_state()
def get_result(self, command_id: int = -1):
return self._pshell.get_result(command_id)
def wait_state(self, state: str, timeout: float) -> None:
self._pshell.wait_state(state, timeout=timeout)
def wait_state_not(self, state: str, timeout: float) -> None:
self._pshell.wait_state_not(state, timeout=timeout)
def wait_events(self, events: dict[str, Any], timeout: float):
return self._pshell.wait_events(events, timeout=timeout)
def eval(self, expr: str):
return self._pshell.eval(expr)
def start_eval(self, expr: str) -> int:
return self._pshell.start_eval(expr)
def run(self, path: str, pars: list[str] | None = None, background: bool = False) -> None:
self._pshell.run(path, pars=pars, background=background)
def abort(self) -> None:
self._pshell.abort()
class SimTellBackend:
def __init__(self):
self._url: str | None = None
self._state = "Ready"
self._last_cmd_id = 1000
self._mounted_sample = ""
self._settings: dict[str, str] = {"mounted_sample_position": ""}
self._results: dict[int, dict[str, Any]] = {}
self._robot_status: dict[str, Any] = {
"powered": True,
"pos": POSITION_PARK,
}
self._current_mA = 30.0
self._pin_offset = 0.0
self._detected_pucks: list[dict[str, Any]] = []
self._system_check_msg = "OK"
self._smart_magnet_state = "Ready"
self._in_mount_position = False
@property
def url(self) -> str | None:
return self._url
def _next_cmd_id(self) -> int:
self._last_cmd_id += 1
return self._last_cmd_id
def _set_ready_soon(self) -> None:
time.sleep(0.01)
self._state = "Ready"
def get_state(self) -> str:
return self._state
def get_result(self, command_id: int = -1):
if command_id == -1:
command_id = self._last_cmd_id
return self._results.get(command_id, {"status": "completed"})
def wait_state(self, state: str, timeout: float) -> None:
if self._state != state:
time.sleep(min(timeout, 0.05))
self._state = state
def wait_state_not(self, state: str, timeout: float) -> None:
if self._state == state:
time.sleep(min(timeout, 0.05))
self._state = "Ready"
def wait_events(self, events: dict[str, Any], timeout: float):
self.wait_state_not("Busy", timeout)
if "Motion Sync" in events:
return "Motion Sync", "Robot Clear after mount"
if "Motion Task" in events:
return "Motion Task", "idle"
return None, self._state
def eval(self, expr: str):
expr = expr.strip()
if expr == "in_mount_position&":
return "true" if self._in_mount_position else "false"
if expr.startswith("in_mount_position = "):
self._in_mount_position = "True" in expr or "true" in expr
return None
if expr.startswith("set_setting("):
match = re.match(r"set_setting\('([^']+)', '([^']*)'\)&", expr)
if match:
key, value = match.groups()
self._settings[key] = value
return None
if expr.startswith("get_setting("):
match = re.match(r"get_setting\('([^']+)'\)&", expr)
if match:
key = match.group(1)
return self._settings.get(key, "")
return ""
if expr == "system_check_msg()&":
return self._system_check_msg
if expr == "robot.state&":
return self._state
if expr == "robot.take()&":
return str(self._robot_status)
if expr == "get_pucks_info()&":
return json.dumps(self._detected_pucks)
if expr == "get_pin_offset()&":
return str(self._pin_offset)
if expr == "smart_magnet.get_current_rb()&":
return str(self._current_mA)
if expr.startswith("smart_magnet.set_current("):
match = re.match(r"smart_magnet\.set_current\(([-+]?\d+(?:\.\d+)?)\)&", expr)
if match:
self._current_mA = float(match.group(1))
return None
if expr == "enable_motion()&":
self._robot_status["powered"] = True
return None
if expr == "smart_magnet.state&":
return self._smart_magnet_state
if expr == "smart_magnet.set_supress(True)&":
return None
if expr == "smart_magnet.set_supress(False)&":
return None
if expr == "smart_magnet.set_resting_current()&":
return None
if expr == "robot.stop_task()&":
self._state = "Ready"
return None
return None
def start_eval(self, expr: str) -> int:
cmd_id = self._next_cmd_id()
self._state = "Busy"
if expr.startswith("mount("):
parts = re.findall(r"'([^']*)'|([^,()]+)", expr)
values = [a if a else b.strip() for a, b in parts]
if len(values) >= 4:
segment = values[0]
puck = values[1]
sample = values[2]
mounted = f"{segment}{puck}{sample}"
self._mounted_sample = mounted
self._settings["mounted_sample_position"] = mounted
self._robot_status["pos"] = POSITION_DEWAR
elif expr.startswith("unmount("):
self._mounted_sample = ""
self._settings["mounted_sample_position"] = ""
self._robot_status["pos"] = POSITION_PARK
elif expr.startswith("move_park("):
self._robot_status["pos"] = POSITION_PARK
elif expr.startswith("move_cold("):
self._robot_status["pos"] = POSITION_COLD
elif expr.startswith("dry("):
self._robot_status["pos"] = POSITION_HEATER
self._results[cmd_id] = {"status": "completed", "command": expr}
self._set_ready_soon()
return cmd_id
def run(self, path: str, pars: list[str] | None = None, background: bool = False) -> None:
_ = background
if path == "data/set_samples_info" and pars:
try:
data = json.loads(pars[0])
self._detected_pucks = []
for item in data:
puck_address = item.get("puckAddress", "")
if puck_address:
self._detected_pucks.append(
{
"puckState": "Present",
"puckAddress": puck_address,
"puckBarcode": item.get("puckBarcode", ""),
}
)
except Exception:
logger.warning("Failed to load simulated samples info")
def abort(self) -> None:
self._state = "Ready"
class LazyTellBackend:
def __init__(self, factory: Callable[[], TellBackend], *, retry_interval_s: float = 2.0):
self._factory = factory
self._backend: TellBackend | None = None
self._retry_interval_s = float(retry_interval_s)
self._last_attempt_ts = 0.0
self._last_error: Exception | None = None
def _get_backend(self) -> TellBackend:
if self._backend is not None:
return self._backend
now = time.monotonic()
if now - self._last_attempt_ts < self._retry_interval_s and self._last_error is not None:
raise self._last_error
self._last_attempt_ts = now
try:
self._backend = self._factory()
self._last_error = None
return self._backend
except TellCommunicationError as e:
self._last_error = e
raise
except Exception as e:
wrapped = TellCommunicationError(
"TELL connection failed",
operation="CONNECT",
)
self._last_error = wrapped
raise wrapped from e
@property
def url(self) -> str | None:
return self._get_backend().url
def get_state(self) -> str:
return self._get_backend().get_state()
def get_result(self, command_id: int = -1):
return self._get_backend().get_result(command_id)
def wait_state(self, state: str, timeout: float) -> None:
self._get_backend().wait_state(state, timeout)
def wait_state_not(self, state: str, timeout: float) -> None:
self._get_backend().wait_state_not(state, timeout)
def wait_events(self, events: dict[str, Any], timeout: float):
return self._get_backend().wait_events(events, timeout)
def eval(self, expr: str):
return self._get_backend().eval(expr)
def start_eval(self, expr: str) -> int:
return self._get_backend().start_eval(expr)
def run(self, path: str, pars: list[str] | None = None, background: bool = False) -> None:
self._get_backend().run(path, pars=pars, background=background)
def abort(self) -> None:
self._get_backend().abort()
+93 -361
View File
@@ -1,14 +1,9 @@
import ast
import json
import random
import re
import time
from typing import List
from urllib.parse import urlparse
import requests
from aare.common.exception_handler import TellCommunicationError
from aare.common.logger_config import setup_logger
from aare.common.beamline import MXBeamline
from aare.common.models import (
PuckLoadedInfo,
DewarAddress,
@@ -16,93 +11,28 @@ from aare.common.models import (
)
from aareDB import PuckWithTellPosition
from aare.common.beamline import MXBeamline # noqa: F401
from pshell import PShellClient
from aare.devices.tell_backend import (
TellBackend,
SimTellBackend,
LazyTellBackend,
PShellTellBackend,
ManualMountException,
POSITION_COLD,
SmartMagnetFaultException,
TellConnectionException,
TellMountFailedException,
TellCommandWhileBusyException,
)
from aare.common.logger_config import setup_logger
logger = setup_logger("aareDAQ")
class ManualMountException(Exception):
"""Custom exception for manual mounting"""
pass
class SmartMagnetFaultException(Exception):
"""Custom exception for smart magnet fault"""
pass
class TellMountFailedException(Exception):
"""Custom exception for mount failure"""
pass
class TellCommandWhileBusyException(Exception):
"""Custom exception for trying to move Tell when it is busy"""
pass
class TellConnectionException(Exception):
"""Custom exception for connection problems"""
pass
VALID_DEWAR_POSITIONS = [f"{p}{n}" for n in "12345" for p in "ABCDEFX"]
def is_valid_dewar_position(position):
"""check if argument is a valid dewar position"""
return position in VALID_DEWAR_POSITIONS
POSITION_PARK = "pPark"
POSITION_COLD = "pCold"
POSITION_AUX = "pAux"
POSITION_DEWAR = "pDewar"
POSITION_HOME = "pHome"
POSITION_HEATER = "pHeatB"
#Nov 26 13:36:00 mx-x06da-queue-01.psi.ch AareDAQ[2944444]: 2025-11-26 13:36:00,388 - aareDAQ - ERROR - Error getting status: ('Connection aborted.', ConnectionResetError(104, 'Connection reset by peer'))
class TellClient:
"""High-level Tell robot API using PShellClient"""
def __init__(self, bl: MXBeamline):
self.__url = None
beamline = bl.value.lower()
"""High-level Tell robot API using a pluggable backend"""
def __init__(self, bl: MXBeamline, backend: TellBackend | None = None):
self.__beamline = bl
if bl == MXBeamline.X06DA:
self.__url = f"http://{beamline}-tell.psi.ch:22222"
elif bl == MXBeamline.X10SA:
self.__url = f"http://PC17488:22222"
elif bl == MXBeamline.X06SA:
self.__url = f""
raise NotImplemented(f"TellClient not implemente for {beamline}")
elif bl == MXBeamline.SIMULATED:
raise NotImplemented(f"Use SimClient, generate tell client using"
f"make_tell_client(beamline)")
else:
raise ValueError(f"Unknown beamline {beamline}")
print(f"Connecting TELL p-shell service at {self.__url} ...", end="")
hostname = urlparse(self.__url).hostname
try:
requests.get(f"{self.__url}/history/0", timeout=1.0)
except requests.exceptions.RequestException as e:
print(f"...connection to {hostname} failed")
raise TellCommunicationError(
f"TELL connection failed ({hostname})",
base_url=self.__url,
endpoint="history/0",
operation="GET",
) from e
except requests.ReadTimeout as e:
print(f"...PShell service {hostname} is down")
raise TellCommunicationError(
f"TELL connection timedout ({hostname})",
base_url=self.__url,
endpoint="history/0",
operation="GET",
) from e
self.pshell = PShellClient(self.__url)
self.backend = backend or PShellTellBackend(bl)
self._aborted = False
self.state = self.get_state()
@@ -112,26 +42,26 @@ class TellClient:
@property
def url(self):
"""returns the configured base url for the Tell robot"""
return self.__url
return self.backend.url
def get_state(self):
"""returns the current state of the robot"""
self.state = self.pshell.get_state()
self.state = self.backend.get_state()
return self.state
def get_result(self, command_id=-1):
"""returns the result of the last command issued to the robot"""
return self.pshell.get_result(command_id)
return self.backend.get_result(command_id)
def wait_ready(self, timeout: float = 360.0):
"""waits until the robot is ready to accept commands returns None if simulation
and raises an exception if the robot is not ready"""
self.pshell.wait_state("Ready", timeout=timeout)
self.backend.wait_state("Ready", timeout=timeout)
def wait_not_busy(self, timeout: float = 360.0):
"""waits until the robot is not busy and returns None if simulation
and raises an exception if the robot is busy"""
self.pshell.wait_state_not("Busy", timeout=timeout)
self.backend.wait_state_not("Busy", timeout=timeout)
state = self.get_state()
if state != "Ready":
if state == "Initializing":
@@ -143,11 +73,11 @@ class TellClient:
def set_in_mount_position(self, value):
"""tells the robot that the beamlien is safe and to set the in mount position flag allowing mounting
:param value """
self.pshell.eval("in_mount_position = " + str(value) + "&")
self.backend.eval("in_mount_position = " + str(value) + "&")
def is_in_mount_position(self) -> bool:
"""checks to see if the robot is in the mount position and returns a boolean"""
return self.pshell.eval("in_mount_position&").lower() == "true"
return self.backend.eval("in_mount_position&").lower() == "true"
def set_samples_info(self, info: List[PuckWithTellPosition]):
"""sets the samples in the robot dewar based on the given list of PuckWithTellPosition objects
@@ -160,7 +90,7 @@ class TellClient:
"userName": x.pgroup,
"dewarName": x.dewar_name or "",
"puckName": x.puck_name,
"puckType": "Unipuck", # could use x.puck_type
"puckType": "Unipuck",
"puckAddress": x.tell_position or "",
"puckBarcode": x.puck_name,
"sampleBarcode": "",
@@ -171,8 +101,7 @@ class TellClient:
}
)
self.pshell.run("data/set_samples_info", pars=[json.dumps(j)], background=True)
# self.pshell.eval("set_samples_info(" + json.dumps(info) + ")&")
self.backend.run("data/set_samples_info", pars=[json.dumps(j)], background=True)
def start_cmd(self, cmd, *argv):
"""starts a command on the robot and returns the command id"""
@@ -180,7 +109,7 @@ class TellClient:
for a in argv:
cmd = cmd + (("'" + a + "'") if type(a) is str else str(a)) + ", "
cmd = cmd + ")"
ret = self.pshell.start_eval(cmd)
ret = self.backend.start_eval(cmd)
self.get_state()
return ret
@@ -191,11 +120,10 @@ class TellClient:
result = self.get_result(self._last_cmd_id)
logger.debug(f"{msg} {result}")
status = result["status"]
if "completed" != status: #FIXME this is very limiting and depends on tell reporting statuses
if "completed" != status:
if "removed" != status:
raise TellMountFailedException(f"{msg} {result}")
else:
return f"{msg} {result}"
return f"{msg} {result}"
def estimate_mounting_time(self, segment) -> int:
"""Adds additional time if cooling/drying is expected based on requested segment,
@@ -206,7 +134,7 @@ class TellClient:
gripper_in_cold = self.is_in_cold()
if current_mounted is None:
unmount_needs_drying = 0 # might not have anything
unmount_needs_drying = 0
unmount_needs_cooling = 0
else:
segment_in_cold = current_mounted.puck.segment in "ABCDEF"
@@ -219,27 +147,18 @@ class TellClient:
needs_cooling = mount_needs_cooling + unmount_needs_cooling
needs_drying = mount_needs_drying + unmount_needs_drying
return needs_cooling * 30 + needs_drying * 120
except:
except Exception:
return 0
def mount(
self,
address: SampleDewarAddress,
force: bool = False, # kept for future
read_dm: bool = False, # read data matrix
auto_unmount: bool = False, # single command, if False it will raise exception
wait: bool = False, # blocking operation
force: bool = False,
read_dm: bool = False,
auto_unmount: bool = False,
wait: bool = False,
timeout: float = 600.0,
):
"""send api request to mount sample from dewer after validating dewer address returns None or repsonse.
If the robot is busy, mount will raise an exception.
:param address: SampleDewarAddress
:param force: bool
:param read_dm: bool
:param auto_unmount: bool
:param wait: bool
:param timeout: float
"""
SampleDewarAddress.model_validate(address)
segment = address.puck.segment
@@ -258,9 +177,16 @@ class TellClient:
wait_timeout = timeout + self.estimate_mounting_time(segment)
logger.info("waiting for mount to complete")
if wait and segment in "ABCDEF":
event, value = self.pshell.wait_events({"state": None, "Motion Task": "dry",
"Gripper detection" : None,
"Motion Sync": "Robot Clear after mount"}, timeout=wait_timeout)
event, value = self.backend.wait_events(
{
"state": None,
"Motion Task": "dry",
"Gripper detection": None,
"Motion Sync": "Robot Clear after mount",
},
timeout=wait_timeout,
)
logger.info(f"event: {event} occurred with value: {value}")
if event is None or event == "state":
logger.info(f"event: {event} occurred with value: {value}, checking command completed okay")
self.check_command_ok(
@@ -299,11 +225,6 @@ class TellClient:
return None
def unmount(self, force=False, wait=False, timeout=360.0):
"""send api request to unmount sample from dewer returns None or repsonse.
:param force: bool Force has a meaning, will unmount even if smart magnet is not detecting sample
:param wait: bool If true will wait until unmount is completed
:timeout: float"""
if self.is_busy():
raise TellCommandWhileBusyException("mount received while robot is busy")
@@ -315,82 +236,60 @@ class TellClient:
return self._last_cmd_id
def dry(self, heat_time=None, speed=None, wait_cold=None, wait=False):
"""send api request to dry tell gripper.
:param: heat_time float if None Tell will use default for drying time
:param: speed float if None Tell will use default for drying speed
:param: wait_cold bool if -1 to go to park after dry. if None Tell will use default time to wait_cold.
:param wait: bool If true will wait until drying is completed
"""
self.pshell.wait_state("Ready", timeout=30.0)
self.backend.wait_state("Ready", timeout=30.0)
self._last_cmd_id = self.start_cmd("dry", heat_time, speed, wait_cold)
if wait:
self.check_command_ok(timeout=360.0, msg=f"Dry failed")
self.check_command_ok(timeout=360.0, msg="Dry failed")
def move_park(self, wait=False):
"""send api request to move robot to park position"""
self._last_cmd_id = self.start_cmd("move_park")
if wait:
self.check_command_ok(timeout=360.0, msg=f"Move to park failed")
self.check_command_ok(timeout=360.0, msg="Move to park failed")
def move_cold(self, reset_timestamp=False, wait=False):
"""send api request to move robot to cold position"""
self._last_cmd_id = self.start_cmd("move_cold", reset_timestamp)
if wait:
self.check_command_ok(timeout=360.0, msg=f"Move to cold failed")
self.check_command_ok(timeout=360.0, msg="Move to cold failed")
def abort_cmd(self):
"""sends an abort pshell requesst and a robot stop task command"""
self.pshell.abort()
self.pshell.eval("robot.stop_task()&")
self.backend.abort()
self.backend.eval("robot.stop_task()&")
def set_setting(self, key: str, value: str):
"""wrapper for pshell eval set_setting command
:param key str, name of a setting in tell
:param value str, the new value of the setting as a string"""
self.pshell.eval(f"set_setting('{key}', '{value}')&")
self.backend.eval(f"set_setting('{key}', '{value}')&")
def get_setting(self, key: str) -> str:
"""wrapper for pshell eval get_setting command, returns the current value for key as a string
:param key str, name of a setting in tell"""
return self.pshell.eval(f"get_setting('{key}')&")
return self.backend.eval(f"get_setting('{key}')&")
def get_mounted_sample(self) -> SampleDewarAddress | None:
"""get the current mounted sample and return a SampleDewarAddress object or None if no sample is mounted"""
ret = self.get_setting('mounted_sample_position').strip()
if not ret or len(ret) == 0:
ret = self.get_setting("mounted_sample_position").strip()
if not ret:
return None
match = re.match(r"([A-Z])(\d)(\d{1,2})", ret)
if match:
segment, puck, sample = match.groups()
dewar_location = DewarAddress(segment=segment, pos=int(puck))
return SampleDewarAddress(puck=dewar_location, pin=int(sample))
else:
logger.warning(f"Failed to decode mounted sample position: {ret}")
return None
logger.warning(f"Failed to decode mounted sample position: {ret}")
return None
def get_system_check(self):
"""returns the current system check status"""
return self.pshell.eval("system_check_msg()&")
return self.backend.eval("system_check_msg()&")
def get_robot_state(self):
"""returns the current robot state"""
return self.pshell.eval("robot.state&")
return self.backend.eval("robot.state&")
def get_robot_status(self):
"""returns the current robot status"""
status = self.pshell.eval("robot.take()&")
return eval(status) # FIXME ALL functions must return a valid JSON object
status = self.backend.eval("robot.take()&")
#return eval(status)
return ast.literal_eval(status)
def get_detected_pucks(self) -> List[PuckLoadedInfo]:
"""returns a list of detected pucks as PuckLoadedInfo objects"""
j = json.loads(self.pshell.eval("get_pucks_info()&"))
j = json.loads(self.backend.eval("get_pucks_info()&"))
output = []
for i in j:
if i["puckState"] == "Present":
puck_address = i["puckAddress"]
@@ -399,90 +298,77 @@ class TellClient:
PuckLoadedInfo(
puck_name=i["puckBarcode"],
location=DewarAddress(
segment=puck_address[0], pos=int(puck_address[1])
segment=puck_address[0],
pos=int(puck_address[1]),
),
),
)
return output
def get_pin_offset(self):
"""get the pin offset for the smart magnet, returns offset as a float"""
try:
offset = float(self.pshell.eval("get_pin_offset()&"))
offset = float(self.backend.eval("get_pin_offset()&"))
except Exception:
offset = 0.0
return offset
def get_current(self):
"""get the current drawn by the smart magnet, returns current as a float in mA"""
current = self.pshell.eval("smart_magnet.get_current_rb()&")
current = self.backend.eval("smart_magnet.get_current_rb()&")
return float(current)
def set_current(self, current: float) -> float:
"""set the current drawn by the smart magnet, returns current as a float in mA"""
self.pshell.eval("smart_magnet.set_current({:.1f})&".format(current))
current = self.pshell.eval("smart_magnet.get_current_rb()&")
self.backend.eval("smart_magnet.set_current({:.1f})&".format(current))
current = self.backend.eval("smart_magnet.get_current_rb()&")
return float(current)
def is_powered(self):
"""returns True if the robot is powered on"""
return self.get_robot_status()["powered"]
def check_enable_motion(self):
"""check if the robot is powered on and enable motion if not"""
if not self.is_powered():
self.pshell.eval("enable_motion()&")
self.backend.eval("enable_motion()&")
def is_in_cold(self):
"""Compare current robot position to the set cold position. Returns True if in cold position, False otherwise."""
return self.is_position(POSITION_COLD)
def is_position(self, position: str) -> bool:
"""Compare current robot position to a given position. Returns True if in position, False otherwise."""
return position == self.get_robot_status()["pos"]
def is_ready(self):
"""returns True if the robot is ready to receive commands"""
return "ready" == self.get_state().lower()
def is_busy(self):
"""returns True if the robot is busy"""
return "busy" == self.get_state().lower()
def check_smart_magnet_mounted(self, timeout: float = 10.0, idle_time: float = 1.0, interval: float = 0.1):
"""Reads smart_magent state and tries to infer if a sample is present
Handles: PAUSED, Fault, Busy and Ready states.
Raises a ManualMountException is the amgnet indicates a sample is present but get_mounted_sample is None.
Raises a SmartMagnetFaultException if the magnet detects no sample but the robot thinks a sample is mounted"""
#TODO tidy up
initial_state = self.pshell.eval("smart_magnet.state&")
#Not sure why unused, potentially can remove them
_ = (timeout, idle_time, interval)
initial_state = self.backend.eval("smart_magnet.state&")
logger.debug(f"checking smart magnet_initial state: {initial_state}")
if initial_state == "Paused":
self.pshell.eval("smart_magnet.set_supress(False)&")
self.pshell.eval("smart_magnet.set_resting_current()&")
self.backend.eval("smart_magnet.set_supress(False)&")
self.backend.eval("smart_magnet.set_resting_current()&")
elif initial_state == "Fault":
logger.error(f"tell smart magnet is in unknown state {initial_state}")
raise SmartMagnetFaultException
state = self.pshell.eval("smart_magnet.state&")
state = self.backend.eval("smart_magnet.state&")
try:
if state == "Busy":
logger.debug('state busy')
self.pshell.eval("smart_magnet.set_supress(True)&")
self.pshell.eval("smart_magnet.state&")
sample_present = True
logger.debug("state busy")
self.backend.eval("smart_magnet.set_supress(True)&")
self.backend.eval("smart_magnet.state&")
if self.get_mounted_sample() is None:
logger.warning("Check mount: A manually mounted sample is detected.")
logger.warning("Remove before mounting with the robot.")
raise ManualMountException
return True
elif state == "Ready":
logger.debug('No sample detected, ready to mount')
sample_present = False
logger.debug("No sample detected, ready to mount")
if self.get_mounted_sample():
logger.error("Check mount: No sample detected, but robot thinks is mounted")
raise SmartMagnetFaultException
@@ -491,174 +377,20 @@ class TellClient:
logger.debug("Smart magnet detection is paused")
return None
else:
self.pshell.eval("smart_magnet.set_supress(True)&")
self.backend.eval("smart_magnet.set_supress(True)&")
logger.error(f"Tell smart magnet is in unknown state {state}")
raise SmartMagnetFaultException
except Exception as e:
logger.error(f"check_smart_magnet_mounted failed: {e}")
raise e
class SimTellClient:
"""
Simulation-only Tell client.
Keeps behavior deterministic-ish and stateful without needing PShellClient.
Implement more methods as your callers need them.
"""
def __init__(self):
self._state = "Ready"
self._last_cmd_id = 1000
self._mounted_sample: str = ""
self._simulated_samples_info = {}
self._simulated_detected_pucks = []
self._simulated_current = 30.0
self._simulated_suppress = True
self._simulated_offset = 0.0
@property
def url(self):
return None
def _next_cmd_id(self) -> int:
self._last_cmd_id += 1
return self._last_cmd_id
def get_state(self) -> str:
return self._state
def is_ready(self) -> bool:
return self._state.lower() == "ready"
def is_busy(self) -> bool:
return self._state.lower() == "busy"
def wait_ready(self, timeout: float = 360.0):
# Keep it simple: flip to Ready quickly.
time.sleep(0.05)
self._state = "Ready"
def mount(
self,
address: SampleDewarAddress,
force: bool = False,
read_dm: bool = False,
auto_unmount: bool = False,
wait: bool = False,
timeout: float = 600.0,
):
SampleDewarAddress.model_validate(address)
if self.is_busy():
raise TellCommandWhileBusyException("mount received while robot is busy")
cmd_id = self._next_cmd_id()
self._state = "Busy"
segment = address.puck.segment
puck = address.puck.pos
sample = address.pin
self._mounted_sample = f"{segment}{puck}{sample}"
if wait:
self.wait_ready(timeout=timeout)
else:
# quickly become ready anyway, but asynchronously-ish
time.sleep(0.01)
self._state = "Ready"
return cmd_id
def unmount(self, force: bool = False, wait: bool = False, timeout: float = 360.0):
if self.is_busy():
raise TellCommandWhileBusyException("unmount received while robot is busy")
cmd_id = self._next_cmd_id()
self._state = "Busy"
self._mounted_sample = ""
if wait:
self.wait_ready(timeout=timeout)
else:
time.sleep(0.01)
self._state = "Ready"
return cmd_id
def get_mounted_sample(self) -> SampleDewarAddress | None:
ret = self._mounted_sample
if not ret:
return None
match = re.match(r"([A-Z])(\d)(\d{1,2})", ret)
if not match:
return None
segment, puck, sample = match.groups()
return SampleDewarAddress(puck=DewarAddress(segment=segment, pos=int(puck)), pin=int(sample))
class TellClientProxy:
"""
Lazy-connecting Tell client proxy that retries periodically.
- Server can start even if TELL is down.
- First use triggers connect; failures raise TellCommunicationError.
"""
def __init__(self, bl: MXBeamline, *, retry_interval_s: float = 2.0):
self._bl = bl
self._client: TellClient | None = None
self._retry_interval_s = float(retry_interval_s)
self._last_attempt_ts = 0.0
self._last_error: Exception | None = None
def _get_client(self) -> TellClient:
if self._client is not None:
return self._client
now = time.monotonic()
if now - self._last_attempt_ts < self._retry_interval_s and self._last_error is not None:
raise self._last_error
self._last_attempt_ts = now
try:
self._client = TellClient(self._bl)
self._last_error = None
return self._client
except TellCommunicationError as e:
self._last_error = e
raise
except Exception as e:
wrapped = TellCommunicationError(
"TELL connection failed",
operation="CONNECT",
)
self._last_error = wrapped
raise wrapped from e
@property
def url(self):
return self._get_client().url
# Delegate methods used by DAQ; add more as needed
def get_mounted_sample(self) -> SampleDewarAddress | None:
return self._get_client().get_mounted_sample()
def get_state(self):
return self._get_client().get_state()
def wait_not_busy(self, timeout: float = 360.0):
return self._get_client().wait_not_busy(timeout=timeout)
def check_enable_motion(self):
return self._get_client().check_enable_motion()
def set_in_mount_position(self, value):
return self._get_client().set_in_mount_position(value)
def mount(self, *args, **kwargs):
return self._get_client().mount(*args, **kwargs)
def unmount(self, *args, **kwargs):
return self._get_client().unmount(*args, **kwargs)
def abort_cmd(self):
return self._get_client().abort_cmd()
def make_tell_client(bl: MXBeamline) -> TellClient | SimTellClient | TellClientProxy:
def make_tell_client(bl: MXBeamline) -> TellClient:
if bl == MXBeamline.SIMULATED:
return SimTellClient()
return TellClientProxy(bl, retry_interval_s=2.0)
backend = SimTellBackend()
else:
backend = LazyTellBackend(
factory=lambda: PShellTellBackend(bl),
retry_interval_s=2.0,
)
return TellClient(bl, backend=backend)
+2 -2
View File
@@ -36,8 +36,8 @@ if __name__ == "__main__":
default_gonio_cam_addr = "axis-accc8ed2972e.psi.ch"
default_gonio_camera_id = 3
case MXBeamline.X10SA:
default_url = "http://127.0.0.1:5210"
default_zmq_addr = "tcp://x10sa-spark-01:9091" #"tcp://x10sa-pserv-01:9089" #
default_url = "http://mx-x10sa-queue-01.psi.ch:5210" #"http://127.0.0.1:5210"#
default_zmq_addr = "tcp://x10sa-spark-01:9091" # "tcp://x10sa-spark-01:9091" #
default_pred_zmq_addr = "tcp://x10sa-spark-01:9091" #"tcp://sls-gpu-003:9089"#""
default_beamline_cam_addr = "axis-accc8eb02488.psi.ch"
default_gonio_cam_addr = "axis-accc8ea5e463.psi.ch"
+357 -12
View File
@@ -1,4 +1,5 @@
import time
import requests
import jwt
from PySide6.QtCore import Qt, Slot, Signal, QTimer, QSettings
@@ -12,6 +13,7 @@ from PySide6.QtWidgets import (
QDockWidget,
QTabWidget, QFrame, QSizePolicy, QLabel)
from aare.common.auth_models import BatonStatus, BatonRequestStatus
from aare.common.coordinate import Coordinate, SmargonCoordinate
from aare.common.diffraction_geometry import DiffractionGeometry
from aare.common.logger_config import setup_logger
@@ -47,13 +49,18 @@ from aare.gui.threads.camera_thread import SampleCameraThread
from aare.gui.threads.prediction_subscriber import PredictionSubscriber
from aare.gui.threads.daq_worker import DAQWorker
from aare.gui.threads.jfjoch_viewer import JFJochDBusClient
from aare.gui.tutorials.tutorial_registration import register_tutorials
from aare.gui.widgets.alert_banner import AlertBanner
from aare.gui.widgets.baton_request_dialog import BatonRequestDialog, BatonPendingDialog
from aare.gui.widgets.camera_image import SampleCameraImageLabel
from aare.gui.widgets.no_wheel_scroll_area import NoWheelScrollArea
from aare.gui.widgets.status_bar import StatusBar
from aare.gui.widgets.video_image import VideoGraphicsView
from aare.gui.panels.fluorescence_panel import FluorescencePanel
from aare.gui.panels.automation_panel import WorkflowPanel
from aare.gui.threads.workflow_sse_client import WorkflowSSEClient
logger = setup_logger("aareGUI")
class MainWindow(QMainWindow):
@@ -69,6 +76,7 @@ class MainWindow(QMainWindow):
gonio_cam_id: int | None
):
super().__init__()
self.__base_url = base_url
self.__token = token
self.__mounting = False
@@ -76,6 +84,11 @@ class MainWindow(QMainWindow):
self._beamline_recovery_dialog = None
self._controls_help_dialog = None
self._cleanup_done = False
self._default_window_state = None
self._waiting_for_baton_response: bool = False
self._baton_request_dialog: BatonRequestDialog | None = None
self._baton_pending_dialog: BatonPendingDialog | None = None
# Tutorial manager (define tutorials after widgets exist)
self._tutorial_event_bus = TutorialEventBus(self)
@@ -114,6 +127,9 @@ class MainWindow(QMainWindow):
self.alert_banner = AlertBanner(parent=root_widget)
root_layout.addWidget(self.alert_banner)
self.alert_banner_secondary = AlertBanner(parent=root_widget)
root_layout.addWidget(self.alert_banner_secondary)
top_widget = QWidget(parent=root_widget)
top_widget_layout = QHBoxLayout(top_widget)
top_widget.setLayout(top_widget_layout)
@@ -221,6 +237,11 @@ class MainWindow(QMainWindow):
self.ref_tools_dock.setAllowedAreas(Qt.DockWidgetArea.BottomDockWidgetArea)
self.addDockWidget(Qt.DockWidgetArea.BottomDockWidgetArea, self.ref_tools_dock)
self.tabifyDockWidget(self.ref_tools_dock, self.tell_samples_dock)
if self.__decoded_token.staff:
self.ref_tools_dock.show()
self.ref_tools_dock.raise_()
else:
self.ref_tools_dock.hide()
self.sample_logic = SampleMountLogic()
@@ -277,11 +298,23 @@ class MainWindow(QMainWindow):
self.addDockWidget(Qt.DockWidgetArea.RightDockWidgetArea, self.smargon_trace_dock)
self.smargon_trace_dock.hide()
# === Workflow Panel ===
self.workflow_panel = WorkflowPanel()
self.workflow_dock = QDockWidget("Workflow", self)
self.workflow_dock.setObjectName("workflow_dock")
self.workflow_dock.setWidget(self.workflow_panel)
self.workflow_dock.setAllowedAreas(
Qt.DockWidgetArea.RightDockWidgetArea | Qt.DockWidgetArea.LeftDockWidgetArea
)
self.addDockWidget(Qt.DockWidgetArea.RightDockWidgetArea, self.workflow_dock)
self.workflow_dock.hide() # Hidden by default
root_layout.addWidget(top_widget)
self.setCentralWidget(root_widget)
self.setWindowTitle("AareGUI")
self.create_menu_bar()
self._capture_default_window_state()
self._restore_window_state()
# Define tutorials now that the UI exists
@@ -300,13 +333,23 @@ class MainWindow(QMainWindow):
self.setStatusBar(self.status_bar)
self.daq = DAQWorker(base_url=self.__base_url, token=self.__token)
self.daq.baton_status_changed.connect(self.status_bar.update_baton_status)
self.daq.baton_status_changed.connect(self._on_baton_status_changed)
self.daq.baton_request_result.connect(self._on_baton_request_result)
self.daq.baton_response_result.connect(self._on_baton_response_result)
self.daq.baton_timeout_checked.connect(self._on_baton_timeout_checked)
self.status_bar.baton_request_received.connect(self._show_baton_request_dialog)
self.status_bar.baton_request_accepted.connect(self._accept_baton_request)
self.status_bar.baton_request_refused.connect(self._refuse_baton_request)
self.daq.spreadsheet.connect(self.tell_samples.new_sample_list)
if self.__decoded_token.staff:
self.daq.reference_tools.connect(self.ref_tools_panel.new_list)
self.beamline.samcam.changed.connect(self.daq.samcam_settings)
self.beamline.samcam.screenshot_requested.connect(self.daq.send_screenshot_db)
self.beamline.loopctr.background.clicked.connect(self.daq.alc_background)
self.beamline.loopctr.find_tip.clicked.connect(self.daq.center_loop)
self.beamline.loopctr.bounding_box.clicked.connect(self.daq.ml_bounding_box)
self.daq.raster_generated_by_ml.connect(self.raster.update_active_grid_request)
@@ -411,9 +454,21 @@ class MainWindow(QMainWindow):
self.data_collection.simple.parameters_changed.connect(self.daq.smart_params)
self.raster.grid_scan_size_changed.connect(self.data_collection.raster.grid_scan_size_change)
self.status_bar.set_pgroup.connect(self.daq.set_pgroup)
self.status_bar.end_session.connect(self.daq.end_session)
self.status_bar.force_session.connect(self.daq.force_session)
self.status_bar.request_baton.connect(self.daq.request_baton)
self.status_bar.cancel_baton_request.connect(self.daq.cancel_baton_request)
self.status_bar.release_baton.connect(self.daq.release_baton)
self.status_bar.baton_request_accepted.connect(
lambda: self.daq.respond_to_baton_request(True)
)
self.status_bar.baton_request_refused.connect(
lambda: self.daq.respond_to_baton_request(False)
)
self.status_bar.dewar_exchange.connect(self.daq.dewar_exchange)
self.status_bar.sample_exchange.connect(self.daq.sample_exchange)
self.status_bar.sample_alignment.connect(self.daq.sample_alignment)
@@ -421,6 +476,7 @@ class MainWindow(QMainWindow):
self.status_bar.close_shutter.connect(self.daq.close_shutter)
self.status_bar.open_shutter.connect(self.daq.open_shutter)
self.rotation.file_ready.connect(self.viewer.load_image)
self.raster.image_selected.connect(self.viewer.load_image)
@@ -474,11 +530,59 @@ class MainWindow(QMainWindow):
self.daq.fluorimeter_spectrum_update.connect(self.fluor_panel.update_plot)
self.daq.fluorimeter_spectrum_update.connect(lambda: self.fluor_panel_dock.setVisible(True))
# === Alert/Status Message Routing ===
# Primary alert banner: Infrastructure devices (Server/Tell/Smargon/Aerotech)
self.daq.polled_devices_status.connect(self.alert_banner.show_message)
# Secondary alert banner: Detector errors (JFJoch)
self.daq.detector_error.connect(self.alert_banner_secondary.show_message)
# Status bar: General status messages (not device connection status)
self.daq.status_message.connect(self.status_bar.show_connection_message)
self.daq.status_message.connect(self.alert_banner.show_message)
register_tutorials(self, self.tutorial_manager)
# Workflow SSE client
if self.__base_url is not None:
self.workflow_sse = WorkflowSSEClient(self.__base_url, self.__token, self)
self.workflow_sse.runtime_changed.connect(self.workflow_panel.update_runtime)
self.workflow_sse.control_changed.connect(self.workflow_panel.update_control)
self.workflow_sse.workflow_event.connect(
lambda e: self.workflow_panel.on_workflow_event(e.event_type, e.message)
)
self.workflow_sse.connect()
else:
self.workflow_sse = None
# Connect workflow panel signals to DAQ worker
self.workflow_panel.request_queue_refresh.connect(self.daq.workflow_load_queue)
self.workflow_panel.request_add_sample.connect(self.daq.workflow_add_sample)
self.workflow_panel.request_add_samples.connect(self.daq.workflow_add_samples)
self.workflow_panel.request_delete_item.connect(self.daq.workflow_delete_item)
self.workflow_panel.request_move_item.connect(self.daq.workflow_move_item)
self.workflow_panel.request_clear_queue.connect(self.daq.workflow_clear_queue)
self.workflow_panel.request_start_item.connect(self.daq.workflow_start_item)
self.workflow_panel.request_next_step.connect(self.daq.workflow_next_step)
self.workflow_panel.request_pause.connect(self.daq.workflow_pause)
self.workflow_panel.request_resume.connect(self.daq.workflow_resume)
self.workflow_panel.request_abort.connect(self.daq.workflow_abort)
self.workflow_panel.request_skip.connect(self.daq.workflow_skip)
self.workflow_panel.request_start_automation.connect(self.daq.workflow_start_automation)
self.workflow_panel.request_stop_automation.connect(self.daq.workflow_stop_automation)
self.workflow_panel.request_skip_sample.connect(self.daq.workflow_skip_sample)
# DAQ worker -> workflow panel
self.daq.workflow_queue_loaded.connect(self.workflow_panel.update_queue)
self.daq.workflow_item_updated.connect(self.workflow_panel.update_current_item)
# Forward sample list to workflow panel for "Add All" feature
self.daq.spreadsheet.connect(
lambda slist: self.workflow_panel.update_available_samples(slist.s)
)
# Initial load
QTimer.singleShot(1000, self.daq.workflow_load_queue)
@Slot(QPixmap)
def _on_samcam_prediction_pixmap(self, pix: QPixmap) -> None:
self._last_pred_image_ts = time.monotonic()
@@ -522,35 +626,41 @@ class MainWindow(QMainWindow):
file_menu.addAction(quit_action)
view_menu = menu_bar.addMenu("View")
show_samples_action = QAction("Show Sample List", self)
show_samples_action.setCheckable(True)
show_samples_action.setChecked(True)
show_samples_action.triggered.connect(lambda: self.tell_samples_dock.setVisible(show_samples_action.isChecked()))
# When dock visibility changes, update the action's checked state
show_samples_action.triggered.connect(lambda checked: self.tell_samples_dock.setVisible(checked))
self.tell_samples_dock.visibilityChanged.connect(show_samples_action.setChecked)
view_menu.addAction(show_samples_action)
if self.__decoded_token.staff:
show_reference_tools_action = QAction("Show Reference Tools", self)
show_reference_tools_action.setCheckable(True)
show_reference_tools_action.setChecked(True)
show_reference_tools_action.triggered.connect(
lambda checked: self.ref_tools_dock.setVisible(checked)
)
self.ref_tools_dock.visibilityChanged.connect(show_reference_tools_action.setChecked)
view_menu.addAction(show_reference_tools_action)
show_job_list_action = QAction("Show job List", self)
show_job_list_action.setCheckable(True)
show_job_list_action.setChecked(True)
show_job_list_action.triggered.connect(self.job_list_dock.setVisible)
# When dock visibility changes, update the action's checked state
show_job_list_action.triggered.connect(lambda: self.job_list_dock.setVisible(show_job_list_action.isChecked()))
show_job_list_action.triggered.connect(lambda checked: self.job_list_dock.setVisible(checked))
self.job_list_dock.visibilityChanged.connect(show_job_list_action.setChecked)
view_menu.addAction(show_job_list_action)
show_manual_sample_action = QAction("Show manual sample", self)
show_manual_sample_action.setCheckable(True)
show_manual_sample_action.setChecked(True)
show_manual_sample_action.triggered.connect(self.manual_sample_dock.setVisible)
# When dock visibility changes, update the action's checked state
show_manual_sample_action.triggered.connect(lambda: self.manual_sample_dock.setVisible(show_manual_sample_action.isChecked()))
show_manual_sample_action.triggered.connect(lambda checked: self.manual_sample_dock.setVisible(checked))
self.manual_sample_dock.visibilityChanged.connect(show_manual_sample_action.setChecked)
view_menu.addAction(show_manual_sample_action)
show_face_panel_action = QAction("Show face detection", self)
show_face_panel_action.setCheckable(True)
show_face_panel_action.setChecked(False)
#show_face_panel_action.triggered.connect(self.face_panel_dock.setVisible)
show_face_panel_action.triggered.connect(lambda checked: self.face_panel_dock.setVisible(checked))
self.face_panel_dock.visibilityChanged.connect(show_face_panel_action.setChecked)
view_menu.addAction(show_face_panel_action)
@@ -566,11 +676,19 @@ class MainWindow(QMainWindow):
show_smargon_trace_action.setCheckable(True)
show_smargon_trace_action.setChecked(False)
show_smargon_trace_action.triggered.connect(lambda checked: self.smargon_trace_dock.setVisible(checked))
self.smargon_trace_dock.visibilityChanged.connect(show_smargon_trace_action.setChecked)
self.smargon_trace_dock.visibilityChanged.connect(
lambda visible: self.smargon_trace_panel.refresh_plot(force=True) if visible else None
)
view_menu.addAction(show_smargon_trace_action)
show_workflow_action = QAction("Show Workflow Panel", self)
show_workflow_action.setCheckable(True)
show_workflow_action.setChecked(False)
show_workflow_action.triggered.connect(lambda checked: self.workflow_dock.setVisible(checked))
self.workflow_dock.visibilityChanged.connect(show_workflow_action.setChecked)
view_menu.addAction(show_workflow_action)
show_log_action = QAction("Show Log", self)
show_log_action.setCheckable(True)
show_log_action.setChecked(False)
@@ -578,6 +696,12 @@ class MainWindow(QMainWindow):
self.log_dock.visibilityChanged.connect(show_log_action.setChecked)
view_menu.addAction(show_log_action)
view_menu.addSeparator()
restore_default_view_action = QAction("Restore Default View", self)
restore_default_view_action.triggered.connect(self.restore_default_view)
view_menu.addAction(restore_default_view_action)
help_menu = menu_bar.addMenu("Help")
about_action = QAction("About", self) # Action for 'About'
about_action.triggered.connect(self.show_about_dialog)
@@ -606,6 +730,32 @@ class MainWindow(QMainWindow):
start_interactive_tutorial_action.triggered.connect(self.start_interactive_tutorial)
help_menu.addAction(start_interactive_tutorial_action)
def _capture_default_window_state(self) -> None:
self._default_window_state = self.saveState()
@Slot()
def restore_default_view(self) -> None:
if self._default_window_state is not None:
self.restoreState(self._default_window_state)
self.tell_samples_dock.setVisible(True)
self.job_list_dock.setVisible(True)
self.manual_sample_dock.setVisible(True)
self.face_panel_dock.setVisible(False)
self.fluor_panel_dock.setVisible(False)
self.smargon_trace_dock.setVisible(False)
self.log_dock.setVisible(False)
if self.__decoded_token.staff:
self.ref_tools_dock.setVisible(True)
self.ref_tools_dock.raise_()
else:
self.ref_tools_dock.setVisible(False)
self.tell_samples_dock.raise_()
self.video_tab.setCurrentIndex(0)
def show_about_dialog(self):
QMessageBox.about(
self,
@@ -674,6 +824,189 @@ class MainWindow(QMainWindow):
self.__mounting = False
self.video_tab.setCurrentIndex(0)
# ========== BATON DIALOG HANDLING ==========
@Slot(dict)
def _show_baton_request_dialog(self, payload: dict):
requester = str(payload.get("requester") or "Another user")
timeout = int(payload.get("timeout") or 30)
if self._baton_request_dialog is not None and self._baton_request_dialog.isVisible():
return
self._baton_request_dialog = BatonRequestDialog(
requester=requester,
timeout_seconds=timeout,
parent=self,
)
self._baton_request_dialog.accepted_signal.connect(self.status_bar._on_baton_dialog_accepted)
self._baton_request_dialog.refused_signal.connect(self.status_bar._on_baton_dialog_refused)
self._baton_request_dialog.show()
self._baton_request_dialog.raise_()
self._baton_request_dialog.activateWindow()
@Slot()
def _accept_baton_request(self):
self.daq.respond_to_baton_request(True)
@Slot()
def _refuse_baton_request(self):
self.daq.respond_to_baton_request(False)
@Slot()
def _accept_baton_request(self):
self.daq.respond_to_baton_request(True)
@Slot()
def _refuse_baton_request(self):
self.daq.respond_to_baton_request(False)
@Slot(BatonStatus)
def _on_baton_status_changed(self, status: BatonStatus):
"""Close the pending dialog immediately if the baton request has been resolved via SSE."""
if self._waiting_for_baton_response and not status.you_have_pending_request:
self._waiting_for_baton_response = False
self._close_baton_pending_dialog()
if status.you_are_holder:
self.alert_banner.show_message("Baton acquired!", False, auto_clear_ms=10000)
else:
self.alert_banner.show_message("Request declined or cancelled", False, auto_clear_ms=10000)
@Slot(dict)
def _on_baton_request_result(self, result: dict):
"""Handle result of our baton request - show waiting banner with countdown."""
if result.get("granted"):
self._waiting_for_baton_response = False
self.alert_banner.show_message("Baton acquired!", False, auto_clear_ms=10000)
logger.info("Baton acquired")
# Close the pending dialog immediately before showing p-group prompt
self._close_baton_pending_dialog()
available_pgroups = [str(p).strip() for p in (self.__decoded_token.pgroups or []) if p is not None and str(p).strip()]
if len(available_pgroups) == 1:
self.status_bar.set_pgroup.emit(available_pgroups[0])
else:
self.status_bar._after_baton_granted_select_pgroup()
elif result.get("pending"):
self._waiting_for_baton_response = True
timeout = result.get("timeout_seconds", 30)
holder = result.get("message", "Waiting for response...")
if getattr(self, "_baton_pending_dialog", None) is None:
target_user = holder.replace("Request sent to ", "")
self._baton_pending_dialog = BatonPendingDialog(target_user=target_user, timeout_seconds=timeout,
parent=self)
self._baton_pending_dialog.cancelled_signal.connect(self.daq.cancel_baton_request)
self._baton_pending_dialog.show()
else:
self._baton_pending_dialog.update_remaining(timeout)
self.alert_banner.show_waiting(f"Requesting control - {holder}", timeout)
logger.info(f"Baton request pending - {timeout}s timeout")
elif result.get("queued"):
self._waiting_for_baton_response = True
self.alert_banner.show_waiting("Control transfer queued - waiting for beamline")
logger.info("Baton transfer queued")
if getattr(self, "_baton_pending_dialog", None) is None:
self._baton_pending_dialog = BatonPendingDialog(target_user="Current Holder", timeout_seconds=0,
parent=self)
self._baton_pending_dialog.cancelled_signal.connect(self.daq.cancel_baton_request)
self._baton_pending_dialog.show()
self._baton_pending_dialog.set_queued_state()
elif result.get("already_holder"):
self._waiting_for_baton_response = False
logger.debug("Already baton holder")
self._close_baton_pending_dialog()
elif result.get("error"):
self._waiting_for_baton_response = False
self.alert_banner.show_message(result.get("message", "Request failed"), True)
logger.warning(f"Baton request failed: {result.get('message')}")
self._close_baton_pending_dialog()
@Slot(dict)
def _on_baton_response_result(self, result: dict):
"""Handle result after we responded to someone else's request."""
logger.debug(f"Baton response result: {result}")
if result.get("accepted"):
self._waiting_for_baton_response = False
self.alert_banner.show_message("Control transferred", False, auto_clear_ms=10000)
self._close_baton_dialog()
self.status_bar.update_baton_status(self.status_bar._baton_status) # refresh label state
elif result.get("refused"):
self._waiting_for_baton_response = False
self.alert_banner.show_message("Request declined", False, auto_clear_ms=10000)
self._close_baton_dialog()
self.status_bar.update_baton_status(self.status_bar._baton_status) # refresh label state
else:
logger.debug(f"replied with {result}")
@Slot(dict)
def _on_baton_timeout_checked(self, result: dict):
"""Refresh waiting UI when the backend confirms timeout state."""
logger.debug(f"Baton timeout checked: {result}")
if result.get("pending"):
remaining = int(result.get("remaining_seconds", 0))
if self._waiting_for_baton_response:
self.alert_banner.show_waiting("Requesting control", remaining)
if getattr(self, "_baton_pending_dialog", None) is not None:
self._baton_pending_dialog.update_remaining(remaining)
elif result.get("granted"):
self._waiting_for_baton_response = False
self.alert_banner.show_message("Baton acquired!", False, auto_clear_ms=10000)
# Close the pending dialog immediately before showing p-group prompt
self._close_baton_pending_dialog()
# P-group logic will be handled automatically by the status_bar stream update
elif result.get("queued"):
self._waiting_for_baton_response = True
self.alert_banner.show_waiting("Control transfer queued - waiting for beamline")
if getattr(self, "_baton_pending_dialog", None) is not None:
self._baton_pending_dialog.set_queued_state()
else:
self._baton_pending_dialog = BatonPendingDialog(target_user="Current Holder", timeout_seconds=0,
parent=self)
self._baton_pending_dialog.cancelled_signal.connect(self.daq.cancel_baton_request)
self._baton_pending_dialog.show()
self._baton_pending_dialog.set_queued_state()
elif result.get("refused"):
self._waiting_for_baton_response = False
self.alert_banner.show_message("Request declined", False, auto_clear_ms=10000)
self._close_baton_pending_dialog()
else:
logger.debug(f"replied with {result}")
self.alert_banner.clear_message()
self._close_baton_pending_dialog()
def _close_baton_dialog(self) -> None:
if getattr(self, "_baton_request_dialog", None) is not None:
try:
self._baton_request_dialog.close()
finally:
self._baton_request_dialog = None
def _close_baton_pending_dialog(self) -> None:
if getattr(self, "_baton_pending_dialog", None) is not None:
try:
if hasattr(self._baton_pending_dialog, '_timer'):
self._baton_pending_dialog._timer.stop()
self._baton_pending_dialog.close()
finally:
self._baton_pending_dialog = None
def _restore_window_state(self) -> None:
settings = QSettings()
geometry = settings.value("main_window/geometry")
@@ -684,6 +1017,9 @@ class MainWindow(QMainWindow):
if state is not None:
self.restoreState(state)
if not self.__decoded_token.staff:
self.ref_tools_dock.hide()
def closeEvent(self, event) -> None:
try:
settings = QSettings()
@@ -692,6 +1028,12 @@ class MainWindow(QMainWindow):
except Exception as e:
logger.warning(f"Failed to save main window state: {e}")
# Release baton before closing
try:
self.daq.release_baton_on_close()
except Exception as e:
logger.warning(f"Failed to release baton on close: {e}")
try:
self.cleanup()
except Exception as e:
@@ -710,6 +1052,9 @@ class MainWindow(QMainWindow):
except Exception as e:
logger.warning(f"Failed to stop _samcam_source_timer: {e}")
if hasattr(self, "workflow_sse") and self.workflow_sse is not None:
self.workflow_sse.disconnect()
for attr_name in (
"camera_thread",
"prediction_thread",
+6 -6
View File
@@ -1,7 +1,7 @@
from PySide6.QtCore import Signal, Slot
from PySide6.QtGui import Qt
from PySide6.QtWidgets import QWidget, QGridLayout, QLabel, QPushButton
from aare.common.coordinate import Coordinate
from aare.common.coordinate import Coordinate, AerotechCoordinate
from aare.common.models import DAQStatusModel
from aare.gui.widgets.button_with_payload import ButtonWithPayload
@@ -11,7 +11,7 @@ from aare.gui.widgets.title_label import TitleLabel
DEFAULT_ABR_STEP_UM = 5
class AbrTweakButtons(QWidget):
abr_tweak = Signal(Coordinate)
abr_tweak = Signal(AerotechCoordinate)
def __init__(self, step_mm, parent=None):
super().__init__(parent)
@@ -66,14 +66,14 @@ class AbrTweakButtons(QWidget):
@Slot(dict)
def abr_button(self, payload: dict):
self.abr_tweak.emit(Coordinate(x=self.__step_mm * payload["x"], y=self.__step_mm * payload["y"], z=self.__step_mm * payload["z"]))
self.abr_tweak.emit(AerotechCoordinate(at_mm=Coordinate(x=self.__step_mm * payload["x"], y=self.__step_mm * payload["y"], z=self.__step_mm * payload["z"])))
@Slot(float)
def set_step(self, val_um: float):
self.__step_mm = val_um / 1000.0
class AbrTweakWidget(QWidget):
abr_tweak = Signal(Coordinate)
abr_tweak = Signal(AerotechCoordinate)
abr_save = Signal()
abr_goto_meas = Signal()
@@ -109,8 +109,8 @@ class AbrTweakWidget(QWidget):
def goto_button_pressed(self):
self.abr_goto_meas.emit()
@Slot(Coordinate)
def abr_button_pressed(self, c: Coordinate):
@Slot(AerotechCoordinate)
def abr_button_pressed(self, c: AerotechCoordinate):
self.abr_tweak.emit(c)
@Slot()
+730
View File
@@ -0,0 +1,730 @@
"""
Workflow automation panel for queue management and step control.
Displays:
- Queue items with status (supports drag-drop from sample list)
- Current step progress
- Control buttons (Start Guided/Start Automation/Pause/Resume/Abort/Skip)
- Live status from SSE
"""
from __future__ import annotations
import json
from typing import Any
from PySide6.QtCore import Qt, Signal, Slot, QTimer, QMimeData
from PySide6.QtGui import QColor, QDragEnterEvent, QDropEvent, QKeySequence, QShortcut
from PySide6.QtWidgets import (
QWidget,
QVBoxLayout,
QHBoxLayout,
QLabel,
QPushButton,
QListWidget,
QListWidgetItem,
QProgressBar,
QGroupBox,
QFrame,
QSizePolicy,
QAbstractItemView,
QMenu,
QMessageBox,
)
from aare.common.automation_models import (
QueueItem,
QueueItemStatus,
RuntimeState,
ControlState,
WorkflowStepRecord,
)
from aare.common.models import SampleShortInfo, SampleShortInfoList
from aare.common.logger_config import setup_logger
logger = setup_logger("aareGUI")
class StepProgressWidget(QWidget):
"""Shows progress through workflow steps."""
def __init__(self, parent: QWidget | None = None):
super().__init__(parent)
self._steps: list[WorkflowStepRecord] = []
self._current_index = 0
self._setup_ui()
def _setup_ui(self) -> None:
layout = QVBoxLayout(self)
layout.setContentsMargins(4, 4, 4, 4)
layout.setSpacing(2)
self._step_labels: list[QLabel] = []
self._container = QWidget()
self._container_layout = QVBoxLayout(self._container)
self._container_layout.setContentsMargins(0, 0, 0, 0)
self._container_layout.setSpacing(2)
layout.addWidget(self._container)
def set_steps(self, steps: list[WorkflowStepRecord], current_index: int) -> None:
"""Update the step display."""
self._steps = steps
self._current_index = current_index
for lbl in self._step_labels:
lbl.deleteLater()
self._step_labels.clear()
for i, step in enumerate(steps):
lbl = QLabel(f"{i + 1}. {step.kind}")
lbl.setStyleSheet(self._style_for_step(i, step.status))
self._container_layout.addWidget(lbl)
self._step_labels.append(lbl)
def update_step(self, index: int, status: str) -> None:
"""Update a single step's status."""
if 0 <= index < len(self._step_labels):
self._step_labels[index].setStyleSheet(self._style_for_step(index, status))
def clear_steps(self) -> None:
"""Clear all step labels."""
for lbl in self._step_labels:
lbl.deleteLater()
self._step_labels.clear()
self._steps = []
def _style_for_step(self, index: int, status: str) -> str:
"""Get stylesheet for step based on status."""
base = "padding: 4px; border-radius: 3px; "
if status == "success":
return base + "background-color: #90EE90; color: #006400;"
elif status == "running":
return base + "background-color: #87CEEB; color: #00008B; font-weight: bold;"
elif status == "failed":
return base + "background-color: #FFB6C1; color: #8B0000;"
elif status == "skipped":
return base + "background-color: #D3D3D3; color: #696969; text-decoration: line-through;"
elif status == "paused":
return base + "background-color: #FFE4B5; color: #8B4513;"
else: # pending
return base + "background-color: #F0F0F0; color: #808080;"
class DraggableQueueListWidget(QListWidget):
"""
QListWidget that accepts drops from TellSamplePanel.
Supports:
- Drag-drop samples from tell_sample_panel
- Internal reordering via drag
- Delete key to remove items
"""
samples_dropped = Signal(list) # list[SampleShortInfo]
item_reordered = Signal(str, int) # item_id, new_index
delete_requested = Signal(list) # list[item_ids]
start_item_requested = Signal(str) # item_id
def __init__(self, parent: QWidget | None = None):
super().__init__(parent)
self.setAcceptDrops(True)
self.setDragEnabled(True)
self.setDragDropMode(QAbstractItemView.DragDropMode.DragDrop)
self.setDefaultDropAction(Qt.DropAction.MoveAction)
self.setSelectionMode(QAbstractItemView.SelectionMode.ExtendedSelection)
self.setContextMenuPolicy(Qt.ContextMenuPolicy.CustomContextMenu)
self.customContextMenuRequested.connect(self._show_context_menu)
# Delete shortcut
self._delete_shortcut = QShortcut(QKeySequence.StandardKey.Delete, self)
self._delete_shortcut.activated.connect(self._on_delete_pressed)
# Store item_id -> row mapping
self._item_ids: list[str] = []
def set_item_ids(self, item_ids: list[str]) -> None:
"""Track item IDs for reordering."""
self._item_ids = item_ids
def dragEnterEvent(self, event: QDragEnterEvent) -> None:
"""Accept drops from sample panels."""
mime = event.mimeData()
if mime.hasText():
try:
text = mime.text()
if text.startswith("{") or text.startswith("["):
event.acceptProposedAction()
return
except Exception:
pass
if event.source() == self:
event.acceptProposedAction()
return
event.ignore()
def dragMoveEvent(self, event) -> None:
"""Show drop indicator."""
if event.mimeData().hasText() or event.source() == self:
event.acceptProposedAction()
else:
event.ignore()
def dropEvent(self, event: QDropEvent) -> None:
"""Handle drop - either samples from panel or internal reorder."""
mime = event.mimeData()
if mime.hasText():
text = mime.text()
try:
sample_list = SampleShortInfoList.model_validate_json(text)
if sample_list.s:
self.samples_dropped.emit(sample_list.s)
event.acceptProposedAction()
return
except Exception:
pass
try:
sample = SampleShortInfo.model_validate_json(text)
self.samples_dropped.emit([sample])
event.acceptProposedAction()
return
except Exception:
pass
if event.source() == self:
drop_row = self.indexAt(event.position().toPoint()).row()
if drop_row < 0:
drop_row = self.count()
selected = self.selectedItems()
if selected and self._item_ids:
for item in selected:
row = self.row(item)
if 0 <= row < len(self._item_ids):
item_id = self._item_ids[row]
self.item_reordered.emit(item_id, drop_row)
event.acceptProposedAction()
return
event.ignore()
def _on_delete_pressed(self) -> None:
"""Handle delete key press."""
selected = self.selectedItems()
if not selected:
return
item_ids = []
for item in selected:
row = self.row(item)
if 0 <= row < len(self._item_ids):
item_ids.append(self._item_ids[row])
if item_ids:
self.delete_requested.emit(item_ids)
def _show_context_menu(self, pos) -> None:
"""Show context menu for queue items."""
item = self.itemAt(pos)
if not item:
return
row = self.row(item)
if row < 0 or row >= len(self._item_ids):
return
menu = QMenu(self)
delete_action = menu.addAction("🗑 Remove from queue")
start_action = menu.addAction("▶ Start this item")
menu.addSeparator()
move_top_action = menu.addAction("⬆ Move to top")
move_bottom_action = menu.addAction("⬇ Move to bottom")
action = menu.exec_(self.mapToGlobal(pos))
if action == delete_action:
self.delete_requested.emit([self._item_ids[row]])
elif action == start_action:
self.start_item_requested.emit(self._item_ids[row])
elif action == move_top_action:
self.item_reordered.emit(self._item_ids[row], 0)
elif action == move_bottom_action:
self.item_reordered.emit(self._item_ids[row], 999999)
class WorkflowPanel(QWidget):
"""
Main workflow panel combining queue display and controls.
Supports:
- Drag-drop samples from TELL sample panel
- Queue management (delete, reorder, clear, add all)
- Step-by-step guided mode (explicit start)
- Full automation mode (explicit start)
"""
# Signals for DAQ worker
request_queue_refresh = Signal()
request_add_sample = Signal(object) # SampleShortInfo
request_add_samples = Signal(list) # list[SampleShortInfo]
request_delete_item = Signal(str) # item_id
request_move_item = Signal(str, int) # item_id, new_order_index
request_clear_queue = Signal()
request_start_item = Signal(str) # item_id
request_next_step = Signal()
request_pause = Signal()
request_resume = Signal()
request_abort = Signal()
request_skip = Signal() # Skip current step
request_skip_sample = Signal() # Skip entire current sample
request_start_automation = Signal()
request_stop_automation = Signal()
request_start_guided = Signal(str) # Start guided mode with item_id
def __init__(self, parent: QWidget | None = None):
super().__init__(parent)
self._queue_items: list[QueueItem] = []
self._all_samples: list[SampleShortInfo] = []
self._runtime: RuntimeState | None = None
self._control: ControlState | None = None
self._automation_enabled = False
self._was_paused_before_automation = False # Track pause state before automation
self._setup_ui()
def _setup_ui(self) -> None:
layout = QVBoxLayout(self)
layout.setContentsMargins(8, 8, 8, 8)
layout.setSpacing(8)
# === Status Section ===
status_group = QGroupBox("Current Status")
status_layout = QVBoxLayout(status_group)
self._status_label = QLabel("⏹️ Idle")
self._status_label.setStyleSheet("font-size: 14px; font-weight: bold;")
status_layout.addWidget(self._status_label)
self._current_item_label = QLabel("No item running")
status_layout.addWidget(self._current_item_label)
self._step_progress = StepProgressWidget()
status_layout.addWidget(self._step_progress)
layout.addWidget(status_group)
# === Start Buttons (explicit start required) ===
start_group = QGroupBox("Start Processing")
start_layout = QHBoxLayout(start_group)
self._start_guided_btn = QPushButton("▶ Start Guided Mode")
self._start_guided_btn.setStyleSheet("background-color: #4CAF50; color: white; font-weight: bold; padding: 8px;")
self._start_guided_btn.setToolTip("Start processing the first pending sample in guided (step-by-step) mode")
self._start_guided_btn.clicked.connect(self._on_start_guided)
self._start_auto_btn = QPushButton("▶▶ Start Automation")
self._start_auto_btn.setStyleSheet("background-color: #2196F3; color: white; font-weight: bold; padding: 8px;")
self._start_auto_btn.setToolTip("Start fully automated processing of all samples")
self._start_auto_btn.clicked.connect(self._on_start_automation)
self._stop_btn = QPushButton("⏹ Stop")
self._stop_btn.setStyleSheet("background-color: #9E9E9E; color: white; font-weight: bold; padding: 8px;")
self._stop_btn.setToolTip("Stop automation mode (current step will complete)")
self._stop_btn.clicked.connect(self._on_stop)
self._stop_btn.setVisible(False)
start_layout.addWidget(self._start_guided_btn)
start_layout.addWidget(self._start_auto_btn)
start_layout.addWidget(self._stop_btn)
layout.addWidget(start_group)
# === Control Buttons ===
controls_group = QGroupBox("Step Controls")
controls_layout = QVBoxLayout(controls_group)
# Step controls row 1
step_layout = QHBoxLayout()
self._next_btn = QPushButton("Next Step")
self._next_btn.setStyleSheet("background-color: #4CAF50; color: white;")
self._next_btn.setToolTip("Execute the next step (guided mode only)")
self._next_btn.clicked.connect(self.request_next_step.emit)
self._skip_btn = QPushButton("Skip Step")
self._skip_btn.setStyleSheet("background-color: #FF9800; color: white;")
self._skip_btn.setToolTip("Skip the current step and move to the next")
self._skip_btn.clicked.connect(self.request_skip.emit)
self._skip_sample_btn = QPushButton("Skip Sample")
self._skip_sample_btn.setStyleSheet("background-color: #FF5722; color: white;")
self._skip_sample_btn.setToolTip("Skip the entire current sample and move to the next")
self._skip_sample_btn.clicked.connect(self.request_skip_sample.emit)
step_layout.addWidget(self._next_btn)
step_layout.addWidget(self._skip_btn)
step_layout.addWidget(self._skip_sample_btn)
controls_layout.addLayout(step_layout)
# Step controls row 2
control_layout2 = QHBoxLayout()
self._pause_btn = QPushButton("⏸ Pause")
self._pause_btn.setStyleSheet("background-color: #2196F3; color: white;")
self._pause_btn.clicked.connect(self._on_pause_resume)
self._abort_btn = QPushButton("🛑 Abort")
self._abort_btn.setStyleSheet("background-color: #F44336; color: white;")
self._abort_btn.clicked.connect(self.request_abort.emit)
control_layout2.addWidget(self._pause_btn)
control_layout2.addWidget(self._abort_btn)
controls_layout.addLayout(control_layout2)
layout.addWidget(controls_group)
# === Queue Section ===
queue_group = QGroupBox("Queue (drag samples here)")
queue_layout = QVBoxLayout(queue_group)
self._drop_hint = QLabel("💡 Drag samples from the Sample List to add them")
self._drop_hint.setStyleSheet("color: #666; font-style: italic;")
queue_layout.addWidget(self._drop_hint)
self._queue_list = DraggableQueueListWidget()
self._queue_list.setMinimumHeight(150)
self._queue_list.samples_dropped.connect(self._on_samples_dropped)
self._queue_list.delete_requested.connect(self._on_delete_requested)
self._queue_list.item_reordered.connect(self._on_item_reordered)
self._queue_list.start_item_requested.connect(self._on_start_specific_item)
queue_layout.addWidget(self._queue_list)
# Queue management buttons
queue_btn_layout = QHBoxLayout()
self._add_all_btn = QPushButton(" Add All")
self._add_all_btn.clicked.connect(self._on_add_all_clicked)
self._add_all_btn.setToolTip("Add all samples from the Sample List")
self._remove_selected_btn = QPushButton("🗑 Remove")
self._remove_selected_btn.clicked.connect(self._on_remove_selected)
self._clear_btn = QPushButton("✖ Clear")
self._clear_btn.clicked.connect(self._on_clear_queue)
self._refresh_btn = QPushButton("🔄")
self._refresh_btn.setFixedWidth(40)
self._refresh_btn.setToolTip("Refresh queue")
self._refresh_btn.clicked.connect(self.request_queue_refresh.emit)
queue_btn_layout.addWidget(self._add_all_btn)
queue_btn_layout.addWidget(self._remove_selected_btn)
queue_btn_layout.addWidget(self._clear_btn)
queue_btn_layout.addStretch()
queue_btn_layout.addWidget(self._refresh_btn)
queue_layout.addLayout(queue_btn_layout)
layout.addWidget(queue_group)
# Initial state
self._update_button_states()
def _on_start_guided(self) -> None:
"""Start guided mode with the first pending sample."""
pending = [i for i in self._queue_items if i.status == QueueItemStatus.PENDING]
if not pending:
QMessageBox.information(self, "No Samples", "No pending samples in the queue.")
return
# Start the first pending item
self.request_start_item.emit(pending[0].item_id)
def _on_start_specific_item(self, item_id: str) -> None:
"""Start guided mode with a specific sample."""
self.request_start_item.emit(item_id)
def _on_start_automation(self) -> None:
"""Start full automation mode."""
pending = [i for i in self._queue_items if i.status == QueueItemStatus.PENDING]
is_running = self._runtime is not None and self._runtime.running
if not pending and not is_running:
QMessageBox.information(self, "No Samples", "No pending samples in the queue.")
return
# Remember if we were paused before starting automation
self._was_paused_before_automation = (
self._control is not None and self._control.pause_requested
)
self._automation_enabled = True
self.request_start_automation.emit()
self._update_button_states()
def _on_stop(self) -> None:
"""Stop automation mode."""
self._automation_enabled = False
self.request_stop_automation.emit()
# Restore pause state if it was paused before automation started
if self._was_paused_before_automation:
self.request_pause.emit()
self._update_button_states()
def _on_pause_resume(self) -> None:
"""Toggle pause/resume."""
if self._runtime and self._runtime.paused:
self.request_resume.emit()
else:
self.request_pause.emit()
def _on_samples_dropped(self, samples: list[SampleShortInfo]) -> None:
"""Handle samples dropped onto the queue."""
logger.info(f"Adding {len(samples)} samples to workflow queue")
for sample in samples:
self.request_add_sample.emit(sample)
QTimer.singleShot(500, self.request_queue_refresh.emit)
def _on_delete_requested(self, item_ids: list[str]) -> None:
"""Handle delete request from queue list."""
for item_id in item_ids:
self.request_delete_item.emit(item_id)
QTimer.singleShot(300, self.request_queue_refresh.emit)
def _on_item_reordered(self, item_id: str, new_index: int) -> None:
"""Handle item reorder."""
self.request_move_item.emit(item_id, new_index)
QTimer.singleShot(300, self.request_queue_refresh.emit)
def _on_remove_selected(self) -> None:
"""Remove selected items from queue."""
selected = self._queue_list.selectedItems()
if not selected:
return
item_ids = self._queue_list._item_ids
for item in selected:
row = self._queue_list.row(item)
if 0 <= row < len(item_ids):
self.request_delete_item.emit(item_ids[row])
QTimer.singleShot(300, self.request_queue_refresh.emit)
def _on_clear_queue(self) -> None:
"""Clear all items from queue."""
if not self._queue_items:
return
reply = QMessageBox.question(
self,
"Clear Queue",
"Remove all non-running items from the queue?",
QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No,
)
if reply == QMessageBox.StandardButton.Yes:
self.request_clear_queue.emit()
QTimer.singleShot(200, self._clear_local_queue)
QTimer.singleShot(500, self.request_queue_refresh.emit)
def _clear_local_queue(self) -> None:
"""Clear local queue state (called after server clear)."""
# Keep only running items
self._queue_items = [i for i in self._queue_items if i.status == QueueItemStatus.RUNNING]
# Update GUI
self._queue_list.clear()
item_ids = []
for item in self._queue_items:
status_emoji = "▶️"
display_text = f"{status_emoji} {item.sample_name or item.item_id}"
list_item = QListWidgetItem(display_text)
list_item.setBackground(QColor("#E6F3FF"))
self._queue_list.addItem(list_item)
item_ids.append(item.item_id)
self._queue_list.set_item_ids(item_ids)
self._update_button_states()
def _on_add_all_clicked(self) -> None:
"""Add all available samples to the queue."""
if not self._all_samples:
QMessageBox.information(
self,
"No Samples",
"No samples available to add. Load samples in the Sample List first.",
)
return
reply = QMessageBox.question(
self,
"Add All Samples",
f"Add all {len(self._all_samples)} samples to the workflow queue?",
QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No,
)
if reply == QMessageBox.StandardButton.Yes:
self.request_add_samples.emit(self._all_samples)
QTimer.singleShot(500, self.request_queue_refresh.emit)
def _update_button_states(self) -> None:
"""Update button enabled/disabled states based on current state."""
is_running = self._runtime is not None and self._runtime.running
is_paused = self._runtime is not None and self._runtime.paused
has_pending = any(i.status == QueueItemStatus.PENDING for i in self._queue_items)
self._start_guided_btn.setVisible(not is_running and not self._automation_enabled)
self._start_auto_btn.setVisible(not self._automation_enabled)
self._stop_btn.setVisible(self._automation_enabled)
self._start_guided_btn.setEnabled(has_pending)
self._start_auto_btn.setEnabled(has_pending or is_running)
self._next_btn.setVisible(not self._automation_enabled)
self._next_btn.setEnabled(is_running and not is_paused)
# Skip buttons - enabled when running
self._skip_btn.setEnabled(is_running)
self._skip_sample_btn.setEnabled(is_running)
self._abort_btn.setEnabled(is_running)
# Pause/Resume button
if is_paused:
self._pause_btn.setText("▶ Resume")
self._pause_btn.setStyleSheet("background-color: #4CAF50; color: white;")
else:
self._pause_btn.setText("⏸ Pause")
self._pause_btn.setStyleSheet("background-color: #2196F3; color: white;")
self._pause_btn.setEnabled(is_running)
# Drop hint
has_items = len(self._queue_items) > 0
self._drop_hint.setVisible(not has_items)
# === Public Update Methods ===
@Slot(list)
def update_queue(self, items: list[QueueItem]) -> None:
"""Update queue display."""
self._queue_items = items
self._queue_list.clear()
item_ids = []
for item in items:
status_emoji = {
QueueItemStatus.PENDING: "",
QueueItemStatus.RUNNING: "▶️",
QueueItemStatus.COMPLETED: "",
QueueItemStatus.FAILED: "",
QueueItemStatus.ABORTED: "🛑",
QueueItemStatus.SKIPPED: "⏭️",
}.get(item.status, "")
display_text = f"{status_emoji} {item.sample_name or item.item_id}"
list_item = QListWidgetItem(display_text)
if item.status == QueueItemStatus.RUNNING:
list_item.setBackground(QColor("#E6F3FF"))
elif item.status == QueueItemStatus.COMPLETED:
list_item.setBackground(QColor("#E6FFE6"))
elif item.status == QueueItemStatus.FAILED:
list_item.setBackground(QColor("#FFE6E6"))
self._queue_list.addItem(list_item)
item_ids.append(item.item_id)
self._queue_list.set_item_ids(item_ids)
self._update_button_states()
@Slot(object)
def update_runtime(self, runtime: RuntimeState) -> None:
"""Update from runtime state."""
self._runtime = runtime
if runtime.running:
if runtime.paused:
self._status_label.setText("⏸️ Paused")
self._status_label.setStyleSheet(
"font-size: 14px; font-weight: bold; color: #FF9800;"
)
else:
mode_str = "Automation" if self._automation_enabled else "Guided"
self._status_label.setText(f"▶️ Running ({mode_str})")
self._status_label.setStyleSheet(
"font-size: 14px; font-weight: bold; color: #4CAF50;"
)
self._current_item_label.setText(
f"Item: {runtime.current_item_id or 'Unknown'}"
)
else:
self._status_label.setText("⏹️ Idle")
self._status_label.setStyleSheet(
"font-size: 14px; font-weight: bold; color: #757575;"
)
self._current_item_label.setText("No item running")
self._step_progress.clear_steps()
# Reset automation enabled when nothing is running
if self._automation_enabled and not runtime.running:
# Check if queue has more pending items
pending = [i for i in self._queue_items if i.status == QueueItemStatus.PENDING]
if not pending:
self._automation_enabled = False
self._update_button_states()
@Slot(object)
def update_control(self, control: ControlState) -> None:
"""Update from control state."""
self._control = control
self._update_button_states()
@Slot(object)
def update_current_item(self, item: QueueItem) -> None:
"""Update step progress for current item."""
self._step_progress.set_steps(item.steps, item.current_step_index)
for i, step in enumerate(item.steps):
self._step_progress.update_step(i, step.status)
@Slot(bool)
def update_automation_enabled(self, enabled: bool) -> None:
"""Update automation mode state."""
self._automation_enabled = enabled
self._update_button_states()
@Slot(list)
def update_available_samples(self, samples: list[SampleShortInfo]) -> None:
"""Update the cached sample list for 'Add All' functionality."""
self._all_samples = samples
@Slot(str, str)
def on_workflow_event(self, event_type: str, message: str) -> None:
"""Handle workflow events from SSE stream."""
logger.debug(f"Workflow event: {event_type} - {message}")
if event_type in ("item_started", "item_completed", "step_completed", "step_skipped", "sample_skipped"):
self.request_queue_refresh.emit()
# Detect automation stop
if event_type == "automation_stopped":
self._automation_enabled = False
# Restore pause state if it was paused before automation started
if self._was_paused_before_automation:
self.request_pause.emit()
self._was_paused_before_automation = False
self._update_button_states()
+1 -1
View File
@@ -30,7 +30,7 @@ class FaceDetectionPanel(QWidget):
self.fig = Figure(figsize=(5, 4))
self.fig.subplots_adjust(
left=0.12,
left=0.18,
right=0.97,
bottom=0.10,
top=0.95,
@@ -14,8 +14,6 @@ class LoopCenteringPanel(QWidget):
self.find_tip = QPushButton("Center", parent=self)
grid_layout.addWidget(self.find_tip, 1, 0)
self.background = QPushButton("Bkg", parent=self)
grid_layout.addWidget(self.background, 1, 1)
self.bounding_box = QPushButton("Box", parent=self)
grid_layout.addWidget(self.bounding_box, 1, 2)
@@ -507,6 +507,7 @@ class SmargonTracePanel(QWidget):
project_root / self._csv_path,
project_root / "src" / "aare" / "daq" / "logs" / "smargon_trace.csv",
project_root / "src" / "aare" / "gui" / "logs" / "smargon_trace.csv",
Path("/sls/mx/applications/logs/smargon_trace.csv"),
]
out: list[Path] = []
+559 -46
View File
@@ -8,11 +8,14 @@ from PySide6.QtCore import Signal, QUrl, Slot, QTimer, QObject, QByteArray
from PySide6.QtNetwork import QNetworkAccessManager, QNetworkRequest, QNetworkReply
from jfjoch_client import ScanResult, ScanResultImagesInner
from aare.common.coordinate import SmargonCoordinate, Coordinate, AerotechCoordinate
from aare.common.auth_models import BatonStatus
from aare.common.coordinate import SmargonCoordinate, Coordinate
from aare.common.error_codes import export_error_codes
from aare.common.exception_handler import JFJochCommunicationError
from aare.common.models import DAQStatusModel, SampleShortInfoList, SampleShortInfo, SampleCameraSettings, \
AutofocusSettings, SimpleScanParameters, FluorescenceSpectrumParameterModel, FluorescenceSpectrumOutputModel
from aare.common.raster_grid import RasterGridRequest, CompletedRasterGrid
from aare.common.raster_grid import RasterGridRequest, CompletedRasterGrid, CompletedRasterGridElem
from aare.common.rotation_scan import RotationScanRequest, CompletedRotationScan
from aare.common.logger_config import setup_logger
@@ -28,8 +31,13 @@ class DAQWorker(QObject):
reference_tools = Signal(SampleShortInfoList)
http_error = Signal(str)
status_message = Signal(str, bool)
# New dedicated signals for polled device errors and request-time errors
polled_devices_status = Signal(str, bool) # (message, is_error)
detector_error = Signal(str, bool) # (message, is_error)
auth_error = Signal()
sample_missing = Signal(str)
automated_scan_done = Signal(int, bool, str) # sample ID, success
run_number_incremented = Signal()
raster_scan_completed = Signal(CompletedRasterGrid)
@@ -45,8 +53,23 @@ class DAQWorker(QObject):
last_error_payload_changed = Signal(dict)
last_error_payloads_changed = Signal(list)
# Workflow signals
workflow_queue_loaded = Signal(list) # list[QueueItem]
workflow_runtime_changed = Signal(object) # RuntimeState
workflow_control_changed = Signal(object) # ControlState
workflow_item_updated = Signal(object) # QueueItem
workflow_automation_status = Signal(bool) # enabled
workflow_event = Signal(object) # WorkflowEvent
baton_status_changed = Signal(BatonStatus)
baton_request_result = Signal(dict)
baton_response_result = Signal(dict)
baton_incoming_request = Signal(dict)
baton_timeout_checked = Signal(dict)
def __init__(self, base_url: str | None, token: str, parent=None):
super().__init__(parent)
self._active_status_error_key = None
self.__token = token
self.__base_url = base_url
self.__net_manager = QNetworkAccessManager()
@@ -77,14 +100,37 @@ class DAQWorker(QObject):
self._last_error_payloads = deque(maxlen=10)
self._last_tell_connected: bool | None = None
self._last_tell_error: str | None = None
self._last_smargon_connected: bool | None = None
self._last_smargon_error: str | None = None
self._last_aerotech_connected: bool | None = None
self._last_aerotech_error: str | None = None
self._server_connected: bool | None = None
self._last_server_error: str | None = None
self._has_seen_disconnection: bool = False
self._server_was_disconnected: bool = False
# Deduplication tracking for polled device messages
self._last_polled_msg: str | None = None
self._last_polled_is_error: bool | None = None
# Deduplication tracking for detector/request-time messages
self._last_detector_msg: str | None = None
self._last_detector_is_error: bool | None = None
self._baton_stream_reply: QNetworkReply | None = None
self._last_baton_status: BatonStatus | None = None
self._baton_timeout_timer = QTimer(self)
self._baton_timeout_timer.setInterval(1000)
self._baton_timeout_timer.timeout.connect(self.check_baton_timeout)
self._face_detection_stream_reply: QNetworkReply | None = None
if self.__base_url is not None:
self.start_face_detection_stream()
self.start_baton_stream()
def get_last_error_payload(self) -> dict:
return dict(self._last_error_payload or {})
@@ -100,6 +146,46 @@ class DAQWorker(QObject):
self.last_error_payload_changed.emit(self.get_last_error_payload())
self.last_error_payloads_changed.emit(self.get_last_error_payloads())
def _emit_polled_device_message(self, msg: str, is_error: bool) -> None:
"""Emit polled device status message with deduplication."""
if (msg, is_error) == (self._last_polled_msg, self._last_polled_is_error):
return
self._last_polled_msg = msg
self._last_polled_is_error = is_error
self.polled_devices_status.emit(msg, is_error)
def _emit_detector_message(self, msg: str, is_error: bool) -> None:
"""Emit detector/request-time error message with deduplication."""
if (msg, is_error) == (self._last_detector_msg, self._last_detector_is_error):
return
self._last_detector_msg = msg
self._last_detector_is_error = is_error
self.detector_error.emit(msg, is_error)
def _emit_status_if_changed(self, key: str | None, message: str | None, is_error: bool) -> None:
"""
Emit polled device status to the primary alert banner.
Used for Server/Tell/Smargon/Aerotech connection status.
"""
if not message:
self._active_status_error_key = None
self._emit_polled_device_message("", False)
return
if is_error:
self._has_seen_disconnection = True
if self._active_status_error_key != key:
self._active_status_error_key = key
self._emit_polled_device_message(message, True)
return
if not self._has_seen_disconnection:
self._active_status_error_key = None
return
self._active_status_error_key = None
self._emit_polled_device_message(message, False)
def _log_smargon_throttled(self, *, endpoint: str | None, message: str) -> None:
"""
Log immediately if endpoint/message changed; otherwise at most every N seconds.
@@ -175,27 +261,15 @@ class DAQWorker(QObject):
logger.error(f"Error in response: {reply.errorString()}")
raise RuntimeError(reply.errorString())
def _emit_status_if_changed(self, key: str | None, message: str | None, is_error: bool) -> None:
if not message:
self._active_status_error_key = None
return
if is_error:
if self._active_status_error_key != key:
self._active_status_error_key = key
self.status_message.emit(message, True)
return
self._active_status_error_key = None
self.status_message.emit(message, False)
def _compose_device_status_message(
self,
*,
tell_conn: bool,
smargon_conn: bool,
aerotech_conn: bool,
tell_changed: bool,
smargon_changed: bool,
aerotech_changed: bool,
) -> tuple[str | None, str | None, bool]:
disconnected: list[str] = []
restored: list[str] = []
@@ -210,25 +284,37 @@ class DAQWorker(QObject):
elif smargon_changed:
restored.append("Smargon")
if not aerotech_conn:
disconnected.append("Aerotech")
elif aerotech_changed:
restored.append("Aerotech")
if disconnected:
if len(disconnected) == 3:
return (
"all-devices-down",
"TELL, Smargon, and Aerotech disconnected.",
True,
)
if len(disconnected) == 2:
return (
"tell+smargon-down",
"TELL and Smargon connection errors, please inform your local contact.",
"+".join(sorted(d.lower() for d in disconnected)) + "-down",
f"{disconnected[0]} and {disconnected[1]} disconnected.",
True,
)
device = disconnected[0]
return (
f"{device.lower()}-down",
f"{device} connection error, please inform your local contact.",
f"{device} disconnected.",
True,
)
if restored:
if len(restored) == 3:
return (None, "TELL, Smargon, and Aerotech reconnected.", False)
if len(restored) == 2:
return (None, "TELL and Smargon connections restored.", False,)
device = restored[0]
return (None, f"{device} connection restored.",False)
return (None, f"{restored[0]} and {restored[1]} reconnected.", False)
return (None, f"{restored[0]} reconnected.", False)
return (None, None, False)
@@ -239,29 +325,68 @@ class DAQWorker(QObject):
parsed_response = DAQStatusModel.model_validate_json(response_data)
self.update.emit(parsed_response)
if self._server_connected is False:
self._emit_status_if_changed(None, "Server connection restored.", False)
# Handle server reconnection - show "Server reconnected" not device messages
if self._server_connected is False and self._has_seen_disconnection:
self._emit_status_if_changed(None, "Server reconnected.", False)
# Reset device states so we don't also emit device reconnection messages
self._last_tell_connected = None
self._last_smargon_connected = None
self._last_aerotech_connected = None
self._server_was_disconnected = False
self._server_connected = True
self._last_server_error = None
tell_conn = bool(getattr(parsed_response, "tell_connected", True))
smargon_conn = bool(getattr(parsed_response, "smargon_connected", True))
tell_err = getattr(parsed_response, "tell_error", None)
smargon_err = getattr(parsed_response, "smargon_error", None)
tell_conn = bool(getattr(parsed_response, "tell_connected", True))
tell_err = getattr(parsed_response, "tell_error", None)
aerotech_conn = bool(getattr(parsed_response, "aerotech_connected", True))
aerotech_err = getattr(parsed_response, "aerotech_error", None)
tell_changed = self._last_tell_connected is not None and self._last_tell_connected != tell_conn
tell_err_text = None if tell_err is None else str(tell_err).strip()
smargon_err_text = None if smargon_err is None else str(smargon_err).strip()
aerotech_err_text = None if aerotech_err is None else str(aerotech_err).strip()
# Skip device status processing if we just reconnected from server down
# (we already showed "Server reconnected")
if self._last_tell_connected is None and self._last_smargon_connected is None and self._last_aerotech_connected is None:
# First status after startup or server reconnect - just record states, don't emit
self._last_tell_connected = tell_conn
self._last_smargon_connected = smargon_conn
self._last_aerotech_connected = aerotech_conn
self._last_tell_error = tell_err_text
self._last_smargon_error = smargon_err_text
self._last_aerotech_error = aerotech_err_text
return
tell_changed = (
self._last_tell_connected != tell_conn
or self._last_tell_error != tell_err_text
)
smargon_changed = (
self._last_smargon_connected is not None and self._last_smargon_connected != smargon_conn
self._last_smargon_connected != smargon_conn
or self._last_smargon_error != smargon_err_text
)
aerotech_changed = (
self._last_aerotech_connected != aerotech_conn
or self._last_aerotech_error != aerotech_err_text
)
self._last_tell_connected = tell_conn
self._last_smargon_connected = smargon_conn
self._last_aerotech_connected = aerotech_conn
self._last_tell_error = tell_err_text
self._last_smargon_error = smargon_err_text
self._last_aerotech_error = aerotech_err_text
status_key, status_msg, is_error = self._compose_device_status_message(
tell_conn=tell_conn,
smargon_conn=smargon_conn,
aerotech_conn=aerotech_conn,
tell_changed=tell_changed,
smargon_changed=smargon_changed,
aerotech_changed=aerotech_changed,
)
self._emit_status_if_changed(status_key, status_msg, is_error)
@@ -269,14 +394,18 @@ class DAQWorker(QObject):
self._log_device_error_throttled(device="tell", message=tell_err)
if not smargon_conn:
self._log_device_error_throttled(device="smargon", message=smargon_err)
if not aerotech_conn:
self._log_device_error_throttled(device="aerotech", message=aerotech_err)
except Exception as e:
err_msg = str(e)
if self._server_connected is not False:
self._has_seen_disconnection = True
self._server_was_disconnected = True
self._emit_status_if_changed(
"server-down",
"Server connection lost. Trying to reconnect...",
"Server disconnected. Reconnecting...",
True,
)
@@ -284,6 +413,7 @@ class DAQWorker(QObject):
self._last_server_error = err_msg
self._last_tell_connected = None
self._last_smargon_connected = None
self._last_aerotech_connected = None
logger.error(f"Exception from status response: {e}")
@@ -449,10 +579,6 @@ class DAQWorker(QObject):
def open_shutter(self):
self.generic_post(f"beamline/shutter?val=true")
@Slot()
def alc_background(self):
self.generic_post("alc/background")
@Slot()
def center_loop(self):
self.generic_post("alc/center_loop")
@@ -547,6 +673,7 @@ class DAQWorker(QObject):
self.generic_delete("access/pgroup")
else:
self.generic_put(f"access/pgroup?val={val}")
self.send_status_request()
@Slot(QNetworkReply)
def _handle_all_pgroups_response(self, reply: QNetworkReply):
@@ -589,12 +716,38 @@ class DAQWorker(QObject):
def handle_rotation_scan_response(self, reply: QNetworkReply):
try:
# Check for HTTP errors first
if reply.error() != QNetworkReply.NetworkError.NoError:
status = reply.attribute(QNetworkRequest.Attribute.HttpStatusCodeAttribute)
err_msg = reply.errorString()
try:
raw_body = reply.readAll().data().decode("utf-8")
if raw_body:
body_json = json.loads(raw_body)
if isinstance(body_json, dict):
err_msg = body_json.get("message", err_msg)
code = body_json.get("code", "")
# Check if this is a JFJoch error
if code == "JFJOCH_UNAVAILABLE" or status == 503:
self._emit_detector_message(f"JFJoch: {err_msg}", is_error=True)
except Exception:
pass
logger.error(f"Rotation scan failed: {err_msg}")
self.http_error.emit(err_msg)
reply.deleteLater()
return
response_data = self.handle_response(reply)
parsed_response = CompletedRotationScan.model_validate_json(response_data)
self.standard_scan_completed.emit(parsed_response)
except Exception as e:
logger.error(f"Exception from rotation scan response: {e}")
self.http_error.emit(str(e))
finally:
reply.deleteLater()
@Slot(RotationScanRequest)
def standard_scan(self, r: RotationScanRequest):
@@ -609,12 +762,40 @@ class DAQWorker(QObject):
def handle_raster_scan_response(self, reply: QNetworkReply):
try:
# Check for HTTP errors first
if reply.error() != QNetworkReply.NetworkError.NoError:
status = reply.attribute(QNetworkRequest.Attribute.HttpStatusCodeAttribute)
raw_body = ""
body_json = None
err_msg = reply.errorString()
try:
raw_body = reply.readAll().data().decode("utf-8")
if raw_body:
body_json = json.loads(raw_body)
if isinstance(body_json, dict):
err_msg = body_json.get("message", err_msg)
code = body_json.get("code", "")
# Check if this is a JFJoch error
if code == "JFJOCH_UNAVAILABLE" or status == 503:
self._emit_detector_message(f"JFJoch: {err_msg}", is_error=True)
except Exception:
pass
logger.error(f"Raster scan failed: {err_msg}")
self.http_error.emit(err_msg)
reply.deleteLater()
return
response_data = self.handle_response(reply)
parsed_response = CompletedRasterGrid.model_validate_json(response_data)
self.raster_scan_completed.emit(parsed_response)
except Exception as e:
logger.error(f"Exception from raster scan response: {e}")
self.http_error.emit(str(e))
finally:
reply.deleteLater()
@Slot(RasterGridRequest)
def raster_scan(self, r: RasterGridRequest):
@@ -633,13 +814,14 @@ class DAQWorker(QObject):
bkg = random.gauss(3.0, 0.1),
spots= random.randint(0, 250),
index= random.randint(0, 1),
mos = random.uniform(0, 0.1),
b= random.uniform(15.0, 80.0)
))
logger.debug("check that this works - raster scan - complete raster grid")
reply = CompletedRasterGrid(request = new_copy,
result = ScanResult(file_prefix=r.file_prefix, images=images))
logger.debug(f"It appears to work {reply}")
raster_elem = CompletedRasterGridElem(
request=new_copy,
result=ScanResult(file_prefix=r.file_prefix, images=images),
centre_of_mass = None,
)
reply = CompletedRasterGrid(r=[raster_elem])
self.raster_scan_completed.emit(reply)
return
@@ -666,13 +848,14 @@ class DAQWorker(QObject):
bkg=random.gauss(3.0, 0.1),
spots=random.randint(0, 250),
index=random.randint(0, 1),
mos=random.uniform(0, 0.1),
b=random.uniform(15.0, 80.0)
))
logger.debug("check that this works - raster scan auto - complete raster grid")
reply = CompletedRasterGrid(request=new_copy,
result=ScanResult(file_prefix=r.file_prefix, images=images))
logger.debug(f"It appears to work {reply}")
raster_elem = CompletedRasterGridElem(
request=new_copy,
result=ScanResult(file_prefix=r.file_prefix, images=images),
centre_of_mass = None,
)
reply = CompletedRasterGrid(r=[raster_elem])
self.raster_scan_completed.emit(reply)
return
@@ -763,8 +946,8 @@ class DAQWorker(QObject):
return
self.generic_post("scan/smart_params", p.model_dump_json())
@Slot(Coordinate)
def abr_tweak(self, c: Coordinate):
@Slot(AerotechCoordinate)
def abr_tweak(self, c: AerotechCoordinate):
self.generic_post("beamline/tweak_abr_meas_pos", c.model_dump_json())
@Slot()
@@ -811,6 +994,10 @@ class DAQWorker(QObject):
def mount(self, s: SampleShortInfo, reference: bool = False):
self.generic_post(f"sample/mount?dbid={s.db_id}&reference={reference}")
@Slot()
def park_and_dry(self):
self.generic_post("sample/park_and_dry")
@Slot(SampleShortInfo)
def sample_manual(self, s: SampleShortInfo):
self.generic_post(f"sample/manual", s.model_dump_json())
@@ -1022,7 +1209,6 @@ class DAQWorker(QObject):
out[str(k)] = str(v)
return out
@Slot()
@Slot()
def get_error_codes(self) -> None:
"""
@@ -1089,4 +1275,331 @@ class DAQWorker(QObject):
query.append(f"message={quote(message)}")
suffix = f"?{'&'.join(query)}" if query else ""
self.generic_post(f"samcam/send_screenshot_db{suffix}")
self.generic_post(f"samcam/send_screenshot_db{suffix}")
def start_baton_stream(self):
"""Start SSE stream for baton status updates."""
if self.__base_url is None:
return
if self._baton_stream_reply is not None:
return
request = QNetworkRequest(QUrl(f"{self.__base_url}/sse/baton"))
request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8"))
reply = self.__net_manager.get(request)
reply.readyRead.connect(lambda: self._read_baton_stream(reply))
reply.finished.connect(self._restart_baton_stream)
self._baton_stream_reply = reply
def _restart_baton_stream(self):
self._baton_stream_reply = None
if self.__base_url is not None:
QTimer.singleShot(1000, self.start_baton_stream)
def _read_baton_stream(self, reply: QNetworkReply):
try:
chunk = reply.readAll().data().decode("utf-8")
for line in chunk.splitlines():
if line.startswith("data:"):
payload = line[5:].strip()
if payload:
status = BatonStatus.model_validate_json(payload)
if status.you_have_pending_request:
if not self._baton_timeout_timer.isActive():
self._baton_timeout_timer.start()
else:
if self._baton_timeout_timer.isActive():
self._baton_timeout_timer.stop()
if (status.incoming_request and
(self._last_baton_status is None or
not self._last_baton_status.incoming_request)):
self.baton_incoming_request.emit({
"requester": status.pending_request.requester_username if status.pending_request else "Unknown",
"timeout": status.pending_request.timeout_seconds if status.pending_request else 30
})
self._last_baton_status = status
self.baton_status_changed.emit(status)
except Exception as e:
logger.error(f"Baton stream parse error: {e}")
@Slot()
def request_baton(self):
"""Request the baton."""
if self.__base_url is None:
logger.info("POST /baton/request")
return
request = QNetworkRequest(QUrl(f"{self.__base_url}/baton/request"))
request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8"))
request.setRawHeader(b"Content-Type", b"application/json")
reply = self.__net_manager.post(request, QByteArray(b""))
reply.finished.connect(lambda: self._handle_baton_request_response(reply))
def _handle_baton_request_response(self, reply: QNetworkReply):
try:
response_data = self.handle_response(reply)
result = json.loads(response_data) if response_data else {}
self.baton_request_result.emit(result)
if result.get("granted"):
self.status_message.emit("Baton acquired", False)
self.send_status_request()
elif result.get("pending"):
self.status_message.emit(
f"Request sent - waiting for response ({result.get('timeout_seconds', 30)}s timeout)",
False
)
elif result.get("error"):
self.status_message.emit(result.get("message", "Request failed"), True)
except Exception as e:
logger.error(f"Baton request failed: {e}")
self.http_error.emit(str(e))
@Slot(bool)
def respond_to_baton_request(self, accept: bool):
"""Respond to an incoming baton request."""
if self.__base_url is None:
logger.info(f"POST /baton/respond?accept={accept}")
return
request = QNetworkRequest(QUrl(f"{self.__base_url}/baton/respond?accept={str(accept).lower()}"))
request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8"))
request.setRawHeader(b"Content-Type", b"application/json")
reply = self.__net_manager.post(request, QByteArray(b""))
reply.finished.connect(lambda: self._handle_baton_response_result(reply))
def _handle_baton_response_result(self, reply: QNetworkReply):
try:
response_data = self.handle_response(reply)
result = json.loads(response_data) if response_data else {}
self.baton_response_result.emit(result)
if result.get("accepted") or result.get("refused"):
self.send_status_request()
self.check_baton_timeout()
self.start_baton_stream()
except Exception as e:
logger.error(f"Baton response failed: {e}")
self.http_error.emit(str(e))
@Slot()
def release_baton(self):
"""Release the baton voluntarily."""
self.generic_post("baton/release")
@Slot()
def cancel_baton_request(self):
"""Cancel your pending baton request."""
self.generic_post("baton/cancel")
@Slot()
def check_baton_timeout(self):
"""Poll to check if timeout has been reached."""
if self.__base_url is None:
return
request = QNetworkRequest(QUrl(f"{self.__base_url}/baton/check_timeout"))
request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8"))
reply = self.__net_manager.get(request)
reply.finished.connect(lambda: self._handle_baton_timeout_response(reply))
def _handle_baton_timeout_response(self, reply: QNetworkReply):
try:
response_data = self.handle_response(reply)
result = json.loads(response_data) if response_data else {}
self.baton_timeout_checked.emit(result)
if result.get("granted") or result.get("queued") or result.get("refused"):
self.send_status_request()
self.start_baton_stream()
except Exception as e:
logger.error(f"Baton timeout check failed: {e}")
self.http_error.emit(str(e))
def release_baton_on_close(self):
"""Release baton when GUI is closed to free the beamline."""
try:
if hasattr(self, "_baton_timeout_timer") and self._baton_timeout_timer is not None:
self._baton_timeout_timer.stop()
self.release_baton()
from PySide6.QtCore import QEventLoop, QTimer
loop = QEventLoop()
QTimer.singleShot(500, loop.quit)
loop.exec()
logger.info("Baton release requested on GUI close")
except Exception as e:
logger.warning(f"Error releasing baton on close: {e}")
# ─────────────────────────────────────────────
# Workflow API methods
# ─────────────────────────────────────────────
@Slot()
def workflow_load_queue(self):
"""Load workflow queue."""
if self.__base_url is None:
return
request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/queue"))
request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode())
reply = self.__net_manager.get(request)
reply.finished.connect(lambda: self._handle_workflow_queue_response(reply))
def _handle_workflow_queue_response(self, reply: QNetworkReply):
try:
response_data = self.handle_response(reply)
data = json.loads(response_data)
from aare.common.automation_models import QueueItem, QueueItemStatus
items = [QueueItem.model_validate(i) for i in data.get("items", [])]
self.workflow_queue_loaded.emit(items)
# Also emit the currently running item for step progress
for item in items:
if item.status == QueueItemStatus.RUNNING:
self.workflow_item_updated.emit(item)
break
except Exception as e:
logger.error(f"Failed to load workflow queue: {e}")
@Slot(object) # SampleShortInfo
def workflow_add_sample(self, sample: "SampleShortInfo"):
"""Add a sample to the workflow queue."""
if self.__base_url is None:
logger.info(f"POST /workflow/queue: {sample.sample_name}")
return
from aare.common.automation_models import CreateQueueItemRequest
request_data = CreateQueueItemRequest(
sample_id=sample.db_id,
sample_name=sample.sample_name,
priority=int(sample.priority or 100),
)
request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/queue"))
request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode())
request.setRawHeader(b"Content-Type", b"application/json")
reply = self.__net_manager.post(request, QByteArray(request_data.model_dump_json().encode()))
reply.finished.connect(lambda: self.handle_req_response(reply))
@Slot(list) # list[SampleShortInfo]
def workflow_add_samples(self, samples: list):
"""Add multiple samples to the workflow queue."""
for sample in samples:
self.workflow_add_sample(sample)
@Slot(str)
def workflow_delete_item(self, item_id: str):
"""Delete an item from the workflow queue."""
if self.__base_url is None:
logger.info(f"DELETE /workflow/queue/{item_id}")
return
self.generic_delete(f"workflow/queue/{item_id}")
@Slot(str, int)
def workflow_move_item(self, item_id: str, new_order_index: int):
"""Move/reorder an item in the workflow queue."""
if self.__base_url is None:
logger.info(f"POST /workflow/queue/{item_id}/move order={new_order_index}")
return
body = json.dumps({"new_order_index": new_order_index})
self.generic_post(f"workflow/queue/{item_id}/move", body)
@Slot()
def workflow_clear_queue(self):
"""Clear all non-running items from the queue via server."""
if self.__base_url is None:
return
request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/queue/clear"))
request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode())
request.setRawHeader(b"Content-Type", b"application/json")
reply = self.__net_manager.deleteResource(request)
reply.finished.connect(lambda: self.handle_req_response(reply))
@Slot()
def workflow_skip_sample(self):
"""Skip the entire current sample."""
self.generic_post("workflow/control/skip_sample")
def _handle_clear_queue_response(self, reply: QNetworkReply):
try:
response_data = self.handle_response(reply)
data = json.loads(response_data)
from aare.common.automation_models import QueueItem, QueueItemStatus
items = [QueueItem.model_validate(i) for i in data.get("items", [])]
# Delete all non-running items
for item in items:
if item.status != QueueItemStatus.RUNNING:
self.workflow_delete_item(item.item_id)
except Exception as e:
logger.error(f"Failed to clear workflow queue: {e}")
@Slot(str)
def workflow_start_item(self, item_id: str):
"""Start processing a queue item."""
self.generic_post(f"workflow/start/{item_id}")
QTimer.singleShot(500, self.workflow_load_queue)
@Slot()
def workflow_next_step(self):
"""Run next step in guided mode."""
if self.__base_url is None:
return
request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/next"))
request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode())
request.setRawHeader(b"Content-Type", b"application/json")
reply = self.__net_manager.post(request, QByteArray(b""))
reply.finished.connect(lambda: self._handle_workflow_step_response(reply))
def _handle_workflow_step_response(self, reply: QNetworkReply):
try:
response_data = self.handle_response(reply)
data = json.loads(response_data)
if data.get("item"):
from aare.common.automation_models import QueueItem
item = QueueItem.model_validate(data["item"])
self.workflow_item_updated.emit(item)
# Also refresh queue
self.workflow_load_queue()
except Exception as e:
logger.error(f"Workflow step error: {e}")
self.http_error.emit(str(e))
@Slot()
def workflow_pause(self):
"""Request pause."""
self.generic_post("workflow/control/pause")
@Slot()
def workflow_resume(self):
"""Request resume."""
self.generic_post("workflow/control/resume")
@Slot()
def workflow_abort(self):
"""Request abort."""
self.generic_post("workflow/control/abort")
@Slot()
def workflow_skip(self):
"""Request skip."""
self.generic_post("workflow/control/skip")
@Slot()
def workflow_start_automation(self):
"""Start automation mode."""
self.generic_post("workflow/automation/start")
@Slot()
def workflow_stop_automation(self):
"""Stop automation mode."""
self.generic_post("workflow/automation/stop")
+175
View File
@@ -0,0 +1,175 @@
"""
SSE client for workflow events.
Connects to the /workflow/sse endpoint and emits signals for:
- Runtime state changes
- Control state changes
- Individual workflow events
"""
from __future__ import annotations
import json
from PySide6.QtCore import QObject, Signal, Slot, QTimer
from PySide6.QtNetwork import QNetworkAccessManager, QNetworkRequest, QNetworkReply
from PySide6.QtCore import QUrl, QByteArray
from aare.common.automation_models import RuntimeState, ControlState, WorkflowEvent
from aare.common.logger_config import setup_logger
logger = setup_logger("aareGUI")
class WorkflowSSEClient(QObject):
"""
SSE client that subscribes to workflow events.
Emits signals when state changes are received.
"""
# Signals
runtime_changed = Signal(object) # RuntimeState
control_changed = Signal(object) # ControlState
workflow_event = Signal(object) # WorkflowEvent
connected = Signal()
disconnected = Signal()
error = Signal(str)
def __init__(
self,
base_url: str,
token: str,
parent: QObject | None = None,
):
super().__init__(parent)
self._base_url = base_url
self._token = token
self._manager = QNetworkAccessManager(self)
self._reply: QNetworkReply | None = None
self._buffer = ""
# Reconnection
self._reconnect_timer = QTimer(self)
self._reconnect_timer.setInterval(5000) # 5 seconds
self._reconnect_timer.timeout.connect(self.connect)
self._should_reconnect = False
def connect(self) -> None:
"""Start SSE connection."""
if self._reply is not None:
return # Already connected
self._should_reconnect = True
url = QUrl(f"{self._base_url}/workflow/sse")
request = QNetworkRequest(url)
request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode())
request.setRawHeader(b"Accept", b"text/event-stream")
request.setRawHeader(b"Cache-Control", b"no-cache")
self._reply = self._manager.get(request)
self._reply.readyRead.connect(self._on_data_ready)
self._reply.finished.connect(self._on_finished)
self._reply.errorOccurred.connect(self._on_error)
self._reconnect_timer.stop()
logger.debug("Workflow SSE: connecting...")
def disconnect(self) -> None:
"""Stop SSE connection."""
self._should_reconnect = False
self._reconnect_timer.stop()
if self._reply is not None:
try:
self._reply.abort()
except Exception:
pass
try:
self._reply.deleteLater()
except Exception:
pass
self._reply = None
self.disconnected.emit()
@Slot()
def _on_data_ready(self) -> None:
"""Handle incoming SSE data."""
if self._reply is None:
return
try:
data = self._reply.readAll().data().decode("utf-8")
self._buffer += data
# Process complete events (separated by double newlines)
while "\n\n" in self._buffer:
event_data, self._buffer = self._buffer.split("\n\n", 1)
self._parse_event(event_data)
except Exception as e:
logger.warning(f"Workflow SSE data read error: {e}")
def _parse_event(self, event_data: str) -> None:
"""Parse a single SSE event."""
event_type = "message"
data_lines = []
for line in event_data.split("\n"):
if line.startswith("event:"):
event_type = line[6:].strip()
elif line.startswith("data:"):
data_lines.append(line[5:].strip())
if not data_lines:
return
data_str = "\n".join(data_lines)
try:
if event_type == "runtime":
runtime = RuntimeState.model_validate_json(data_str)
self.runtime_changed.emit(runtime)
elif event_type == "control":
control = ControlState.model_validate_json(data_str)
self.control_changed.emit(control)
elif event_type == "workflow_event":
event = WorkflowEvent.model_validate_json(data_str)
self.workflow_event.emit(event)
except Exception as e:
logger.warning(f"Workflow SSE: failed to parse {event_type}: {e}")
@Slot()
def _on_finished(self) -> None:
"""Handle connection finished."""
if self._reply is not None:
try:
self._reply.deleteLater()
except Exception:
pass
self._reply = None
self._buffer = ""
self.disconnected.emit()
# Reconnect if desired
if self._should_reconnect:
logger.debug("Workflow SSE: disconnected, will reconnect...")
self._reconnect_timer.start()
@Slot(QNetworkReply.NetworkError)
def _on_error(self, error: QNetworkReply.NetworkError) -> None:
"""Handle connection error."""
error_msg = ""
if self._reply is not None:
try:
error_msg = self._reply.errorString()
except Exception:
error_msg = str(error)
else:
error_msg = str(error)
logger.warning(f"Workflow SSE error: {error_msg}")
self.error.emit(error_msg)
+84 -2
View File
@@ -15,6 +15,13 @@ class AlertBanner(QFrame):
self._clear_timer.setSingleShot(True)
self._clear_timer.timeout.connect(self.clear_message)
# Countdown timer for "waiting" state
self._countdown_timer = QTimer(self)
self._countdown_timer.setInterval(1000)
self._countdown_timer.timeout.connect(self._tick_countdown)
self._countdown_remaining = 0
self._countdown_base_message = ""
self._label = QLabel("", self)
self._label.setWordWrap(True)
self._label.setAlignment(Qt.AlignmentFlag.AlignCenter)
@@ -33,7 +40,9 @@ class AlertBanner(QFrame):
self.setSizePolicy(QSizePolicy.Policy.Expanding, QSizePolicy.Policy.Fixed)
@Slot(str, bool)
def show_message(self, msg: str, is_error: bool = True):
def show_message(self, msg: str, is_error: bool = True, auto_clear_ms: int | None = None):
"""Show error (red) or success (green) message."""
self._stop_countdown()
self._clear_timer.stop()
if not msg:
@@ -72,13 +81,86 @@ class AlertBanner(QFrame):
" padding: 2px 6px 2px 6px;"
"}"
)
self._clear_timer.start(5000)
timeout = 5000 if auto_clear_ms is None else int(auto_clear_ms)
self._clear_timer.start(timeout)
self._current_message = msg
self._current_is_error = is_error
self._label.setText(decorated)
self.setVisible(True)
@Slot(str, int)
def show_waiting(self, msg: str, countdown_seconds: int = 0):
"""
Show a waiting/pending message (yellow) with optional countdown.
Args:
msg: Base message to display
countdown_seconds: If > 0, append countdown and auto-update
"""
self._clear_timer.stop()
self._stop_countdown()
if not msg:
self.clear_message()
return
self._countdown_base_message = msg
self._countdown_remaining = countdown_seconds
self._apply_waiting_style()
self._update_waiting_text()
if countdown_seconds > 0:
self._countdown_timer.start()
self.setVisible(True)
def _apply_waiting_style(self):
"""Apply yellow/waiting style."""
self.setStyleSheet(
"QFrame {"
" background-color: #fff8e1;"
" border: 2px solid #ffb300;"
" border-radius: 12px;"
" margin: 8px 12px 8px 12px;"
"}"
"QLabel {"
" color: #e65100;"
" font-weight: 700;"
" font-size: 20px;"
" padding: 2px 6px 2px 6px;"
"}"
)
def _update_waiting_text(self):
"""Update the waiting message text, including countdown if active."""
if self._countdown_remaining > 0:
decorated = f"{self._countdown_base_message} ({self._countdown_remaining}s) ⏳"
else:
decorated = f"{self._countdown_base_message}"
self._label.setText(decorated)
def _tick_countdown(self):
"""Called every second during countdown."""
self._countdown_remaining -= 1
if self._countdown_remaining <= 0:
self._stop_countdown()
self.clear_message()
return
self._update_waiting_text()
def _stop_countdown(self):
"""Stop the countdown timer."""
self._countdown_timer.stop()
self._countdown_remaining = 0
self._countdown_base_message = ""
@Slot()
def clear_message(self):
self._current_message = None
self._current_is_error = None
self._clear_timer.stop()
self._label.clear()
self.setVisible(False)
@@ -0,0 +1,348 @@
from PySide6.QtCore import Qt, Signal, QTimer
from PySide6.QtWidgets import (
QDialog, QVBoxLayout, QHBoxLayout, QLabel,
QPushButton, QProgressBar, QFrame
)
from PySide6.QtGui import QFont
class BatonRequestDialog(QDialog):
"""
Dialog shown to current baton holder when someone requests control.
Based on the workflow diagram:
- User can Accept (transfer immediately or queue if busy)
- User can Refuse (deny the request)
- If user ignores/closes, timeout causes auto-transfer
"""
accepted_signal = Signal()
refused_signal = Signal()
def __init__(self, requester: str, timeout_seconds: int = 30, parent=None):
super().__init__(parent)
self.setWindowTitle("⚡ Baton Request")
self.setModal(False) # Non-modal so user can see beamline status
self.setMinimumWidth(400)
self.setWindowFlags(
self.windowFlags() |
Qt.WindowType.WindowStaysOnTopHint
)
self._timeout = timeout_seconds
self._remaining = timeout_seconds
self._requester = requester
self._setup_ui()
self._start_timer()
def _setup_ui(self):
layout = QVBoxLayout(self)
layout.setSpacing(15)
# Header
header = QLabel("🔔 Control Request")
header_font = QFont()
header_font.setPointSize(14)
header_font.setBold(True)
header.setFont(header_font)
header.setAlignment(Qt.AlignmentFlag.AlignCenter)
layout.addWidget(header)
# Separator
line = QFrame()
line.setFrameShape(QFrame.Shape.HLine)
line.setFrameShadow(QFrame.Shadow.Sunken)
layout.addWidget(line)
# Message
self.message_label = QLabel(
f"<b>{self._requester}</b> is requesting control of the beamline."
)
self.message_label.setWordWrap(True)
self.message_label.setAlignment(Qt.AlignmentFlag.AlignCenter)
layout.addWidget(self.message_label)
# Timeout progress
progress_layout = QVBoxLayout()
self.progress = QProgressBar()
self.progress.setRange(0, self._timeout)
self.progress.setValue(self._timeout)
self.progress.setTextVisible(False)
self.progress.setFixedHeight(8)
self.progress.setStyleSheet("""
QProgressBar {
border: 1px solid #ccc;
border-radius: 4px;
background-color: #f0f0f0;
}
QProgressBar::chunk {
background-color: #4CAF50;
border-radius: 3px;
}
""")
progress_layout.addWidget(self.progress)
self.time_label = QLabel(f"{self._timeout} seconds remaining")
self.time_label.setAlignment(Qt.AlignmentFlag.AlignCenter)
self.time_label.setStyleSheet("color: #666;")
progress_layout.addWidget(self.time_label)
layout.addLayout(progress_layout)
# Warning about auto-transfer
self.warning_label = QLabel(
"⚠️ If you don't respond, control will transfer automatically."
)
self.warning_label.setWordWrap(True)
self.warning_label.setAlignment(Qt.AlignmentFlag.AlignCenter)
self.warning_label.setStyleSheet("color: #ff9800; font-style: italic;")
layout.addWidget(self.warning_label)
# Buttons
button_layout = QHBoxLayout()
button_layout.setSpacing(20)
self.accept_btn = QPushButton("✓ Accept")
self.accept_btn.setMinimumHeight(40)
self.accept_btn.setStyleSheet("""
QPushButton {
background-color: #4CAF50;
color: white;
border: none;
border-radius: 5px;
font-weight: bold;
font-size: 13px;
}
QPushButton:hover {
background-color: #45a049;
}
QPushButton:pressed {
background-color: #3d8b40;
}
""")
self.accept_btn.clicked.connect(self._on_accept)
button_layout.addWidget(self.accept_btn)
self.refuse_btn = QPushButton("✗ Refuse")
self.refuse_btn.setMinimumHeight(40)
self.refuse_btn.setStyleSheet("""
QPushButton {
background-color: #f44336;
color: white;
border: none;
border-radius: 5px;
font-weight: bold;
font-size: 13px;
}
QPushButton:hover {
background-color: #da190b;
}
QPushButton:pressed {
background-color: #c41000;
}
""")
self.refuse_btn.clicked.connect(self._on_refuse)
button_layout.addWidget(self.refuse_btn)
layout.addLayout(button_layout)
# Info text
info_label = QLabel(
"<small>If the beamline is busy, transfer will occur after "
"the current operation completes.</small>"
)
info_label.setWordWrap(True)
info_label.setAlignment(Qt.AlignmentFlag.AlignCenter)
info_label.setStyleSheet("color: #999;")
layout.addWidget(info_label)
def _start_timer(self):
self._timer = QTimer(self)
self._timer.setInterval(1000)
self._timer.timeout.connect(self._tick)
self._timer.start()
def _tick(self):
self._remaining -= 1
self.progress.setValue(self._remaining)
self.time_label.setText(f"{self._remaining} seconds remaining")
# Change progress bar color as time runs out
if self._remaining <= 10:
self.progress.setStyleSheet("""
QProgressBar {
border: 1px solid #ccc;
border-radius: 4px;
background-color: #f0f0f0;
}
QProgressBar::chunk {
background-color: #ff9800;
border-radius: 3px;
}
""")
if self._remaining <= 5:
self.progress.setStyleSheet("""
QProgressBar {
border: 1px solid #ccc;
border-radius: 4px;
background-color: #f0f0f0;
}
QProgressBar::chunk {
background-color: #f44336;
border-radius: 3px;
}
""")
self.time_label.setStyleSheet("color: #f44336; font-weight: bold;")
if self._remaining <= 0:
self._timer.stop()
# Timeout = auto-accept (as per your diagram: "Ignores request" → auto transfer)
self._on_accept()
def _on_accept(self):
self._timer.stop()
self.accepted_signal.emit()
self.accept()
def _on_refuse(self):
self._timer.stop()
self.refused_signal.emit()
self.reject()
def closeEvent(self, event):
"""Closing the dialog counts as ignoring = auto-accept on timeout."""
# Don't emit anything here - let the timeout handle it
# or the SSE stream will close the dialog when resolved
self._timer.stop()
super().closeEvent(event)
class BatonPendingDialog(QDialog):
"""
Dialog shown to the user who requested the baton while they wait for a response
or for the beamline queue to clear.
"""
cancelled_signal = Signal()
def __init__(self, target_user: str, timeout_seconds: int = 30, parent=None):
super().__init__(parent)
self.setWindowTitle("⏳ Baton Request Pending")
self.setModal(False)
self.setMinimumWidth(400)
self.setWindowFlags(self.windowFlags() | Qt.WindowType.WindowStaysOnTopHint)
self._timeout = timeout_seconds
self._remaining = timeout_seconds
self._target_user = target_user
self._setup_ui()
self._start_timer()
def _setup_ui(self):
layout = QVBoxLayout(self)
layout.setSpacing(15)
self.header = QLabel("⏳ Requesting Control")
header_font = QFont()
header_font.setPointSize(14)
header_font.setBold(True)
self.header.setFont(header_font)
self.header.setAlignment(Qt.AlignmentFlag.AlignCenter)
layout.addWidget(self.header)
line = QFrame()
line.setFrameShape(QFrame.Shape.HLine)
line.setFrameShadow(QFrame.Shadow.Sunken)
layout.addWidget(line)
self.message_label = QLabel(
f"Waiting for <b>{self._target_user}</b> to respond..."
)
self.message_label.setWordWrap(True)
self.message_label.setAlignment(Qt.AlignmentFlag.AlignCenter)
layout.addWidget(self.message_label)
self.progress_layout = QVBoxLayout()
self.progress = QProgressBar()
self.progress.setRange(0, max(1, self._timeout))
self.progress.setValue(self._timeout)
self.progress.setTextVisible(False)
self.progress.setFixedHeight(8)
self.progress.setStyleSheet("""
QProgressBar {
border: 1px solid #ccc;
border-radius: 4px;
background-color: #f0f0f0;
}
QProgressBar::chunk {
background-color: #2196F3;
border-radius: 3px;
}
""")
self.progress_layout.addWidget(self.progress)
self.time_label = QLabel(f"{self._timeout} seconds remaining")
self.time_label.setAlignment(Qt.AlignmentFlag.AlignCenter)
self.time_label.setStyleSheet("color: #666;")
self.progress_layout.addWidget(self.time_label)
layout.addLayout(self.progress_layout)
button_layout = QHBoxLayout()
self.cancel_btn = QPushButton("✗ Cancel Request")
self.cancel_btn.setMinimumHeight(40)
self.cancel_btn.setStyleSheet("""
QPushButton {
background-color: #f44336;
color: white;
border: none;
border-radius: 5px;
font-weight: bold;
font-size: 13px;
}
QPushButton:hover { background-color: #da190b; }
QPushButton:pressed { background-color: #c41000; }
""")
self.cancel_btn.clicked.connect(self._on_cancel)
button_layout.addWidget(self.cancel_btn)
layout.addLayout(button_layout)
def _start_timer(self):
self._timer = QTimer(self)
self._timer.setInterval(1000)
self._timer.timeout.connect(self._tick)
self._timer.start()
def _tick(self):
self._remaining -= 1
if self._remaining < 0:
self._remaining = 0
self.progress.setValue(self._remaining)
self.time_label.setText(f"{self._remaining} seconds remaining")
if self._remaining <= 0:
self._timer.stop()
def update_remaining(self, remaining: int):
self._remaining = remaining
self.progress.setValue(self._remaining)
self.time_label.setText(f"{self._remaining} seconds remaining")
def set_queued_state(self):
self._timer.stop()
self.header.setText("⏳ Transfer Queued")
self.message_label.setText("Waiting for current action to finish before receiving baton...")
self.progress.hide()
self.time_label.hide()
# Keep cancel button so they can abort the wait if they change their mind
def _on_cancel(self):
self._timer.stop()
self.cancelled_signal.emit()
self.reject()
def closeEvent(self, event):
self._timer.stop()
super().closeEvent(event)
+84 -16
View File
@@ -1,5 +1,8 @@
from PySide6.QtCore import Qt
from PySide6.QtWidgets import QDialog, QLineEdit, QVBoxLayout, QPushButton, QLabel, QComboBox, QCompleter
from PySide6.QtWidgets import (
QDialog, QVBoxLayout, QPushButton, QLabel,
QComboBox, QCompleter, QMessageBox
)
class PGroupDialog(QDialog):
@@ -8,42 +11,107 @@ class PGroupDialog(QDialog):
self.setWindowTitle("Change current p-group")
self.setMinimumWidth(300)
# Create a layout
self._pgroups = [str(p).strip() for p in (pgroups or []) if p is not None and str(p).strip()]
layout = QVBoxLayout(self)
# Add a label
self.label = QLabel("Set p-group:", self)
layout.addWidget(self.label)
self.combo = QComboBox(self)
self.combo.setEditable(True)
items = [str(p) for p in (pgroups or []) if p is not None and str(p).strip()]
self.combo.addItems(items)
self.combo.setInsertPolicy(QComboBox.InsertPolicy.NoInsert)
self.combo.addItems(self._pgroups)
completer = QCompleter(items, self)
completer = QCompleter(self._pgroups, self)
completer.setCaseSensitivity(Qt.CaseInsensitive)
completer.setFilterMode(Qt.MatchFlag.MatchContains) # requires Qt import; fallback to default if not desired
completer.setFilterMode(Qt.MatchFlag.MatchContains)
completer.setCompletionMode(QCompleter.CompletionMode.PopupCompletion)
self.combo.setCompleter(completer)
if curr_pgroup and curr_pgroup in items:
self.combo.setCurrentText(curr_pgroup)
elif curr_pgroup:
default_pgroup = self._latest_pgroup(self._pgroups)
if curr_pgroup and curr_pgroup in self._pgroups:
self.combo.setCurrentText(curr_pgroup)
elif default_pgroup is not None:
self.combo.setCurrentText(default_pgroup)
layout.addWidget(self.combo)
# Create buttons
self.ok_button = QPushButton("OK", self)
self.cancel_button = QPushButton("Cancel", self)
# Add buttons to the layout
layout.addWidget(self.ok_button)
layout.addWidget(self.cancel_button)
# Connect button signals
self.ok_button.clicked.connect(self.accept)
self.ok_button.clicked.connect(self._validate_and_accept)
self.cancel_button.clicked.connect(self.reject)
if self.combo.lineEdit() is not None:
self.combo.lineEdit().textEdited.connect(self._live_validate)
self._live_validate(self.combo.currentText())
@staticmethod
def _latest_pgroup(pgroups: list[str]) -> str | None:
"""
Return the numerically largest p-group, e.g. p16371 over p01234.
Falls back to lexicographic max if parsing fails.
"""
if not pgroups:
return None
def _key(pg: str):
s = str(pg).strip()
if s.startswith("p") and s[1:].isdigit():
return (1, int(s[1:]), s)
return (0, -1, s)
return max(pgroups, key=_key)
def _set_error_state(self, is_error: bool, message: str | None = None) -> None:
if is_error:
self.combo.setStyleSheet("border: 2px solid #d9534f;")
if message:
self.label.setText(f"Set p-group: <span style='color:#d9534f;'>{message}</span>")
else:
self.label.setText("Set p-group:")
else:
self.combo.setStyleSheet("")
self.label.setText("Set p-group:")
def _live_validate(self, text: str) -> None:
text = (text or "").strip()
if not text:
self._set_error_state(True, "Select a p-group")
return
if self._pgroups and text not in self._pgroups:
self._set_error_state(True, "Not in allowed list")
return
self._set_error_state(False)
def _validate_and_accept(self) -> None:
entered_text = (self.combo.currentText() or "").strip()
if not entered_text:
QMessageBox.warning(self, "Invalid P-Group", "You must select a p-group.")
self.combo.setFocus()
return
if self._pgroups and entered_text not in self._pgroups:
QMessageBox.warning(
self,
"Invalid P-Group",
f"P-group '{entered_text}' is not in your allowed list.\n"
f"Please select from: {', '.join(self._pgroups)}"
)
self.combo.setFocus()
if self.combo.lineEdit() is not None:
self.combo.lineEdit().selectAll()
self._set_error_state(True, "Not in allowed list")
return
self._set_error_state(False)
self.accept()
def get_input(self):
"""Return the input text when the dialog is accepted."""
#return self.text_entry.text()
return self.combo.currentText()
return (self.combo.currentText() or "").strip()
+227 -39
View File
@@ -5,8 +5,10 @@ from PySide6.QtGui import QFont
from PySide6.QtWidgets import QStatusBar, QDialog, QMenu, QMessageBox, QLabel, QSizePolicy
from aare.common.models import TokenData, BeamlineStateEnum, DAQStatusModel, SessionsStateEnum
from aare.gui.widgets.baton_request_dialog import BatonRequestDialog
from aare.gui.widgets.clickable_label import ClickableLabel
from aare.gui.widgets.pgroup_dialog import PGroupDialog
from aare.common.auth_models import BatonStatus, BatonRequestStatus
from aare.gui.widgets.value_label import ValueLabel
from aare.common.logger_config import setup_logger
@@ -22,10 +24,18 @@ class StatusBar(QStatusBar):
force_session = Signal()
end_session = Signal()
close_shutter = Signal()
open_shutter = Signal()
request_baton = Signal()
cancel_baton_request = Signal()
release_baton = Signal()
baton_request_accepted = Signal()
baton_request_refused = Signal()
get_all_pgroups = Signal()
staff_pgroups_loaded = Signal(list)
baton_request_received = Signal(dict)
close_shutter = Signal()
open_shutter = Signal()
def __init__(self, token: TokenData, parent=None):
super().__init__(parent)
@@ -38,6 +48,11 @@ class StatusBar(QStatusBar):
self._message_clear_timer.setSingleShot(True)
self._message_clear_timer.timeout.connect(self.clear_connection_message)
self._baton_status: BatonStatus | None = None
self._has_pending_request: bool = False
self._pgroup_dialog_for_baton: PGroupDialog | None = None
self._baton_request_dialog: BatonRequestDialog | None = None
self.message_label = QLabel("", self)
self.message_label.setVisible(False)
self.message_label.setSizePolicy(QSizePolicy.Policy.Maximum, QSizePolicy.Policy.Preferred)
@@ -200,26 +215,147 @@ class StatusBar(QStatusBar):
html_content_session = f"""Session: {session_flag}"""
self.session_label.setText(html_content_session)
@Slot(BatonStatus)
def update_baton_status(self, status: BatonStatus):
"""Update baton status from SSE stream."""
prev_incoming = bool(self._baton_status and self._baton_status.incoming_request)
# Detect if we just became the holder (e.g., from a queue resolving)
was_holder = bool(self._baton_status and self._baton_status.you_are_holder)
now_holder = bool(status and status.you_are_holder)
self._baton_status = status
self._has_pending_request = status.you_have_pending_request if status else False
self._update_session_display()
# If we just received the baton (and weren't the holder a moment ago)
if now_holder and not was_holder:
self._after_baton_granted_select_pgroup()
incoming = bool(status and status.incoming_request)
if incoming and not prev_incoming:
self._emit_incoming_baton_request(status)
if self._baton_request_dialog is not None and self._baton_request_dialog.isVisible():
if not incoming:
self._baton_request_dialog.close()
self._baton_request_dialog = None
def _emit_incoming_baton_request(self, status: BatonStatus) -> None:
requester = "Another user"
timeout = 30
if status.pending_request is not None:
requester = status.pending_request.requester_username or requester
timeout = int(status.pending_request.timeout_seconds or timeout)
self.baton_request_received.emit({
"requester": requester,
"timeout": timeout,
})
@Slot()
def _on_baton_dialog_accepted(self):
self.baton_request_accepted.emit()
self._baton_request_dialog = None
@Slot()
def _on_baton_dialog_refused(self):
self.baton_request_refused.emit()
self._baton_request_dialog = None
def _update_session_display(self):
"""Update session label based on current status."""
if self.__status is None:
return
session_state = self.__status.session.session
# Base text
if session_state == SessionsStateEnum.OwnedByYou:
text = "Session: You"
if self._baton_status and self._baton_status.incoming_request:
text = "Session: You (⚡ Request)"
elif session_state == SessionsStateEnum.OwnedByElse:
holder_name = ""
if self._baton_status and self._baton_status.holder:
holder_name = self._baton_status.holder.username
text = f"Session: {holder_name or 'Other'}"
if self._has_pending_request:
text += " (⏳ Waiting)"
else:
text = "Session: Vacant"
self.session_label.setText(text)
def show_session_menu(self):
menu = QMenu(self)
is_busy = self.__status and self.__status.busy
is_vacant = self.__status and self.__status.session.session == SessionsStateEnum.Vacant
action_1 = menu.addAction("Grab")
action_1.setEnabled(bool(not is_busy or self.__is_staff or is_vacant))
action_1.triggered.connect(self._on_grab_clicked)
action_2 = menu.addAction("End")
action_2.setEnabled(bool(not is_busy or self.__is_staff))
action_2.triggered.connect(self.end_session_clicked)
is_yours = self.__status and self.__status.session.session == SessionsStateEnum.OwnedByYou
is_other = self.__status and self.__status.session.session == SessionsStateEnum.OwnedByElse
# Determine holder info from baton status
holder_is_staff = (
self._baton_status and
self._baton_status.holder and
self._baton_status.holder.is_staff
)
# --- GRAB / REQUEST ---
if is_vacant:
# Vacant - simple grab
action_grab = menu.addAction("Grab")
action_grab.setEnabled(True)
action_grab.triggered.connect(self._on_grab_clicked)
elif is_other:
# Someone else has it
if self._has_pending_request:
# Already have a pending request - show cancel option
action_cancel = menu.addAction("Cancel Request")
action_cancel.triggered.connect(self._on_cancel_request_clicked)
elif self.__is_staff:
# Staff can always grab (override)
action_grab = menu.addAction("Grab (Override)")
action_grab.setEnabled(not is_busy) # Still respect busy for safety
action_grab.triggered.connect(self._on_grab_clicked)
elif holder_is_staff:
# Non-staff cannot request from staff
allowed = self._baton_status and getattr(self._baton_status, "allow_non_staff_request", False)
action_grab = menu.addAction("Request from Staff")
action_grab.setEnabled(allowed)
if allowed:
action_grab.triggered.connect(self._on_grab_clicked)
else:
# Same level - request with timeout
action_request = menu.addAction("Request Control")
action_request.setEnabled(True)
action_request.triggered.connect(self._on_grab_clicked)
elif is_yours:
# You have it - show release option
action_release = menu.addAction("Release")
action_release.setEnabled(not is_busy)
action_release.triggered.connect(self._on_release_clicked)
menu.addSeparator()
# --- END SESSION (cleanup) ---
action_end = menu.addAction("End Session")
action_end.setEnabled(bool((is_yours and not is_busy) or self.__is_staff))
action_end.triggered.connect(self.end_session_clicked)
# --- STAFF: FORCE GRAB (emergency) ---
if self.__is_staff and is_other:
menu.addSeparator()
action_force = menu.addAction("⚠️ Force Take Over")
action_force.triggered.connect(self._on_force_session_clicked)
label_geometry = self.session_label.geometry()
menu.move(self.mapToGlobal(label_geometry.topLeft()) - QPoint(0, menu.sizeHint().height()))
menu.setFixedWidth(label_geometry.width())
menu.exec()
def show_pgroup_menu(self):
in_curr = self.__status and self.__status.session.current_pgroup in (self.__allowed_pgroups or [])
in_curr = self.__status
logger.info(f"in_curr is {in_curr}")
@@ -301,40 +437,90 @@ class StatusBar(QStatusBar):
menu.exec()
def _latest_pgroup(self, pgroups: list[str]) -> str | None:
if not pgroups:
return None
def _key(pg: str):
s = str(pg).strip()
if s.startswith("p") and s[1:].isdigit():
return (1, int(s[1:]), s)
return (0, -1, s)
return max(pgroups, key=_key)
def _after_baton_granted_select_pgroup(self) -> None:
"""
After baton grant:
- if exactly one allowed p-group, apply it automatically
- otherwise prompt user to choose from their allowed list
"""
pgroups = [str(p).strip() for p in (self.__allowed_pgroups or []) if p is not None and str(p).strip()]
if not pgroups:
return
if len(pgroups) == 1:
self.set_pgroup.emit(pgroups[0])
return
default_pgroup = self._latest_pgroup(pgroups)
curr = None
if self.__status and self.__status.session:
curr = self.__status.session.current_pgroup or default_pgroup
self._pgroup_dialog_for_baton = PGroupDialog(
curr_pgroup=curr,
pgroups=pgroups,
parent=self,
)
if self._pgroup_dialog_for_baton.exec() == QDialog.DialogCode.Accepted:
selected_pgroup = self._pgroup_dialog_for_baton.get_input()
if selected_pgroup:
self.set_pgroup.emit(selected_pgroup)
self._pgroup_dialog_for_baton = None
def _show_post_grant_pgroup_dialog(self, available_pgroups: list[str]) -> None:
curr_pgroup = None
if self.__status and self.__status.session:
curr_pgroup = self.__status.session.current_pgroup
if self._pgroup_dialog_for_baton is not None and self._pgroup_dialog_for_baton.isVisible():
return
self._pgroup_dialog_for_baton = PGroupDialog(
curr_pgroup=curr_pgroup,
pgroups=available_pgroups,
parent=self
)
if self._pgroup_dialog_for_baton.exec() == QDialog.DialogCode.Accepted:
selected_pgroup = (self._pgroup_dialog_for_baton.get_input() or "").strip()
if selected_pgroup:
self.set_pgroup.emit(selected_pgroup)
self._pgroup_dialog_for_baton = None
def _on_grab_clicked(self):
self.grab_session_clicked()
"""Handle grab/request click - baton first, p-group after grant."""
self.request_baton.emit()
def after_grab():
def _on_release_clicked(self):
"""Handle release click."""
self.release_baton.emit()
if not self.__status:
logger.info("status is None when session is grabbed")
return
def _on_cancel_request_clicked(self):
"""Handle cancel request click."""
self.cancel_baton_request.emit()
check_state = self.__status.state in (
BeamlineStateEnum.SampleAlignment,
BeamlineStateEnum.SampleExchange,
BeamlineStateEnum.DewarTransfer,
BeamlineStateEnum.Maintenance,
)
session_ownership = self.__status.session.session == SessionsStateEnum.OwnedByYou
have_pgroup = (self.__status.session.current_pgroup in (self.__allowed_pgroups or []))
allowed = (not self.__status.busy and check_state and session_ownership and have_pgroup) or self.__is_staff
if allowed:
self.show_change_dialog()
else:
logger.debug(
"not allowed to change pgroup due to; "
f"session ownership: {session_ownership}, allowed_pgroup: {have_pgroup}, "
f"beamline busy: {self.__status.busy}, beamline state: {check_state}"
)
QTimer.singleShot(300, after_grab)
def _on_force_session_clicked(self):
"""Staff emergency force take over (bypasses baton protocol)."""
self.force_session.emit()
def grab_session_clicked(self):
self.force_session.emit()
"""Legacy method - now routes to baton request."""
self.request_baton.emit()
def end_session_clicked(self):
self.end_session.emit()
@@ -359,6 +545,7 @@ class StatusBar(QStatusBar):
return
def _generate_pgroup_dialogue(self, curr: str | None = None, pgroups: list | None = None):
logger.info(pgroups)
dialog = PGroupDialog(curr_pgroup=curr, pgroups=pgroups)
if dialog.exec() == QDialog.DialogCode.Accepted:
entered_text = dialog.get_input()
@@ -370,6 +557,7 @@ class StatusBar(QStatusBar):
f"P-group '{entered_text}' is not in your allowed list.\n"
f"Please select from: {', '.join(pgroups)}"
)
self._generate_pgroup_dialogue(curr=curr, pgroups=pgroups)
return
self.set_pgroup.emit(entered_text)