Major update from x10sa #48
+1
-1
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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),
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
@@ -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
@@ -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"}
|
||||
@@ -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",
|
||||
},
|
||||
)
|
||||
@@ -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
@@ -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
@@ -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
@@ -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:
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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",
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
@@ -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] = []
|
||||
|
||||
@@ -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")
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user