From 886bdcc0327a824461752e951d200a037fe1a6b5 Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Fri, 13 Mar 2026 09:25:11 +0100 Subject: [PATCH 01/56] Enhance SSE connection handling in `tellupdater` for compatibility and improved error handling --- src/aare/daq/tellupdater.py | 12 ++++++++---- 1 file changed, 8 insertions(+), 4 deletions(-) diff --git a/src/aare/daq/tellupdater.py b/src/aare/daq/tellupdater.py index e9815dec..449a0fcc 100644 --- a/src/aare/daq/tellupdater.py +++ b/src/aare/daq/tellupdater.py @@ -40,13 +40,17 @@ def listen_to_sse(): return sse_url = tell_client.url + "/events" try: - #response = requests.get(sse_url, stream=True) - client = sseclient.SSEClient(sse_url) + # Build a streaming HTTP response first to support common sseclient variants. + response = requests.get(sse_url, stream=True, timeout=(5, 60)) + response.raise_for_status() + client = sseclient.SSEClient(response) print("[SSE][listen_to_sse] Initial detected pucks fetch on connect") handle_tell_change_event() - for event in client.events(): + # 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) @@ -164,4 +168,4 @@ def main(): time.sleep(5) if __name__ == "__main__": - main() \ No newline at end of file + main() -- 2.54.0 From 68df2930b3c0ecd7ec33be5c97d8d30ec03c8498 Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Fri, 13 Mar 2026 09:35:21 +0100 Subject: [PATCH 02/56] Streamline SSE client setup in `tellupdater` by removing redundant HTTP stream handling --- src/aare/daq/tellupdater.py | 7 ++----- 1 file changed, 2 insertions(+), 5 deletions(-) diff --git a/src/aare/daq/tellupdater.py b/src/aare/daq/tellupdater.py index 449a0fcc..defc8dd1 100644 --- a/src/aare/daq/tellupdater.py +++ b/src/aare/daq/tellupdater.py @@ -4,7 +4,6 @@ import threading import websocket import sseclient -import requests import time from aareDB.models import PuckWithTellPosition @@ -40,10 +39,8 @@ def listen_to_sse(): return sse_url = tell_client.url + "/events" try: - # Build a streaming HTTP response first to support common sseclient variants. - response = requests.get(sse_url, stream=True, timeout=(5, 60)) - response.raise_for_status() - client = sseclient.SSEClient(response) + # Use URL directly for sseclient variants that manage their own HTTP stream. + client = sseclient.SSEClient(sse_url) print("[SSE][listen_to_sse] Initial detected pucks fetch on connect") handle_tell_change_event() -- 2.54.0 From 6d0cf1fc4711e26c622c75963b0d7aef33eb462d Mon Sep 17 00:00:00 2001 From: appleb_m Date: Fri, 13 Mar 2026 09:38:11 +0100 Subject: [PATCH 03/56] Tell Client: temp fix of bug --- src/aare/devices/tell_client.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/src/aare/devices/tell_client.py b/src/aare/devices/tell_client.py index 9ff31f99..9ef0b5f5 100755 --- a/src/aare/devices/tell_client.py +++ b/src/aare/devices/tell_client.py @@ -633,6 +633,9 @@ class TellClientProxy: def url(self): return self._get_client().url + def __getattr__(self, name): + return getattr(self._get_client(), name) + # Delegate methods used by DAQ; add more as needed def get_mounted_sample(self) -> SampleDewarAddress | None: return self._get_client().get_mounted_sample() -- 2.54.0 From 133a4a68766b14643bd947db7d057197220b4d87 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Fri, 13 Mar 2026 09:40:06 +0100 Subject: [PATCH 04/56] TellUpdater: remove unessacery print statement --- src/aare/daq/tellupdater.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/aare/daq/tellupdater.py b/src/aare/daq/tellupdater.py index defc8dd1..5959fe45 100644 --- a/src/aare/daq/tellupdater.py +++ b/src/aare/daq/tellupdater.py @@ -48,7 +48,7 @@ def listen_to_sse(): # 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}") + #print(f"event = {event.event} with data: {event.data}") if event.event == "DewarContentUpdate": on_sse_event(event) except Exception as exc: -- 2.54.0 From 6a7e61c1a7abc98fc2c967766d8470593f32913f Mon Sep 17 00:00:00 2001 From: appleb_m Date: Fri, 13 Mar 2026 09:40:22 +0100 Subject: [PATCH 05/56] SmargonTrace: change location of csv to /sls/mx/applications/logs --- src/aare/daq/daq.py | 2 +- src/aare/gui/panels/smargon_trace_panel.py | 1 + 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/src/aare/daq/daq.py b/src/aare/daq/daq.py index 9a72ec0e..db2a5d3b 100644 --- a/src/aare/daq/daq.py +++ b/src/aare/daq/daq.py @@ -62,7 +62,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 diff --git a/src/aare/gui/panels/smargon_trace_panel.py b/src/aare/gui/panels/smargon_trace_panel.py index fe958702..7f031cc4 100644 --- a/src/aare/gui/panels/smargon_trace_panel.py +++ b/src/aare/gui/panels/smargon_trace_panel.py @@ -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] = [] -- 2.54.0 From 3c769106bf8254edd0c9081561deca17f7c518a3 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Fri, 13 Mar 2026 09:40:44 +0100 Subject: [PATCH 06/56] Aerotech: think we are removing automation1 api, WIP --- src/aare/devices/aerotech.py | 290 +++++++++++++++++------------------ 1 file changed, 145 insertions(+), 145 deletions(-) diff --git a/src/aare/devices/aerotech.py b/src/aare/devices/aerotech.py index 5884e40b..bed65a09 100644 --- a/src/aare/devices/aerotech.py +++ b/src/aare/devices/aerotech.py @@ -5,7 +5,7 @@ import time from enum import Enum from typing import Union -import automation1 as a1 +#import automation1 as a1 from epics import PV, Motor, caput, caget from aare.common.beamline import MXBeamline, mx_beamline @@ -173,150 +173,150 @@ class AerotechControllerEpics: 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 - 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() +# +# 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() #aerotech_epics = AerotechControllerEpics(beamline) #aerotech_epics.omega.speed = 80.0 -- 2.54.0 From 6110a24d61a75e867359c40f9a92499edc254bc5 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Fri, 13 Mar 2026 09:40:52 +0100 Subject: [PATCH 07/56] Aerotech: think we are removing automation1 api, WIP --- src/aare/daq/devices.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/aare/daq/devices.py b/src/aare/daq/devices.py index 0455beea..075aa57c 100644 --- a/src/aare/daq/devices.py +++ b/src/aare/daq/devices.py @@ -28,7 +28,7 @@ class BeamlineDevices: 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(controller_ip="129.129.118.96") self.__smargon = smargon.Smargon(beamline) self.__ring_current_pv = PV(f"ARS07-DPCT-0100:CURR") -- 2.54.0 From a3d4027b087865f950707459e82942ae78610d86 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Fri, 13 Mar 2026 09:41:27 +0100 Subject: [PATCH 08/56] GUI: Use Pserv, spark and x10sa-queue --- src/aare/gui/gui.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/aare/gui/gui.py b/src/aare/gui/gui.py index 284ec734..de22821b 100644 --- a/src/aare/gui/gui.py +++ b/src/aare/gui/gui.py @@ -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-pserv-01:9089" # "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" -- 2.54.0 From b44fd51610f4a7bdcd1dfebab1d16d371ca7dfe9 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Fri, 13 Mar 2026 09:43:06 +0100 Subject: [PATCH 09/56] DAQ: added park_and_dry functionality - needs testing --- src/aare/daq/daq.py | 15 +++++++++++++++ src/aare/daq/server.py | 9 +++++++++ src/aare/gui/threads/daq_worker.py | 4 ++++ 3 files changed, 28 insertions(+) diff --git a/src/aare/daq/daq.py b/src/aare/daq/daq.py index 9a72ec0e..9e011953 100644 --- a/src/aare/daq/daq.py +++ b/src/aare/daq/daq.py @@ -405,6 +405,21 @@ 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 diff --git a/src/aare/daq/server.py b/src/aare/daq/server.py index 45f57d4f..f47216f7 100644 --- a/src/aare/daq/server.py +++ b/src/aare/daq/server.py @@ -316,6 +316,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): diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py index 16f77b8d..90a7f2d6 100644 --- a/src/aare/gui/threads/daq_worker.py +++ b/src/aare/gui/threads/daq_worker.py @@ -811,6 +811,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()) -- 2.54.0 From 12412c29d812200ebcad51be2db79d68067abf69 Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Fri, 20 Mar 2026 11:55:38 +0100 Subject: [PATCH 10/56] Improve SSE client resilience in `tellupdater` by adding retry mechanism and enhanced error logging --- src/aare/daq/tellupdater.py | 32 +++++++++++++++++++------------- 1 file changed, 19 insertions(+), 13 deletions(-) diff --git a/src/aare/daq/tellupdater.py b/src/aare/daq/tellupdater.py index 5959fe45..2616ff4e 100644 --- a/src/aare/daq/tellupdater.py +++ b/src/aare/daq/tellupdater.py @@ -38,21 +38,27 @@ 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: - # Use URL directly for sseclient variants that manage their own HTTP stream. - client = sseclient.SSEClient(sse_url) + + while True: + try: + print(f"[SSE][INFO] Attempting to connect to {sse_url}...") + # Use URL directly for sseclient variants that manage their own HTTP stream. + client = sseclient.SSEClient(sse_url) - print("[SSE][listen_to_sse] Initial detected pucks fetch on connect") - handle_tell_change_event() + 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] Failed to connect to {sse_url}: {exc}") + # 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): """ -- 2.54.0 From 9db0e45f8b8c49e5257a7491f60e00581c9a9baa Mon Sep 17 00:00:00 2001 From: appleb_m Date: Fri, 20 Mar 2026 16:56:35 +0100 Subject: [PATCH 11/56] GUI: FD panel bug fix --- src/aare/common/automation_models.py | 156 ++++++++++++++++ src/aare/devices/automation1_python_api.py | 187 ++++++++++++++++++++ src/aare/gui/panels/face_detection_panel.py | 2 +- 3 files changed, 344 insertions(+), 1 deletion(-) create mode 100644 src/aare/common/automation_models.py create mode 100644 src/aare/devices/automation1_python_api.py diff --git a/src/aare/common/automation_models.py b/src/aare/common/automation_models.py new file mode 100644 index 00000000..67be794b --- /dev/null +++ b/src/aare/common/automation_models.py @@ -0,0 +1,156 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from enum import Enum +from typing import Any + + +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) + + +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) \ No newline at end of file diff --git a/src/aare/devices/automation1_python_api.py b/src/aare/devices/automation1_python_api.py new file mode 100644 index 00000000..f3ca8c8f --- /dev/null +++ b/src/aare/devices/automation1_python_api.py @@ -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) \ No newline at end of file diff --git a/src/aare/gui/panels/face_detection_panel.py b/src/aare/gui/panels/face_detection_panel.py index 9551551d..2335d0f6 100644 --- a/src/aare/gui/panels/face_detection_panel.py +++ b/src/aare/gui/panels/face_detection_panel.py @@ -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, -- 2.54.0 From c25e20a10c8094103ec5979b58ad1c7c7cdbccec Mon Sep 17 00:00:00 2001 From: appleb_m Date: Fri, 20 Mar 2026 16:57:32 +0100 Subject: [PATCH 12/56] GUI: main window add restore default view and way to recover refernece samples panel, also capture default window state. --- src/aare/gui/main_window.py | 70 +++++++++++++++++++++++++++++++------ 1 file changed, 60 insertions(+), 10 deletions(-) diff --git a/src/aare/gui/main_window.py b/src/aare/gui/main_window.py index 06d0c67f..458de6f1 100644 --- a/src/aare/gui/main_window.py +++ b/src/aare/gui/main_window.py @@ -47,6 +47,7 @@ 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.camera_image import SampleCameraImageLabel from aare.gui.widgets.no_wheel_scroll_area import NoWheelScrollArea @@ -76,6 +77,7 @@ class MainWindow(QMainWindow): self._beamline_recovery_dialog = None self._controls_help_dialog = None self._cleanup_done = False + self._default_window_state = None # Tutorial manager (define tutorials after widgets exist) self._tutorial_event_bus = TutorialEventBus(self) @@ -221,6 +223,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() @@ -282,6 +289,7 @@ class MainWindow(QMainWindow): self.setWindowTitle("AareGUI") self.create_menu_bar() + self._capture_default_window_state() self._restore_window_state() # Define tutorials now that the UI exists @@ -522,35 +530,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,6 +580,7 @@ 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 ) @@ -578,6 +593,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 +627,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, @@ -684,6 +731,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() -- 2.54.0 From 322370ac087aa18feec26a3313f5cd56214f3614 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Tue, 24 Mar 2026 09:28:11 +0100 Subject: [PATCH 13/56] ALC: remove background function and button --- src/aare/daq/daq.py | 25 - src/aare/daq/server.py | 7 - src/aare/devices/aerotech.py | 504 ++++++-------------- src/aare/gui/main_window.py | 1 - src/aare/gui/panels/loop_centering_panel.py | 2 - 5 files changed, 145 insertions(+), 394 deletions(-) diff --git a/src/aare/daq/daq.py b/src/aare/daq/daq.py index 9e011953..6fc07354 100644 --- a/src/aare/daq/daq.py +++ b/src/aare/daq/daq.py @@ -754,31 +754,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 diff --git a/src/aare/daq/server.py b/src/aare/daq/server.py index f47216f7..faf1bfbb 100644 --- a/src/aare/daq/server.py +++ b/src/aare/daq/server.py @@ -611,13 +611,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") diff --git a/src/aare/devices/aerotech.py b/src/aare/devices/aerotech.py index 5884e40b..47ae6cad 100644 --- a/src/aare/devices/aerotech.py +++ b/src/aare/devices/aerotech.py @@ -1,386 +1,172 @@ -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 -class AxisEnum(Enum): - X = "gmx" - Y = "gmy" - Z = "gmz" - OMEGA = "Omega" - -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 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 enable_all(self): - self.__enable_all.put(1) - self.__enable_all.put(0) - - def acknowledge_all(self): - self.__acknowledge_all.put(1) - self.__acknowledge_all.put(0) - - def _disable_all(self): - self.__disable_all.put(1) - self.__disable_all.put(0) - - def stop_all(self): - self.__stop_all.put(1) - self.__stop_all.put(0) - - 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): - try: - axis.put(attr=attr, value=value) - except Exception as e: - raise ValueError(f"Error setting {axis.name} {attr} to {value}: {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) +AEROTECH_HOME = AerotechCoordinate(at_mm=Coordinate(x=0, y=0, z=0), omega_deg=0) - 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) +class AerotechController(object): - 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 + 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 ValueError(f"Invalid type {type(value)} for global variable") + raise Exception("unknown beamline") - def get_global_variable(self, index: int, var_type:VariableTypeEnum): - return self.__read_global_from_index(index, var_type.name) + if not self.__simulated: + self.__client = ApiClient(Configuration(host=self.__base)) + self.__api = DefaultApi(self.__client) - def set_global_variable(self, index: int, value: Union[int, float, str]): - self.__set_global(index, value) + def __make_aerotech_target( + self, + coord: AerotechCoordinate, + wait: bool = False, + incremental: bool = False, + ) -> Target: + at_mm = coord.at_mm + + 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, + var_async=not wait, + incremental=incremental, + ) + + 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 cancel(self): + self.__api.cancel_post() + + def is_idle(self) -> bool: + status = self.__api.status_get() + return status.state == 'Idle' + + def get_position(self) -> AerotechCoordinate: + status = self.__api.status_get() + return AerotechCoordinate( + at_mm=Coordinate( + x=status.x.pos, + y=status.y.pos, + z=status.z.pos, + ), + omega_deg=status.u.pos, + ) -class AerotechController: - def __init__(self, controller_ip: str): - if controller_ip is None: - self.controller = None - else: - self.controller = a1.Controller.connect(controller_ip) + 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), + ) + return self.__api.status_get() - self.status_item_configuration = a1.StatusItemConfiguration() - self.start_controller() + 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 start_controller(self): - self.controller.start() + def home_aerotech(self): + return self.__api.home_post() - def disconnect(self): - self.controller.disconnect() + def wait_till_done(self, timeout=60): + return self.__api.wait_till_done_post(timeout=timeout) - def enable_motion(self, axis: str): - self.controller.runtime.commands.motion.enable(axis.upper()) + def position( + self, + target: AerotechCoordinate, + /, + wait: bool = True, + incremental: bool = False, + ): + if self.__simulated: + return self.__pos - 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()) + payload = self.__make_aerotech_target(target, wait=wait, incremental=incremental) + return self.__api.position_post(payload) - 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 rotation_scan(self, rotation_deg: float | int, + time_sec: float | int, + start_pos_deg: float | int, + var_async: bool = False): + payload = RotationRequest( + rotation_deg=rotation_deg, + time_sec=time_sec, + start_pos_deg=start_pos_deg, + var_async=var_async + ) + if self.__simulated: + return payload - 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}") + return self.__api.rotation_scan_post(payload) - 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 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, + var_async: Optional[bool] = False - 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") + ): + 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, + var_async=var_async, + ) + if self.__simulated: + return payload + return self.__api.grid_scan_post(payload) - 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 screening_scan(self, + rotation_deg: float | int, + wedge_deg: float | int, + time_sec: float | int, + steps: int, + var_async: bool = False + ): + payload = ScreenRequest( + rotation_deg=rotation_deg, + wedge_deg=wedge_deg, + time_sec=time_sec, + steps=steps, + var_async=var_async + ) + if self.__simulated: + return payload + return self.__api.screening_post(payload) - 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() - - #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()) diff --git a/src/aare/gui/main_window.py b/src/aare/gui/main_window.py index 458de6f1..4d1bc2a7 100644 --- a/src/aare/gui/main_window.py +++ b/src/aare/gui/main_window.py @@ -314,7 +314,6 @@ class MainWindow(QMainWindow): 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) diff --git a/src/aare/gui/panels/loop_centering_panel.py b/src/aare/gui/panels/loop_centering_panel.py index 41de4548..5c34d340 100644 --- a/src/aare/gui/panels/loop_centering_panel.py +++ b/src/aare/gui/panels/loop_centering_panel.py @@ -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) -- 2.54.0 From cacc98678504c191e04ee7d35605f4241de9632e Mon Sep 17 00:00:00 2001 From: appleb_m Date: Tue, 24 Mar 2026 09:28:56 +0100 Subject: [PATCH 14/56] Workflows: use ABR_POS_MOUNT --- src/aare/daq/workflows.py | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/src/aare/daq/workflows.py b/src/aare/daq/workflows.py index 04addece..094da900 100644 --- a/src/aare/daq/workflows.py +++ b/src/aare/daq/workflows.py @@ -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') -- 2.54.0 From ba559c5efa728c6bd1aff676d3c72a75a77670b3 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Tue, 24 Mar 2026 09:30:25 +0100 Subject: [PATCH 15/56] GUI: fixed tell error flashing - hopefully --- src/aare/gui/threads/daq_worker.py | 19 ++++++++++++------- src/aare/gui/widgets/alert_banner.py | 6 ++++++ 2 files changed, 18 insertions(+), 7 deletions(-) diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py index 90a7f2d6..7bf240cb 100644 --- a/src/aare/gui/threads/daq_worker.py +++ b/src/aare/gui/threads/daq_worker.py @@ -47,6 +47,7 @@ class DAQWorker(QObject): 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,6 +78,7 @@ 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._server_connected: bool | None = None self._last_server_error: str | None = None @@ -244,12 +246,19 @@ class DAQWorker(QObject): 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) + tell_err_text = None if tell_err is None else str(tell_err).strip() + + + tell_changed = ( + self._last_tell_connected is None + or self._last_tell_connected != tell_conn + or self._last_tell_error != tell_err_text + ) - tell_changed = self._last_tell_connected is not None and self._last_tell_connected != tell_conn smargon_changed = ( self._last_smargon_connected is not None and self._last_smargon_connected != smargon_conn ) @@ -449,10 +458,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") diff --git a/src/aare/gui/widgets/alert_banner.py b/src/aare/gui/widgets/alert_banner.py index fe10f4c7..5a440057 100644 --- a/src/aare/gui/widgets/alert_banner.py +++ b/src/aare/gui/widgets/alert_banner.py @@ -11,6 +11,8 @@ class AlertBanner(QFrame): def __init__(self, parent=None): super().__init__(parent) + self._current_message: str | None = None + self._clear_timer = QTimer(self) self._clear_timer.setSingleShot(True) self._clear_timer.timeout.connect(self.clear_message) @@ -40,6 +42,9 @@ class AlertBanner(QFrame): self.clear_message() return + if msg == self._current_message: + return + if is_error: decorated = f"🛑 {msg} 🛑" self.setStyleSheet( @@ -79,6 +84,7 @@ class AlertBanner(QFrame): @Slot() def clear_message(self): + self._current_message = None self._clear_timer.stop() self._label.clear() self.setVisible(False) -- 2.54.0 From eaa6f0cd0c8ca8592e31b35c06cea13600f094a1 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Tue, 24 Mar 2026 09:31:09 +0100 Subject: [PATCH 16/56] DAQ: throttle userrightsexception. so it does not repeatedly log the same error --- src/aare/common/exception_handler.py | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/src/aare/common/exception_handler.py b/src/aare/common/exception_handler.py index 7098c362..9a84f6f2 100644 --- a/src/aare/common/exception_handler.py +++ b/src/aare/common/exception_handler.py @@ -104,6 +104,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 +117,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 -- 2.54.0 From 95010444fd62a44205b315291ad41374adfc9f46 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Tue, 24 Mar 2026 09:31:16 +0100 Subject: [PATCH 17/56] DAQ: throttle userrightsexception. so it does not repeatedly log the same error --- src/aare/common/exception_handler.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/src/aare/common/exception_handler.py b/src/aare/common/exception_handler.py index 9a84f6f2..9a81665d 100644 --- a/src/aare/common/exception_handler.py +++ b/src/aare/common/exception_handler.py @@ -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 -- 2.54.0 From 3edf9200fd6be59cf2edaaeaf3cd326615ced4e3 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Tue, 24 Mar 2026 09:41:22 +0100 Subject: [PATCH 18/56] Tell: start of rework - needs testing --- src/aare/devices/tell_backend.py | 410 ++++++++++++++++++++++++++++ src/aare/devices/tell_client.py | 454 +++++++------------------------ 2 files changed, 503 insertions(+), 361 deletions(-) create mode 100644 src/aare/devices/tell_backend.py diff --git a/src/aare/devices/tell_backend.py b/src/aare/devices/tell_backend.py new file mode 100644 index 00000000..0a6f962e --- /dev/null +++ b/src/aare/devices/tell_backend.py @@ -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() \ No newline at end of file diff --git a/src/aare/devices/tell_client.py b/src/aare/devices/tell_client.py index 9ff31f99..f7a4d8c9 100755 --- a/src/aare/devices/tell_client.py +++ b/src/aare/devices/tell_client.py @@ -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) \ No newline at end of file + backend = SimTellBackend() + else: + backend = LazyTellBackend( + factory=lambda: PShellTellBackend(bl), + retry_interval_s=2.0, + ) + return TellClient(bl, backend=backend) \ No newline at end of file -- 2.54.0 From e9f368895297d0d3caf1d970f198e46e3d3df745 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Tue, 24 Mar 2026 10:59:16 +0100 Subject: [PATCH 19/56] Models: added aerotech connected and error to DAQStatusModel --- src/aare/common/models.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/src/aare/common/models.py b/src/aare/common/models.py index f90ad14d..1db6f879 100644 --- a/src/aare/common/models.py +++ b/src/aare/common/models.py @@ -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 -- 2.54.0 From 5634554c00db00aaa9725a1bdf0ffa8003b2511b Mon Sep 17 00:00:00 2001 From: appleb_m Date: Tue, 24 Mar 2026 11:06:32 +0100 Subject: [PATCH 20/56] GUI/DAQ: fixed connection error handling for tell, daq, smargon. Added handling for aerotech. Fixed bug where connection messaged flashed when polling status, fixed bug where tell connected would quash other messages and appear as an error. --- src/aare/daq/daq.py | 13 ++++++++ src/aare/gui/main_window.py | 20 +++++++++++- src/aare/gui/threads/daq_worker.py | 46 +++++++++++++++++++++++----- src/aare/gui/widgets/alert_banner.py | 8 ++++- 4 files changed, 78 insertions(+), 9 deletions(-) diff --git a/src/aare/daq/daq.py b/src/aare/daq/daq.py index 6fc07354..67c47868 100644 --- a/src/aare/daq/daq.py +++ b/src/aare/daq/daq.py @@ -1622,6 +1622,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. @@ -1704,6 +1714,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, @@ -1721,6 +1732,8 @@ 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): diff --git a/src/aare/gui/main_window.py b/src/aare/gui/main_window.py index 4d1bc2a7..05069a1d 100644 --- a/src/aare/gui/main_window.py +++ b/src/aare/gui/main_window.py @@ -78,6 +78,8 @@ class MainWindow(QMainWindow): self._controls_help_dialog = None self._cleanup_done = False self._default_window_state = None + self._last_status_message: str | None = None + self._last_status_is_error: bool | None = None # Tutorial manager (define tutorials after widgets exist) self._tutorial_event_bus = TutorialEventBus(self) @@ -482,10 +484,26 @@ class MainWindow(QMainWindow): self.daq.fluorimeter_spectrum_update.connect(lambda: self.fluor_panel_dock.setVisible(True)) self.daq.status_message.connect(self.status_bar.show_connection_message) - self.daq.status_message.connect(self.alert_banner.show_message) + self.daq.status_message.connect(self._show_status_message_once) register_tutorials(self, self.tutorial_manager) + @Slot(str, bool) + def _show_status_message_once(self, message: str, is_error: bool = True) -> None: + message = (message or "").strip() + if not message: + self.alert_banner.clear_message() + self._last_status_message = None + self._last_status_is_error = None + return + + if message == self._last_status_message and is_error == self._last_status_is_error: + return + + self._last_status_message = message + self._last_status_is_error = is_error + self.alert_banner.show_message(message, is_error=is_error) + @Slot(QPixmap) def _on_samcam_prediction_pixmap(self, pix: QPixmap) -> None: self._last_pred_image_ts = time.monotonic() diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py index 7bf240cb..36fb7464 100644 --- a/src/aare/gui/threads/daq_worker.py +++ b/src/aare/gui/threads/daq_worker.py @@ -80,6 +80,9 @@ class DAQWorker(QObject): 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 @@ -196,8 +199,10 @@ class DAQWorker(QObject): *, 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] = [] @@ -212,11 +217,22 @@ 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 connection errors, please inform your local contact.", + 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]} connection errors, please inform your local contact.", True, ) device = disconnected[0] @@ -227,10 +243,11 @@ class DAQWorker(QObject): ) if restored: + if len(restored) == 3: + return (None, "TELL, Smargon, and Aerotech connections restored.", 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]} connections restored.", False) + return (None, f"{restored[0]} connection restored.", False) return (None, None, False) @@ -258,19 +275,31 @@ class DAQWorker(QObject): or 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 is None + or self._last_smargon_connected != smargon_conn + or self._last_smargon_error != smargon_err_text + ) + aerotech_changed = ( + self._last_aerotech_connected is None + or 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) @@ -278,6 +307,8 @@ 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) @@ -293,6 +324,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}") diff --git a/src/aare/gui/widgets/alert_banner.py b/src/aare/gui/widgets/alert_banner.py index 5a440057..86e5c791 100644 --- a/src/aare/gui/widgets/alert_banner.py +++ b/src/aare/gui/widgets/alert_banner.py @@ -12,6 +12,7 @@ class AlertBanner(QFrame): super().__init__(parent) self._current_message: str | None = None + self._current_is_error: bool | None = None self._clear_timer = QTimer(self) self._clear_timer.setSingleShot(True) @@ -42,7 +43,7 @@ class AlertBanner(QFrame): self.clear_message() return - if msg == self._current_message: + if self._current_message == msg and self._current_is_error == is_error: return if is_error: @@ -62,6 +63,8 @@ class AlertBanner(QFrame): "}" ) else: + if msg == self._current_message: + return decorated = f"✅ {msg} ✅" self.setStyleSheet( "QFrame {" @@ -79,12 +82,15 @@ class AlertBanner(QFrame): ) self._clear_timer.start(5000) + self._current_message = msg + self._current_is_error = is_error self._label.setText(decorated) self.setVisible(True) @Slot() def clear_message(self): self._current_message = None + self._current_is_error = None self._clear_timer.stop() self._label.clear() self.setVisible(False) -- 2.54.0 From 9a0f7abba0a73af79769fdb446eead09bdf3b7ca Mon Sep 17 00:00:00 2001 From: appleb_m Date: Tue, 24 Mar 2026 11:11:02 +0100 Subject: [PATCH 21/56] BIG AEROTECH UPGRADE: transferred to aerotech api. Ported several functions from Coordinate to aerotech coordinate. Nees thorough testing as this could lead to strange behaviour --- src/aare/daq/config.py | 18 ++- src/aare/daq/daq.py | 199 +++++++++++++++++++++---- src/aare/daq/devices.py | 55 ++++--- src/aare/daq/server.py | 5 +- src/aare/devices/aerotech.py | 16 +- src/aare/gui/panels/abr_tweak_panel.py | 12 +- src/aare/gui/threads/daq_worker.py | 6 +- 7 files changed, 232 insertions(+), 79 deletions(-) diff --git a/src/aare/daq/config.py b/src/aare/daq/config.py index 61196563..4b353ae0 100644 --- a/src/aare/daq/config.py +++ b/src/aare/daq/config.py @@ -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, @@ -23,8 +23,13 @@ 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,6 +79,7 @@ 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 @@ -458,16 +464,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 diff --git a/src/aare/daq/daq.py b/src/aare/daq/daq.py index 67c47868..1a105633 100644 --- a/src/aare/daq/daq.py +++ b/src/aare/daq/daq.py @@ -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 @@ -358,10 +360,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 @@ -426,7 +428,6 @@ class AareDAQ: 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() @@ -691,30 +692,97 @@ 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, + index=0, + mos=0.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}') + 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 + #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.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 + #result = self.__jfjoch.wait_till_done(60) - return None - #return CompletedRasterGridElem(request=copy.deepcopy(r), result=, centre_of_mass=None) + #return None + + return self._build_fake_raster_result(request=request) + #return CompletedRasterGridElem(request=copy.deepcopy(request), result=result, centre_of_mass=None) def measure_raster(self, r: RasterGridRequest, auto: bool) -> CompletedRasterGrid: self.__cfg.try_set_busy(timeout=ceil(360)) @@ -723,25 +791,98 @@ 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.") + # result = self._build_fake_rotation_result(request) + + #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: + #logger.info(f"rotation scan {request}") + 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, + ) + #logger.info(f"rotation scan in progress {request}") + #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) + + #logger.info(f"wait till rotation scan done {request}") + self.__devs.aerotech.wait_till_done(timeout=int(round(total_time + 60,0))) + #logger.info(f"move aerotech to omega start: {omega_start}") + self.__devs.aerotech_omega = omega_start + + #result = self.__jfjoch.wait_till_done(60) + # self.__aare.sample_collected(self.sample) + + # if result is None: + # logger.warning("JFJoch returned no ScanResult; using fake result for rotation scan.") + # result = self._build_fake_rotation_result(request) + + 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 Exception as e: + + logger.error(f"Exception during rotation scan: {e}") + #result = self._build_fake_rotation_result(request) + + return self._build_fake_rotation_result(request) + #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: @@ -813,8 +954,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), diff --git a/src/aare/daq/devices.py b/src/aare/daq/devices.py index 0455beea..bf2ee040 100644 --- a/src/aare/daq/devices.py +++ b/src/aare/daq/devices.py @@ -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") @@ -262,8 +257,6 @@ class BeamlineDevices: """ return int(self.__sample_cam.uid.get()) - - # Detector Z @property def dtz(self) -> float: @@ -284,32 +277,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: diff --git a/src/aare/daq/server.py b/src/aare/daq/server.py index faf1bfbb..2665f44d 100644 --- a/src/aare/daq/server.py +++ b/src/aare/daq/server.py @@ -7,7 +7,7 @@ import json import cv2 import urllib3 import uvicorn -from aare.common.coordinate import SmargonCoordinate, Coordinate +from aare.common.coordinate import SmargonCoordinate, Coordinate, AerotechCoordinate from aare.common.error_codes import export_error_codes, export_error_codes_grouped from aare.common.logger_config import setup_logger from aare.common.models import SampleShortInfo, DAQStatusModel, BeamlineStateEnum, BeamlineSettingsModel, \ @@ -205,8 +205,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" diff --git a/src/aare/devices/aerotech.py b/src/aare/devices/aerotech.py index 47ae6cad..2d236f7e 100644 --- a/src/aare/devices/aerotech.py +++ b/src/aare/devices/aerotech.py @@ -44,7 +44,7 @@ class AerotechController(object): 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, - var_async=not wait, + run_async=not wait, incremental=incremental, ) @@ -112,12 +112,12 @@ class AerotechController(object): def rotation_scan(self, rotation_deg: float | int, time_sec: float | int, start_pos_deg: float | int, - var_async: bool = False): + run_async: bool = False): payload = RotationRequest( rotation_deg=rotation_deg, time_sec=time_sec, start_pos_deg=start_pos_deg, - var_async=var_async + run_async=run_async ) if self.__simulated: return payload @@ -130,16 +130,16 @@ class AerotechController(object): time_sec: int|float, grid_elem_size_x_um: Optional[Union[float, int]] = None, grid_elem_count_x: Optional[int] = None, - var_async: Optional[bool] = False + 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, - var_async=var_async, + run_async=run_async, ) if self.__simulated: return payload @@ -150,14 +150,14 @@ class AerotechController(object): wedge_deg: float | int, time_sec: float | int, steps: int, - var_async: bool = False + run_async: bool = False ): payload = ScreenRequest( rotation_deg=rotation_deg, wedge_deg=wedge_deg, time_sec=time_sec, steps=steps, - var_async=var_async + run_async=run_async ) if self.__simulated: return payload diff --git a/src/aare/gui/panels/abr_tweak_panel.py b/src/aare/gui/panels/abr_tweak_panel.py index 8741f0dd..9eb589cf 100644 --- a/src/aare/gui/panels/abr_tweak_panel.py +++ b/src/aare/gui/panels/abr_tweak_panel.py @@ -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() diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py index 36fb7464..619d6cec 100644 --- a/src/aare/gui/threads/daq_worker.py +++ b/src/aare/gui/threads/daq_worker.py @@ -8,7 +8,7 @@ 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 +from aare.common.coordinate import SmargonCoordinate, Coordinate, AerotechCoordinate from aare.common.error_codes import export_error_codes from aare.common.models import DAQStatusModel, SampleShortInfoList, SampleShortInfo, SampleCameraSettings, \ AutofocusSettings, SimpleScanParameters, FluorescenceSpectrumParameterModel, FluorescenceSpectrumOutputModel @@ -800,8 +800,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() -- 2.54.0 From 130f275197084465a1a7f84dfede5f5fae525933 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Tue, 24 Mar 2026 11:11:52 +0100 Subject: [PATCH 22/56] devices: added magnet position sensor functions and moved set_transmission function --- src/aare/daq/devices.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/src/aare/daq/devices.py b/src/aare/daq/devices.py index bf2ee040..78fdd330 100644 --- a/src/aare/daq/devices.py +++ b/src/aare/daq/devices.py @@ -86,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: -- 2.54.0 From ae52d411e571384938b6428946b567542fc46215 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Tue, 24 Mar 2026 11:12:14 +0100 Subject: [PATCH 23/56] jfjoch.py: added simulated mode. --- src/aare/devices/jfjoch.py | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/src/aare/devices/jfjoch.py b/src/aare/devices/jfjoch.py index 964cd761..e2d6f59a 100644 --- a/src/aare/devices/jfjoch.py +++ b/src/aare/devices/jfjoch.py @@ -10,6 +10,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 +18,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") @@ -143,7 +148,9 @@ class JFJochWrapper: ) self.__api.start_post(dataset_settings=dataset_settings) - def wait_till_done(self, timeout : int | float) -> jfjoch_client.models.ScanResult: + def wait_till_done(self, timeout : int | float) -> jfjoch_client.models.ScanResult | None: + if self.__simulated: + return None self.__api.wait_till_done_post_with_http_info(timeout=math.ceil(timeout)) return self.__api.result_scan_get() -- 2.54.0 From ecf4ccffb3c3d8d9e77d69e7912bd74969093289 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Tue, 24 Mar 2026 11:12:43 +0100 Subject: [PATCH 24/56] daq_worker: added missing aerotech and tell changes to status respone handler --- src/aare/gui/threads/daq_worker.py | 12 ++++++++---- 1 file changed, 8 insertions(+), 4 deletions(-) diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py index 619d6cec..c0c68f1d 100644 --- a/src/aare/gui/threads/daq_worker.py +++ b/src/aare/gui/threads/daq_worker.py @@ -267,13 +267,17 @@ class DAQWorker(QObject): 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) - tell_err_text = None if tell_err is None else str(tell_err).strip() + aerotech_conn = bool(getattr(parsed_response, "aerotech_connected", True)) + aerotech_err = getattr(parsed_response, "aerotech_error", None) + 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() tell_changed = ( - self._last_tell_connected is None - or self._last_tell_connected != tell_conn - or self._last_tell_error != tell_err_text + self._last_tell_connected is None + or self._last_tell_connected != tell_conn + or self._last_tell_error != tell_err_text ) smargon_changed = ( self._last_smargon_connected is None -- 2.54.0 From 3b440cd1fa18d696517345b9e3a957f0a6e078d0 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Tue, 24 Mar 2026 11:13:09 +0100 Subject: [PATCH 25/56] aerotech models --- src/aare/common/aerotech_models.py | 155 +++++++++++++++++++++++++++++ 1 file changed, 155 insertions(+) create mode 100644 src/aare/common/aerotech_models.py diff --git a/src/aare/common/aerotech_models.py b/src/aare/common/aerotech_models.py new file mode 100644 index 00000000..786ea1c3 --- /dev/null +++ b/src/aare/common/aerotech_models.py @@ -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) \ No newline at end of file -- 2.54.0 From af4cd798d45e3006ddba13f32ea45b3e1ebe7f39 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Tue, 24 Mar 2026 12:44:14 +0100 Subject: [PATCH 26/56] Daq_worker - updating error message text for users if devices disconnect --- src/aare/gui/threads/daq_worker.py | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py index c0c68f1d..0787d2c9 100644 --- a/src/aare/gui/threads/daq_worker.py +++ b/src/aare/gui/threads/daq_worker.py @@ -226,28 +226,28 @@ class DAQWorker(QObject): if len(disconnected) == 3: return ( "all-devices-down", - "TELL, Smargon, and Aerotech connection errors, please inform your local contact.", + "TELL, Smargon, and Aerotech disconnected.", True, ) if len(disconnected) == 2: return ( "+".join(sorted(d.lower() for d in disconnected)) + "-down", - f"{disconnected[0]} and {disconnected[1]} connection errors, please inform your local contact.", + 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 connections restored.", False) + return (None, "TELL, Smargon, and Aerotech reconnected.", False) if len(restored) == 2: - return (None, f"{restored[0]} and {restored[1]} connections restored.", False) - return (None, f"{restored[0]} connection restored.", False) + return (None, f"{restored[0]} and {restored[1]} reconnected.", False) + return (None, f"{restored[0]} reconnected.", False) return (None, None, False) @@ -259,7 +259,7 @@ class DAQWorker(QObject): self.update.emit(parsed_response) if self._server_connected is False: - self._emit_status_if_changed(None, "Server connection restored.", False) + self._emit_status_if_changed(None, "Server reconnected.", False) self._server_connected = True self._last_server_error = None @@ -320,7 +320,7 @@ class DAQWorker(QObject): if self._server_connected is not False: self._emit_status_if_changed( "server-down", - "Server connection lost. Trying to reconnect...", + "Server disconnected. Reconnecting...", True, ) -- 2.54.0 From 2e4adb5eff521d862ae2a13822a7166ce7fe1d55 Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Tue, 24 Mar 2026 15:01:46 +0100 Subject: [PATCH 27/56] Enable mTLS support for WebSocket connections in `tellupdater` by configuring certificate and key files. --- src/aare/daq/tellupdater.py | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) diff --git a/src/aare/daq/tellupdater.py b/src/aare/daq/tellupdater.py index 2616ff4e..6addb2b5 100644 --- a/src/aare/daq/tellupdater.py +++ b/src/aare/daq/tellupdater.py @@ -147,6 +147,7 @@ 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) @@ -163,7 +164,16 @@ def main(): on_close=on_close, on_open=on_open, ) - ws.run_forever(sslopt={"cert_reqs": 0}) + + # --- mTLS CERTIFICATES ADDED HERE -- + import ssl + ssl_settings = { + "certfile": "/etc/ssl/certs/secrets/mx-x10sa-queue-01.crt", + "keyfile": "/etc/ssl/certs/secrets/mx-x10sa-queue-01.key", + "cert_reqs": ssl.CERT_REQUIRED + } + + ws.run_forever(sslopt=ssl_settings) except Exception as e: print(f"[MAIN][ERROR] ws.run_forever() crashed with: {e}") -- 2.54.0 From a9650876690bf2967a0f5a4dccdaec5f75a0c456 Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Tue, 24 Mar 2026 15:33:42 +0100 Subject: [PATCH 28/56] Add CA certificate configuration to SSL settings in `tellupdater` to complete mTLS setup --- src/aare/daq/tellupdater.py | 1 + 1 file changed, 1 insertion(+) diff --git a/src/aare/daq/tellupdater.py b/src/aare/daq/tellupdater.py index 6addb2b5..6c21f535 100644 --- a/src/aare/daq/tellupdater.py +++ b/src/aare/daq/tellupdater.py @@ -170,6 +170,7 @@ def main(): ssl_settings = { "certfile": "/etc/ssl/certs/secrets/mx-x10sa-queue-01.crt", "keyfile": "/etc/ssl/certs/secrets/mx-x10sa-queue-01.key", + "ca_certs": "/etc/ssl/certs/secrets/mx-db-01_DigiCert_Global_Root_G2.pem", "cert_reqs": ssl.CERT_REQUIRED } -- 2.54.0 From 6a0761250980111e6dfbb7627cafd2772593830f Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Tue, 24 Mar 2026 15:39:21 +0100 Subject: [PATCH 29/56] Enable mandatory mTLS for SSE connections in `tellupdater` by configuring client certificates and CA root files. --- src/aare/daq/tellupdater.py | 19 ++++++++++++++++--- 1 file changed, 16 insertions(+), 3 deletions(-) diff --git a/src/aare/daq/tellupdater.py b/src/aare/daq/tellupdater.py index 6c21f535..13bf63af 100644 --- a/src/aare/daq/tellupdater.py +++ b/src/aare/daq/tellupdater.py @@ -38,12 +38,25 @@ 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" - + while True: try: print(f"[SSE][INFO] Attempting to connect to {sse_url}...") - # Use URL directly for sseclient variants that manage their own HTTP stream. - client = sseclient.SSEClient(sse_url) + + # --- MANDATORY mTLS FOR mx-db-01 --- + import requests + cert_pair = ( + "/etc/ssl/certs/secrets/mx-x10sa-queue-01.crt", + "/etc/ssl/certs/secrets/mx-x10sa-queue-01.key" + ) + # The DigiCert root you have on the machine + ca_root = "/etc/ssl/certs/secrets/mx-db-01_DigiCert_Global_Root_G2.pem" + + # Open the stream using the certificates required by Nginx + response = requests.get(sse_url, stream=True, cert=cert_pair, verify=ca_root) + response.raise_for_status() + + client = sseclient.SSEClient(response) print("[SSE][listen_to_sse] Initial detected pucks fetch on connect") handle_tell_change_event() -- 2.54.0 From 4e873fefc358c8346d8e2b0538d17dbb42d54c90 Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Tue, 24 Mar 2026 16:04:34 +0100 Subject: [PATCH 30/56] Update CA certificate path in `tellupdater` to use the full chain file for improved SSL configuration --- src/aare/daq/tellupdater.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/aare/daq/tellupdater.py b/src/aare/daq/tellupdater.py index 13bf63af..f4e796ca 100644 --- a/src/aare/daq/tellupdater.py +++ b/src/aare/daq/tellupdater.py @@ -50,7 +50,7 @@ def listen_to_sse(): "/etc/ssl/certs/secrets/mx-x10sa-queue-01.key" ) # The DigiCert root you have on the machine - ca_root = "/etc/ssl/certs/secrets/mx-db-01_DigiCert_Global_Root_G2.pem" + ca_root = "/etc/ssl/certs/secrets/mx-db-01_Full_Chain_CA.pem" # Open the stream using the certificates required by Nginx response = requests.get(sse_url, stream=True, cert=cert_pair, verify=ca_root) @@ -183,7 +183,7 @@ def main(): ssl_settings = { "certfile": "/etc/ssl/certs/secrets/mx-x10sa-queue-01.crt", "keyfile": "/etc/ssl/certs/secrets/mx-x10sa-queue-01.key", - "ca_certs": "/etc/ssl/certs/secrets/mx-db-01_DigiCert_Global_Root_G2.pem", + "ca_certs": "/etc/ssl/certs/secrets/mx-db-01_Full_Chain_CA.pem", "cert_reqs": ssl.CERT_REQUIRED } -- 2.54.0 From 87df42a179ded99705baa6dec4745abccba23c3f Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Tue, 24 Mar 2026 16:17:45 +0100 Subject: [PATCH 31/56] Conditionally apply mTLS for HTTPS connections in `tellupdater` to avoid certificate usage on plain HTTP requests. --- src/aare/daq/tellupdater.py | 20 ++++++++++++-------- 1 file changed, 12 insertions(+), 8 deletions(-) diff --git a/src/aare/daq/tellupdater.py b/src/aare/daq/tellupdater.py index f4e796ca..14f47bf6 100644 --- a/src/aare/daq/tellupdater.py +++ b/src/aare/daq/tellupdater.py @@ -45,15 +45,19 @@ def listen_to_sse(): # --- MANDATORY mTLS FOR mx-db-01 --- import requests - cert_pair = ( - "/etc/ssl/certs/secrets/mx-x10sa-queue-01.crt", - "/etc/ssl/certs/secrets/mx-x10sa-queue-01.key" - ) - # The DigiCert root you have on the machine - ca_root = "/etc/ssl/certs/secrets/mx-db-01_Full_Chain_CA.pem" + # --- FIX: Only use certs for HTTPS (DB) connections --- + if sse_url.startswith("https:"): + cert_pair = ( + "/etc/ssl/certs/secrets/mx-x10sa-queue-01.crt", + "/etc/ssl/certs/secrets/mx-x10sa-queue-01.key" + ) + # Use the combined bundle (Intermediate + Root) if you made it + ca_root = "/etc/ssl/certs/secrets/mx-db-01_DigiCert_Global_Root_G2.pem" + response = requests.get(sse_url, stream=True, cert=cert_pair, verify=ca_root) + else: + # Robot connection (PC17488): No certs allowed on plain HTTP + response = requests.get(sse_url, stream=True) - # Open the stream using the certificates required by Nginx - response = requests.get(sse_url, stream=True, cert=cert_pair, verify=ca_root) response.raise_for_status() client = sseclient.SSEClient(response) -- 2.54.0 From a4b6c2575ee3296244a7ef1a4e7ea34cc3580589 Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Tue, 24 Mar 2026 17:18:42 +0100 Subject: [PATCH 32/56] Update certificate and key file paths in `tellupdater` to use database-specific mTLS credentials --- src/aare/daq/tellupdater.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/aare/daq/tellupdater.py b/src/aare/daq/tellupdater.py index 14f47bf6..b1394475 100644 --- a/src/aare/daq/tellupdater.py +++ b/src/aare/daq/tellupdater.py @@ -185,8 +185,8 @@ def main(): # --- mTLS CERTIFICATES ADDED HERE -- import ssl ssl_settings = { - "certfile": "/etc/ssl/certs/secrets/mx-x10sa-queue-01.crt", - "keyfile": "/etc/ssl/certs/secrets/mx-x10sa-queue-01.key", + "certfile": "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_db-01.crt", + "keyfile": "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_db-01.key", "ca_certs": "/etc/ssl/certs/secrets/mx-db-01_Full_Chain_CA.pem", "cert_reqs": ssl.CERT_REQUIRED } -- 2.54.0 From d6d20d598ff4133acea579de0d4fbf3ad6667222 Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Tue, 24 Mar 2026 17:28:11 +0100 Subject: [PATCH 33/56] Simplify mTLS handling in `tellupdater` by consolidating certificate paths and streamlining HTTP/HTTPS logic. --- src/aare/daq/tellupdater.py | 27 ++++++++++----------------- 1 file changed, 10 insertions(+), 17 deletions(-) diff --git a/src/aare/daq/tellupdater.py b/src/aare/daq/tellupdater.py index b1394475..9ab979ba 100644 --- a/src/aare/daq/tellupdater.py +++ b/src/aare/daq/tellupdater.py @@ -42,25 +42,18 @@ def listen_to_sse(): while True: try: print(f"[SSE][INFO] Attempting to connect to {sse_url}...") - - # --- MANDATORY mTLS FOR mx-db-01 --- - import requests - # --- FIX: Only use certs for HTTPS (DB) connections --- - if sse_url.startswith("https:"): - cert_pair = ( - "/etc/ssl/certs/secrets/mx-x10sa-queue-01.crt", - "/etc/ssl/certs/secrets/mx-x10sa-queue-01.key" - ) - # Use the combined bundle (Intermediate + Root) if you made it - ca_root = "/etc/ssl/certs/secrets/mx-db-01_DigiCert_Global_Root_G2.pem" + if sse_url.startswith("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") + ca_root = "/etc/ssl/certs/secrets/mx-db-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 connection (PC17488): No certs allowed on plain HTTP - response = requests.get(sse_url, stream=True) - - response.raise_for_status() - - client = sseclient.SSEClient(response) + # Robot path (PC17488) - Pass URL STRING directly + # SSEClient will handle the simple HTTP GET itself + client = sseclient.SSEClient(sse_url) print("[SSE][listen_to_sse] Initial detected pucks fetch on connect") handle_tell_change_event() -- 2.54.0 From eaf3083eca6142e2c263fb1ffe8923d978486985 Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Tue, 24 Mar 2026 20:35:58 +0100 Subject: [PATCH 34/56] Bump `aaredb` dependency version to `0.1.1a43` in `pyproject.toml`. --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index 044818a6..e4529c7e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", -- 2.54.0 From 0225bb1e3b26eb7c0a0b06d464968e5857c670a3 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Wed, 25 Mar 2026 13:21:02 +0100 Subject: [PATCH 35/56] Exception handler p- added exceptions for JFJOCH, Aerotech and AareDB communication errors --- src/aare/common/exception_handler.py | 101 +++++++++++++++++++++++++++ 1 file changed, 101 insertions(+) diff --git a/src/aare/common/exception_handler.py b/src/aare/common/exception_handler.py index 9a81665d..7012fc78 100644 --- a/src/aare/common/exception_handler.py +++ b/src/aare/common/exception_handler.py @@ -195,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 \ No newline at end of file -- 2.54.0 From 624a8365a99efe5fb560c54dfcd1b668e95691cf Mon Sep 17 00:00:00 2001 From: appleb_m Date: Wed, 25 Mar 2026 13:21:38 +0100 Subject: [PATCH 36/56] Sever_exception_handler: added handlers for JFJOCH and Aerotech errors --- src/aare/daq/server_exception_handler.py | 37 ++++++++++++++++++++++-- 1 file changed, 35 insertions(+), 2 deletions(-) diff --git a/src/aare/daq/server_exception_handler.py b/src/aare/daq/server_exception_handler.py index 530864db..7b4307df 100644 --- a/src/aare/daq/server_exception_handler.py +++ b/src/aare/daq/server_exception_handler.py @@ -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"), -- 2.54.0 From 3d4e252438511085285629ddd3d992506a862a61 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Wed, 25 Mar 2026 13:23:20 +0100 Subject: [PATCH 37/56] Fixes to alert banner --- src/aare/gui/main_window.py | 33 ++++---- src/aare/gui/threads/daq_worker.py | 112 +++++++++++++++++++++------ src/aare/gui/widgets/alert_banner.py | 9 ++- 3 files changed, 105 insertions(+), 49 deletions(-) diff --git a/src/aare/gui/main_window.py b/src/aare/gui/main_window.py index 05069a1d..9d05a250 100644 --- a/src/aare/gui/main_window.py +++ b/src/aare/gui/main_window.py @@ -70,6 +70,7 @@ class MainWindow(QMainWindow): gonio_cam_id: int | None ): super().__init__() + self.__base_url = base_url self.__token = token self.__mounting = False @@ -78,8 +79,6 @@ class MainWindow(QMainWindow): self._controls_help_dialog = None self._cleanup_done = False self._default_window_state = None - self._last_status_message: str | None = None - self._last_status_is_error: bool | None = None # Tutorial manager (define tutorials after widgets exist) self._tutorial_event_bus = TutorialEventBus(self) @@ -115,9 +114,12 @@ class MainWindow(QMainWindow): root_layout.setContentsMargins(0, 0, 0, 0) root_layout.setSpacing(0) - self.alert_banner = AlertBanner(parent=root_widget) + self.alert_banner = AlertBanner(parent=root_widget, error_timeout_ms=30000, recover_timeout_ms=5000) root_layout.addWidget(self.alert_banner) + self.alert_banner_secondary = AlertBanner(parent=root_widget, error_timeout_ms=10000, recover_timeout_ms=5000) + 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) @@ -483,27 +485,18 @@ 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._show_status_message_once) register_tutorials(self, self.tutorial_manager) - @Slot(str, bool) - def _show_status_message_once(self, message: str, is_error: bool = True) -> None: - message = (message or "").strip() - if not message: - self.alert_banner.clear_message() - self._last_status_message = None - self._last_status_is_error = None - return - - if message == self._last_status_message and is_error == self._last_status_is_error: - return - - self._last_status_message = message - self._last_status_is_error = is_error - self.alert_banner.show_message(message, is_error=is_error) - @Slot(QPixmap) def _on_samcam_prediction_pixmap(self, pix: QPixmap) -> None: self._last_pred_image_ts = time.monotonic() diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py index 0787d2c9..43cdc327 100644 --- a/src/aare/gui/threads/daq_worker.py +++ b/src/aare/gui/threads/daq_worker.py @@ -10,9 +10,10 @@ from jfjoch_client import ScanResult, ScanResultImagesInner from aare.common.coordinate import SmargonCoordinate, Coordinate, AerotechCoordinate 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 +29,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) @@ -86,6 +92,17 @@ class DAQWorker(QObject): 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._face_detection_stream_reply: QNetworkReply | None = None if self.__base_url is not None: @@ -105,6 +122,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. @@ -180,20 +237,6 @@ 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, *, @@ -258,8 +301,15 @@ class DAQWorker(QObject): parsed_response = DAQStatusModel.model_validate_json(response_data) self.update.emit(parsed_response) - if self._server_connected is 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 @@ -274,20 +324,29 @@ class DAQWorker(QObject): 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 is None - or self._last_tell_connected != tell_conn - or self._last_tell_error != tell_err_text + self._last_tell_connected != tell_conn + or self._last_tell_error != tell_err_text ) smargon_changed = ( - self._last_smargon_connected is None - or self._last_smargon_connected != smargon_conn - or self._last_smargon_error != smargon_err_text + self._last_smargon_connected != smargon_conn + or self._last_smargon_error != smargon_err_text ) aerotech_changed = ( - self._last_aerotech_connected is None - or self._last_aerotech_connected != aerotech_conn - or self._last_aerotech_error != aerotech_err_text + self._last_aerotech_connected != aerotech_conn + or self._last_aerotech_error != aerotech_err_text ) self._last_tell_connected = tell_conn @@ -318,6 +377,8 @@ class DAQWorker(QObject): 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 disconnected. Reconnecting...", @@ -326,6 +387,7 @@ class DAQWorker(QObject): self._server_connected = False self._last_server_error = err_msg + # Clear device states so we show "Server reconnected" on recovery self._last_tell_connected = None self._last_smargon_connected = None self._last_aerotech_connected = None diff --git a/src/aare/gui/widgets/alert_banner.py b/src/aare/gui/widgets/alert_banner.py index 86e5c791..a0ea3f78 100644 --- a/src/aare/gui/widgets/alert_banner.py +++ b/src/aare/gui/widgets/alert_banner.py @@ -8,11 +8,13 @@ logger = setup_logger("aareGUI") class AlertBanner(QFrame): - def __init__(self, parent=None): + def __init__(self, parent=None, error_timeout_ms: int = 15000, recover_timeout_ms: int = 5000): super().__init__(parent) self._current_message: str | None = None self._current_is_error: bool | None = None + self._error_timeout_ms = error_timeout_ms + self._recovery_timeout_ms = recover_timeout_ms self._clear_timer = QTimer(self) self._clear_timer.setSingleShot(True) @@ -62,9 +64,8 @@ class AlertBanner(QFrame): " padding: 2px 6px 2px 6px;" "}" ) + self._clear_timer.start(self._error_timeout_ms) else: - if msg == self._current_message: - return decorated = f"✅ {msg} ✅" self.setStyleSheet( "QFrame {" @@ -80,7 +81,7 @@ class AlertBanner(QFrame): " padding: 2px 6px 2px 6px;" "}" ) - self._clear_timer.start(5000) + self._clear_timer.start(self._recovery_timeout_ms) self._current_message = msg self._current_is_error = is_error -- 2.54.0 From 98f665370a6cf680eb567ff9090bd9a43c939bb1 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Wed, 25 Mar 2026 13:23:44 +0100 Subject: [PATCH 38/56] Aerotech: added custom exception for communication error --- src/aare/devices/aerotech.py | 98 +++++++++++++++++++++++++++++++----- 1 file changed, 86 insertions(+), 12 deletions(-) diff --git a/src/aare/devices/aerotech.py b/src/aare/devices/aerotech.py index 2d236f7e..463087f7 100644 --- a/src/aare/devices/aerotech.py +++ b/src/aare/devices/aerotech.py @@ -7,6 +7,7 @@ from aare.common.beamline import MXBeamline, mx_beamline from aarescan_client import ApiClient, Status, RotationRequest, Configuration, DefaultApi, AxisStatus, Target from aare.common.coordinate import AerotechCoordinate, Coordinate +from aare.common.exception_handler import AerotechCommunicationError AEROTECH_HOME = AerotechCoordinate(at_mm=Coordinate(x=0, y=0, z=0), omega_deg=0) @@ -52,14 +53,30 @@ class AerotechController(object): return AerotechCoordinate(at_mm=Coordinate(x=target.x,y=target.y,z=target.z), omega_deg=target.u) def cancel(self): - self.__api.cancel_post() + try: + self.__api.cancel_post() + except Exception as e: + raise AerotechCommunicationError( + "Aerotech cancel failed", + endpoint="cancel_post", + base_url=self.__base, + operation="POST", + ) from e def is_idle(self) -> bool: - status = self.__api.status_get() - return status.state == 'Idle' + try: + status = self.__api.status_get() + return status.state == 'Idle' + except Exception as e: + raise AerotechCommunicationError( + "Aerotech status check failed", + endpoint="status_get", + base_url=self.__base, + operation="GET", + ) from e def get_position(self) -> AerotechCoordinate: - status = self.__api.status_get() + status = self.status() return AerotechCoordinate( at_mm=Coordinate( x=status.x.pos, @@ -69,7 +86,6 @@ class AerotechController(object): omega_deg=status.u.pos, ) - def status(self) -> Status: if self.__simulated: return Status(state=Status.State.IDLE, @@ -82,7 +98,15 @@ class AerotechController(object): u=AxisStatus(pos=self.__pos.u, vel=self.__vel, enabled=False, homed=False, moving=False,fault=False), ) - return self.__api.status_get() + 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_home(self, wait:bool=True, incremental:bool=False): if self.__simulated: @@ -91,10 +115,26 @@ class AerotechController(object): return self.position(AEROTECH_HOME, wait=wait, incremental=incremental) def home_aerotech(self): - return self.__api.home_post() + 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): - return self.__api.wait_till_done_post(timeout=timeout) + 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, @@ -107,7 +147,15 @@ class AerotechController(object): return self.__pos payload = self.__make_aerotech_target(target, wait=wait, incremental=incremental) - return self.__api.position_post(payload) + 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, @@ -122,7 +170,16 @@ class AerotechController(object): if self.__simulated: return payload - return self.__api.rotation_scan_post(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, @@ -143,7 +200,15 @@ class AerotechController(object): ) if self.__simulated: return payload - return self.__api.grid_scan_post(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, @@ -161,7 +226,15 @@ class AerotechController(object): ) if self.__simulated: return payload - return self.__api.screening_post(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 if __name__ == "__main__": @@ -170,3 +243,4 @@ if __name__ == "__main__": controller = AerotechController(beamline) #controller.print_status(colored=True, compact=True) print(controller.get_position()) + controller.cancel() -- 2.54.0 From fbc9890fe4a830937ba1c4b051711119ceed3e27 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Wed, 25 Mar 2026 13:24:05 +0100 Subject: [PATCH 39/56] jfjoch: added custom excpetions --- src/aare/devices/jfjoch.py | 36 ++++++++++++++++++++++++++++++++---- 1 file changed, 32 insertions(+), 4 deletions(-) diff --git a/src/aare/devices/jfjoch.py b/src/aare/devices/jfjoch.py index e2d6f59a..e0671e8b 100644 --- a/src/aare/devices/jfjoch.py +++ b/src/aare/devices/jfjoch.py @@ -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 @@ -98,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, @@ -146,13 +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 | None: if self.__simulated: return None - self.__api.wait_till_done_post_with_http_info(timeout=math.ceil(timeout)) - return self.__api.result_scan_get() + 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() -- 2.54.0 From c9ec2e50289e3275c12ee1bc76112d1905bb2b9f Mon Sep 17 00:00:00 2001 From: appleb_m Date: Wed, 25 Mar 2026 13:24:31 +0100 Subject: [PATCH 40/56] DAQ/Daq_worker - changes to rotation and raster as well as handling of responses --- src/aare/daq/daq.py | 108 ++++++++++++++++++----------- src/aare/gui/threads/daq_worker.py | 76 +++++++++++++++++--- 2 files changed, 132 insertions(+), 52 deletions(-) diff --git a/src/aare/daq/daq.py b/src/aare/daq/daq.py index 1a105633..57968015 100644 --- a/src/aare/daq/daq.py +++ b/src/aare/daq/daq.py @@ -46,7 +46,10 @@ from aare.common.exception_handler import ( MountingFailed, WarningTellException, CriticalTellException, - AXCFailed, SmargonCommunicationError, TellCommunicationError, + AXCFailed, + SmargonCommunicationError, + TellCommunicationError, + JFJochCommunicationError ) logger = setup_logger("aareDAQ") @@ -474,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 @@ -727,8 +730,9 @@ class AareDAQ: efficiency=1.0, bkg=0.0, spots=0, + spots_low_res=0, + spots_indexed=0, index=0, - mos=0.0, b=0.0, ) for i in range(max(1, image_count)) @@ -765,24 +769,46 @@ class AareDAQ: logger.info(f'raster grid request: {request}') total_time = request.exp_time_s*request.n_x*request.n_y - #self.__jfjoch.measure_raster(r, status) + 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}") - 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) + else: + logger.info("Simulated detector mode enabled; returning fake zero raster result.") - self.__devs.aerotech.wait_till_done(timeout=int(round(total_time*2,0))) + 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_pos = self.__cfg.abr_meas_pos + self.__devs.aerotech.wait_till_done(timeout=int(round(total_time*2,0))) - #result = self.__jfjoch.wait_till_done(60) - #return None + self.__devs.aerotech_pos = self.__cfg.abr_meas_pos - return self._build_fake_raster_result(request=request) - #return CompletedRasterGridElem(request=copy.deepcopy(request), result=result, centre_of_mass=None) + 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)) @@ -809,20 +835,18 @@ class AareDAQ: raise e def __rotation(self, request: RotationScanRequest) -> CompletedRotationScan: - 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.") - # result = self._build_fake_rotation_result(request) - - #else: - #self.__jfjoch.measure_rotation(request, status, self.__cfg.xrf) + 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( @@ -833,14 +857,13 @@ class AareDAQ: run_async=True, ) else: - #logger.info(f"rotation scan {request}") 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, ) - #logger.info(f"rotation scan in progress {request}") + #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) @@ -852,33 +875,34 @@ class AareDAQ: ) time.sleep(smargon_time_step) - #logger.info(f"wait till rotation scan done {request}") self.__devs.aerotech.wait_till_done(timeout=int(round(total_time + 60,0))) - #logger.info(f"move aerotech to omega start: {omega_start}") self.__devs.aerotech_omega = omega_start - #result = self.__jfjoch.wait_till_done(60) # self.__aare.sample_collected(self.sample) - # if result is None: - # logger.warning("JFJoch returned no ScanResult; using fake result for rotation scan.") - # result = self._build_fake_rotation_result(request) + if self.__cfg.simulated_detector: + logger.warning("Detector in simulation mode, returning fake zero rotation result.") + result = self._build_fake_rotation_result(request) - if self.sample is not None and self.sample.db_id is not None: - self.save_screenshot_db(self.sample.db_id, "after_dc") + 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}") - #result = self._build_fake_rotation_result(request) + raise - return self._build_fake_rotation_result(request) - #return CompletedRotationScan(request=copy.deepcopy(request), result=result) + return CompletedRotationScan(request=copy.deepcopy(request), result=result) def measure_rotation(self, request: RotationScanRequest) -> CompletedRotationScan: total_time = request.exp_time_s * request.steps @@ -1879,7 +1903,7 @@ class AareDAQ: 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): diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py index 43cdc327..7984bf57 100644 --- a/src/aare/gui/threads/daq_worker.py +++ b/src/aare/gui/threads/daq_worker.py @@ -692,12 +692,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): @@ -712,12 +738,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): @@ -736,13 +790,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 @@ -769,13 +824,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 -- 2.54.0 From 100895f56bc8d27ba56a1fd6ad12b78b99943ac7 Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Wed, 25 Mar 2026 13:49:43 +0100 Subject: [PATCH 41/56] Added mTLS certificates authentication to aaredb --- src/aare/daq/aaredb.py | 12 ++++++++++-- src/aare/daq/spreadsheetupdater.py | 22 ++++++++++++++++++++-- src/aare/daq/tellupdater.py | 14 +++++++++----- 3 files changed, 39 insertions(+), 9 deletions(-) diff --git a/src/aare/daq/aaredb.py b/src/aare/daq/aaredb.py index e165f7b1..0ee85151 100644 --- a/src/aare/daq/aaredb.py +++ b/src/aare/daq/aaredb.py @@ -38,6 +38,7 @@ from jfjoch_client.models import ScanResult logger = setup_logger("aareDAQ") + class AareWrapper: def __init__( self, @@ -45,14 +46,21 @@ class AareWrapper: host: str = "https://mx-db-01.psi.ch/dispatcher", ): configuration = aareDB.Configuration(host=host) - configuration.verify_ssl = False # Disable SSL verification + + # --- mTLS & SSL CONFIGURATION --- + configuration.cert_file = "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_db-01.crt" + configuration.key_file = "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_db-01.key" + + # Point this to the CA bundle that signed the server's certificate + configuration.verify_ssl = True + configuration.ssl_ca_cert = "/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem" self.client = aareDB.ApiClient(configuration) 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 diff --git a/src/aare/daq/spreadsheetupdater.py b/src/aare/daq/spreadsheetupdater.py index 4f7c65f8..c1cd28af 100644 --- a/src/aare/daq/spreadsheetupdater.py +++ b/src/aare/daq/spreadsheetupdater.py @@ -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") @@ -142,7 +144,23 @@ def main(): on_close=on_close, on_open=on_open, ) - ws.run_forever(sslopt={"cert_reqs": 0}) + + # --- mTLS CONFIGURATION --- + import ssl + ssl_settings = { + # This is the "Machine Certificate" NGINX will verify + "certfile": "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_aaredb-dmz-01.crt", + "keyfile": "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_aaredb-dmz-01.key", + + # This allows Python to trust the Server (NGINX/PSI) + "ca_certs": "/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem", + + # Standard TLS requirements + "cert_reqs": ssl.CERT_REQUIRED, + "check_hostname": True, + } + + ws.run_forever(sslopt=ssl_settings) except Exception as e: print(f"[MAIN][ERROR] ws.run_forever() crashed with: {e}") diff --git a/src/aare/daq/tellupdater.py b/src/aare/daq/tellupdater.py index 9ab979ba..1023ad76 100644 --- a/src/aare/daq/tellupdater.py +++ b/src/aare/daq/tellupdater.py @@ -18,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) @@ -42,18 +43,21 @@ def listen_to_sse(): while True: try: print(f"[SSE][INFO] Attempting to connect to {sse_url}...") - if sse_url.startswith("https://mx-db-01"): + 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") - ca_root = "/etc/ssl/certs/secrets/mx-db-01_Full_Chain_CA.pem" + #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) + client = sseclient.SSEClient(sse_url) print("[SSE][listen_to_sse] Initial detected pucks fetch on connect") handle_tell_change_event() -- 2.54.0 From b24e3aa12aba6517aa4ced2b8c59ae9b2d3056d6 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Wed, 25 Mar 2026 13:51:39 +0100 Subject: [PATCH 42/56] ml_box: changed to new model yolo26l-seg-overlap-false_2026-03-16.engine --- src/aare/daq/mlbox.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/aare/daq/mlbox.py b/src/aare/daq/mlbox.py index 094233c2..84d00fc7 100644 --- a/src/aare/daq/mlbox.py +++ b/src/aare/daq/mlbox.py @@ -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}") -- 2.54.0 From 6e3e14a837367a7b9eb5b9a5a1b0dd646f804324 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Wed, 25 Mar 2026 13:57:10 +0100 Subject: [PATCH 43/56] aaredb: changed certification, host and logging statements --- src/aare/daq/aaredb.py | 42 ++++++++++++++++++++---------------------- 1 file changed, 20 insertions(+), 22 deletions(-) diff --git a/src/aare/daq/aaredb.py b/src/aare/daq/aaredb.py index 0ee85151..1e70683b 100644 --- a/src/aare/daq/aaredb.py +++ b/src/aare/daq/aaredb.py @@ -38,22 +38,20 @@ from jfjoch_client.models import ScanResult logger = setup_logger("aareDAQ") - 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) - # --- mTLS & SSL CONFIGURATION --- - configuration.cert_file = "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_db-01.crt" - configuration.key_file = "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_db-01.key" # Point this to the CA bundle that signed the server's certificate configuration.verify_ssl = True configuration.ssl_ca_cert = "/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem" + 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" self.client = aareDB.ApiClient(configuration) self.client.default_headers["X-Shared-Password"] = os.getenv("AAREDB_SHARED_PASSWORD") @@ -78,7 +76,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 @@ -92,7 +90,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: @@ -102,7 +100,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: @@ -112,7 +110,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: @@ -122,7 +120,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: @@ -133,7 +131,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: @@ -144,7 +142,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: @@ -155,7 +153,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: @@ -166,7 +164,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: @@ -296,9 +294,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: @@ -360,9 +358,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, @@ -386,7 +384,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, @@ -429,7 +427,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, @@ -453,7 +451,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, @@ -471,5 +469,5 @@ class AareWrapper: return payload except Exception as e: - print(e) + logger.error(e) raise e -- 2.54.0 From f1a9147f107727abe117e1212cbd1b59f10e0742 Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Wed, 25 Mar 2026 13:59:02 +0100 Subject: [PATCH 44/56] modify tellupdater mTLS certificates --- src/aare/daq/tellupdater.py | 24 +++++++++++++++--------- 1 file changed, 15 insertions(+), 9 deletions(-) diff --git a/src/aare/daq/tellupdater.py b/src/aare/daq/tellupdater.py index 1023ad76..c87e1059 100644 --- a/src/aare/daq/tellupdater.py +++ b/src/aare/daq/tellupdater.py @@ -170,6 +170,19 @@ def main(): # Main thread runs websocket client loop while True: try: + import ssl + # 1. Create a proper SSL Context + context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + + # 2. Load the CA to trust the server + context.load_verify_locations(cafile="/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem") + + # 3. Load the Client Certificate (mTLS) + 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" + ) + ws = websocket.WebSocketApp( WS_URL, header=WS_HEADERS, @@ -179,16 +192,9 @@ def main(): on_open=on_open, ) - # --- mTLS CERTIFICATES ADDED HERE -- - import ssl - ssl_settings = { - "certfile": "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_db-01.crt", - "keyfile": "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_db-01.key", - "ca_certs": "/etc/ssl/certs/secrets/mx-db-01_Full_Chain_CA.pem", - "cert_reqs": ssl.CERT_REQUIRED - } + # 4. Pass the context object to run_forever + ws.run_forever(sslopt={"context": context}) - ws.run_forever(sslopt=ssl_settings) except Exception as e: print(f"[MAIN][ERROR] ws.run_forever() crashed with: {e}") -- 2.54.0 From fd20ae50d805a0065c0778701d83d5f087f12131 Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Wed, 25 Mar 2026 14:08:00 +0100 Subject: [PATCH 45/56] modify tellupdater mTLS certificates --- src/aare/daq/aaredb.py | 13 +++++++++---- src/aare/daq/tellupdater.py | 21 ++++++++++++++------- 2 files changed, 23 insertions(+), 11 deletions(-) diff --git a/src/aare/daq/aaredb.py b/src/aare/daq/aaredb.py index 1e70683b..2d1b9785 100644 --- a/src/aare/daq/aaredb.py +++ b/src/aare/daq/aaredb.py @@ -46,14 +46,19 @@ class AareWrapper: ): configuration = aareDB.Configuration(host=host) - - # Point this to the CA bundle that signed the server's certificate + # --- 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" - 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" + # 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) diff --git a/src/aare/daq/tellupdater.py b/src/aare/daq/tellupdater.py index c87e1059..9792741d 100644 --- a/src/aare/daq/tellupdater.py +++ b/src/aare/daq/tellupdater.py @@ -167,22 +167,28 @@ def main(): 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 proper SSL Context + # 1. Create a modern SSL Context for a TLS Client context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) - # 2. Load the CA to trust the server + # 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 (mTLS) + # 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, @@ -192,13 +198,14 @@ def main(): on_open=on_open, ) - # 4. Pass the context object to run_forever + # 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] ws.run_forever() crashed with: {e}") + print(f"[MAIN][ERROR] WebSocket connection failed: {e}") - print("[WS][INFO] WebSocket connection lost. Reconnecting in 5 seconds...") + print("[WS][INFO] Reconnecting in 5 seconds...") time.sleep(5) if __name__ == "__main__": -- 2.54.0 From 25163a5b4dfc17c27d185247af510f5de600ddd7 Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Wed, 25 Mar 2026 14:10:17 +0100 Subject: [PATCH 46/56] modify spreadsheetupdater.py mTLS certificates --- src/aare/daq/spreadsheetupdater.py | 41 +++++++++++++++++------------- 1 file changed, 24 insertions(+), 17 deletions(-) diff --git a/src/aare/daq/spreadsheetupdater.py b/src/aare/daq/spreadsheetupdater.py index c1cd28af..204a8102 100644 --- a/src/aare/daq/spreadsheetupdater.py +++ b/src/aare/daq/spreadsheetupdater.py @@ -2,6 +2,7 @@ import os import json import websocket import time +import ssl from aareDB.models import PuckWithTellPosition from aare.common.models import SampleShortInfoList, SampleShortInfo, DewarAddress from config import BeamlineConfig @@ -136,6 +137,25 @@ def main(): """ while True: try: + # 1. Create the Context + # Use PROTOCOL_TLS_CLIENT as it's the modern standard + context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + + # 2. Load the Server Trust (CA) + context.load_verify_locations(cafile="/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem") + + # 3. Load the Client Identity (The key/cert that worked in 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" + ) + + # 4. Critical: If NGINX is using a self-signed or internal CA, + # and the hostname doesn't perfectly match the Cert SN/SAN, + # you might need to toggle this, but keep it True for now. + context.check_hostname = True + context.verify_mode = ssl.CERT_REQUIRED + ws = websocket.WebSocketApp( WS_URL, header=WS_HEADERS, @@ -145,26 +165,13 @@ def main(): on_open=on_open, ) - # --- mTLS CONFIGURATION --- - import ssl - ssl_settings = { - # This is the "Machine Certificate" NGINX will verify - "certfile": "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_aaredb-dmz-01.crt", - "keyfile": "/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_aaredb-dmz-01.key", + # 5. Connect using ONLY the context + # We remove other keys to ensure websocket-client uses our pre-configured context + ws.run_forever(sslopt={"context": context}) - # This allows Python to trust the Server (NGINX/PSI) - "ca_certs": "/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem", - - # Standard TLS requirements - "cert_reqs": ssl.CERT_REQUIRED, - "check_hostname": True, - } - - ws.run_forever(sslopt=ssl_settings) except Exception as e: - print(f"[MAIN][ERROR] ws.run_forever() crashed with: {e}") + print(f"[MAIN][ERROR] WebSocket setup failed: {e}") - print("[WS][INFO] WebSocket connection lost. Reconnecting in 5 seconds...") time.sleep(5) if __name__ == "__main__": -- 2.54.0 From 97d7e561251579b4af5901937348db238223ba66 Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Wed, 25 Mar 2026 14:21:07 +0100 Subject: [PATCH 47/56] modify spreadsheetupdater.py mTLS certificates --- src/aare/daq/spreadsheetupdater.py | 16 ++++------------ 1 file changed, 4 insertions(+), 12 deletions(-) diff --git a/src/aare/daq/spreadsheetupdater.py b/src/aare/daq/spreadsheetupdater.py index 204a8102..fca2ff7a 100644 --- a/src/aare/daq/spreadsheetupdater.py +++ b/src/aare/daq/spreadsheetupdater.py @@ -137,25 +137,18 @@ def main(): """ while True: try: - # 1. Create the Context - # Use PROTOCOL_TLS_CLIENT as it's the modern standard + # Initialize with the modern client protocol context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) - # 2. Load the Server Trust (CA) + # Trust the PSI DMZ CA context.load_verify_locations(cafile="/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem") - # 3. Load the Client Identity (The key/cert that worked in curl) + # Load your machine identity 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" ) - # 4. Critical: If NGINX is using a self-signed or internal CA, - # and the hostname doesn't perfectly match the Cert SN/SAN, - # you might need to toggle this, but keep it True for now. - context.check_hostname = True - context.verify_mode = ssl.CERT_REQUIRED - ws = websocket.WebSocketApp( WS_URL, header=WS_HEADERS, @@ -165,8 +158,7 @@ def main(): on_open=on_open, ) - # 5. Connect using ONLY the context - # We remove other keys to ensure websocket-client uses our pre-configured context + print(f"[WS][INFO] Connecting to {WS_URL}...") ws.run_forever(sslopt={"context": context}) except Exception as e: -- 2.54.0 From 12b749c4e63d9ff5f951b5269045afd7036d8ca2 Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Wed, 25 Mar 2026 14:24:48 +0100 Subject: [PATCH 48/56] modify spreadsheetupdater.py mTLS certificates --- src/aare/daq/spreadsheetupdater.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/src/aare/daq/spreadsheetupdater.py b/src/aare/daq/spreadsheetupdater.py index fca2ff7a..e7dc1ee3 100644 --- a/src/aare/daq/spreadsheetupdater.py +++ b/src/aare/daq/spreadsheetupdater.py @@ -149,6 +149,8 @@ def main(): keyfile="/etc/ssl/certs/secrets/mx-x10sa-queue-01_from_dmz-01.key" ) + context.check_hostname = True + ws = websocket.WebSocketApp( WS_URL, header=WS_HEADERS, @@ -158,7 +160,7 @@ def main(): on_open=on_open, ) - print(f"[WS][INFO] Connecting to {WS_URL}...") + print(f"[WS][INFO] Connecting to {WS_URL} using mTLS...") ws.run_forever(sslopt={"context": context}) except Exception as e: -- 2.54.0 From 0de5cbfc01d652a8d899a07a215b641045415c97 Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Wed, 25 Mar 2026 14:28:16 +0100 Subject: [PATCH 49/56] modify spreadsheetupdater.py mTLS certificates --- src/aare/daq/spreadsheetupdater.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/aare/daq/spreadsheetupdater.py b/src/aare/daq/spreadsheetupdater.py index e7dc1ee3..8f4b1f06 100644 --- a/src/aare/daq/spreadsheetupdater.py +++ b/src/aare/daq/spreadsheetupdater.py @@ -2,7 +2,6 @@ import os import json import websocket import time -import ssl from aareDB.models import PuckWithTellPosition from aare.common.models import SampleShortInfoList, SampleShortInfo, DewarAddress from config import BeamlineConfig @@ -137,6 +136,7 @@ def main(): """ while True: try: + import ssl # Initialize with the modern client protocol context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) -- 2.54.0 From d98769d67ca34b5d976b15ac8cc32524484c9ddf Mon Sep 17 00:00:00 2001 From: GotthardG <51994228+GotthardG@users.noreply.github.com> Date: Wed, 25 Mar 2026 14:32:43 +0100 Subject: [PATCH 50/56] modify spreadsheetupdater.py mTLS certificates --- src/aare/daq/spreadsheetupdater.py | 22 +++++++++++++--------- 1 file changed, 13 insertions(+), 9 deletions(-) diff --git a/src/aare/daq/spreadsheetupdater.py b/src/aare/daq/spreadsheetupdater.py index 8f4b1f06..5cad17fb 100644 --- a/src/aare/daq/spreadsheetupdater.py +++ b/src/aare/daq/spreadsheetupdater.py @@ -137,19 +137,23 @@ def main(): while True: try: import ssl - # Initialize with the modern client protocol - context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + import websocket - # Trust the PSI DMZ CA + # 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") - - # Load your machine identity 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" ) - context.check_hostname = True + # 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, @@ -160,11 +164,11 @@ def main(): on_open=on_open, ) - print(f"[WS][INFO] Connecting to {WS_URL} using mTLS...") - ws.run_forever(sslopt={"context": context}) + print(f"[WS][INFO] Connecting to {WS_URL}...") + ws.run_forever(sslopt=ssl_opt) except Exception as e: - print(f"[MAIN][ERROR] WebSocket setup failed: {e}") + print(f"[MAIN][ERROR] WebSocket connection failed: {e}") time.sleep(5) -- 2.54.0 From 3dd5f46f1ebcab3c7cec81931ab4bc67356bf8f8 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Wed, 25 Mar 2026 15:41:57 +0100 Subject: [PATCH 51/56] Automation 2.0: WIP added endpoints to backend, frontend connected and panel made, now debugging --- src/aare/common/automation_models.py | 194 +++--- src/aare/common/automation_queue_manager.py | 240 +++++++ src/aare/common/automation_workflow.py | 351 ++++++++++ src/aare/daq/auth.py | 8 +- src/aare/daq/automation_api_router.py | 694 ++++++++++++++++++++ src/aare/daq/automation_runner.py | 427 ++++++++++++ src/aare/daq/server.py | 31 + src/aare/gui/main_window.py | 64 ++ src/aare/gui/panels/automation_panel.py | 640 ++++++++++++++++++ src/aare/gui/threads/daq_worker.py | 168 ++++- src/aare/gui/threads/workflow_sse_client.py | 158 +++++ 11 files changed, 2887 insertions(+), 88 deletions(-) create mode 100644 src/aare/common/automation_queue_manager.py create mode 100644 src/aare/common/automation_workflow.py create mode 100644 src/aare/daq/automation_api_router.py create mode 100644 src/aare/daq/automation_runner.py create mode 100644 src/aare/gui/panels/automation_panel.py create mode 100644 src/aare/gui/threads/workflow_sse_client.py diff --git a/src/aare/common/automation_models.py b/src/aare/common/automation_models.py index 67be794b..7431e794 100644 --- a/src/aare/common/automation_models.py +++ b/src/aare/common/automation_models.py @@ -3,7 +3,10 @@ from __future__ import annotations from dataclasses import dataclass, field from enum import Enum from typing import Any +import time +import uuid +from pydantic import BaseModel, Field class WorkflowMode(str, Enum): FLEXIBLE_MANUAL = "flexible_manual" @@ -64,93 +67,116 @@ class WorkflowContext: metadata: dict[str, Any] = field(default_factory=dict) -STATE_REGISTRY: dict[WorkflowStateKind, StateDefinition] = { - WorkflowStateKind.MOUNT: StateDefinition( - kind=WorkflowStateKind.MOUNT, - description="Mount the sample", - transitions=( - TransitionRule( - to_state=WorkflowStateKind.LOOP_CENTRE, - allowed_modes=frozenset({ - WorkflowMode.FLEXIBLE_MANUAL, - WorkflowMode.GUIDED_MANUAL, - WorkflowMode.AUTOMATION, - }), - ), - ), - ), - WorkflowStateKind.LOOP_CENTRE: StateDefinition( - kind=WorkflowStateKind.LOOP_CENTRE, - description="Centre the loop", - transitions=( - TransitionRule( - to_state=WorkflowStateKind.RASTER, - allowed_modes=frozenset({ - WorkflowMode.FLEXIBLE_MANUAL, - WorkflowMode.GUIDED_MANUAL, - WorkflowMode.AUTOMATION, - }), - ), - TransitionRule( - to_state=WorkflowStateKind.DATA_COLLECTION, - allowed_modes=frozenset({ - WorkflowMode.FLEXIBLE_MANUAL, - WorkflowMode.GUIDED_MANUAL, - }), - optional=True, - ), - ), - ), - WorkflowStateKind.RASTER: StateDefinition( - kind=WorkflowStateKind.RASTER, - description="Run raster scan", - transitions=( - TransitionRule( - to_state=WorkflowStateKind.DATA_COLLECTION, - allowed_modes=frozenset({ - WorkflowMode.FLEXIBLE_MANUAL, - WorkflowMode.GUIDED_MANUAL, - WorkflowMode.AUTOMATION, - }), - ), - ), - ), - WorkflowStateKind.DATA_COLLECTION: StateDefinition( - kind=WorkflowStateKind.DATA_COLLECTION, - description="Collect diffraction data", - transitions=(), - ), -} +class QueueItemStatus(str, Enum): + PENDING = "pending" + RUNNING = "running" + COMPLETED = "completed" + FAILED = "failed" + SKIPPED = "skipped" + ABORTED = "aborted" -def get_state_definition(kind: WorkflowStateKind) -> StateDefinition: - try: - return STATE_REGISTRY[kind] - except KeyError as exc: - raise KeyError(f"Unknown workflow state: {kind}") from exc +class WorkflowStepRecord(BaseModel): + kind: str + status: str = "pending" + message: str = "" + started_at: float | None = None + completed_at: float | None = None + error_detail: str | None = None -def get_allowed_next_states( - kind: WorkflowStateKind, - mode: WorkflowMode | None = None, -) -> list[WorkflowStateKind]: - definition = get_state_definition(kind) - out: list[WorkflowStateKind] = [] - - for transition in definition.transitions: - if mode is None: - out.append(transition.to_state) - continue - - if not transition.allowed_modes or mode in transition.allowed_modes: - out.append(transition.to_state) - - return out +class QueueItem(BaseModel): + item_id: str + beamline: str + sample_id: int | None = None + sample_name: str = "" + owner_pgroup: str = "" + created_by: str = "" + created_at: float = Field(default_factory=time.time) + priority: int = 100 + order_index: int = 0 + status: QueueItemStatus = QueueItemStatus.PENDING + steps: list[WorkflowStepRecord] = Field(default_factory=list) + current_step_index: int = 0 + recipe: dict[str, Any] = Field(default_factory=dict) + metadata: dict[str, Any] = Field(default_factory=dict) -def can_transition( - from_state: WorkflowStateKind, - to_state: WorkflowStateKind, - mode: WorkflowMode | None = None, -) -> bool: - return to_state in get_allowed_next_states(from_state, mode=mode) \ No newline at end of file +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 \ No newline at end of file diff --git a/src/aare/common/automation_queue_manager.py b/src/aare/common/automation_queue_manager.py new file mode 100644 index 00000000..54e4d2e2 --- /dev/null +++ b/src/aare/common/automation_queue_manager.py @@ -0,0 +1,240 @@ +from __future__ import annotations + +import time +from typing import Any + +import redis + +from aare.common.automation_models import ( + QueueItem, + WorkflowEvent, + ControlState, + RuntimeState, + QueueItemStatus, + WorkflowStepRecord, + WorkflowStateKind, +) + + +def build_default_steps() -> list[WorkflowStepRecord]: + return [ + WorkflowStepRecord(kind=WorkflowStateKind.MOUNT.value), + WorkflowStepRecord(kind=WorkflowStateKind.LOOP_CENTRE.value), + WorkflowStepRecord(kind=WorkflowStateKind.RASTER.value), + WorkflowStepRecord(kind=WorkflowStateKind.DATA_COLLECTION.value), + ] + + +class WorkflowRedisManager: + def __init__(self, client: redis.Redis, beamline: str): + self._client = client + self._bl = beamline.lower() + + def _key(self, suffix: str) -> str: + return f"{self._bl}:workflow:{suffix}" + + def _item_key(self, item_id: str) -> str: + return self._key(f"item:{item_id}") + + def _queue_key(self) -> str: + return self._key("queue") + + def _runtime_key(self) -> str: + return self._key("runtime") + + def _control_key(self) -> str: + return self._key("control") + + def _events_key(self) -> str: + return self._key("events") + + def _next_item_id(self) -> str: + n = int(self._client.incr(self._key("item_seq"))) + return f"wf_{int(time.time())}_{n:06d}" + + # ───────────────────────────────────────────── + # Queue item operations + # ───────────────────────────────────────────── + + def create_item(self, item: QueueItem) -> QueueItem: + if not item.item_id: + item.item_id = self._next_item_id() + + pipe = self._client.pipeline(transaction=True) + pipe.set(self._item_key(item.item_id), item.model_dump_json()) + pipe.zadd(self._queue_key(), {item.item_id: float(item.order_index)}) + pipe.execute() + + self.append_event(WorkflowEvent( + beamline=self._bl, + item_id=item.item_id, + event_type="item_created", + message="Queue item created", + payload={"status": item.status.value}, + )) + + return item + + def get_item(self, item_id: str) -> QueueItem | None: + raw = self._client.get(self._item_key(item_id)) + if raw is None: + return None + return QueueItem.model_validate_json(raw) + + def update_item(self, item_id: str, patch: dict[str, Any]) -> QueueItem: + item = self.get_item(item_id) + if item is None: + raise KeyError(f"Queue item not found: {item_id}") + + updated = item.model_copy(update=patch) + self._client.set(self._item_key(item_id), updated.model_dump_json()) + return updated + + def delete_item(self, item_id: str) -> None: + pipe = self._client.pipeline(transaction=True) + pipe.delete(self._item_key(item_id)) + pipe.zrem(self._queue_key(), item_id) + pipe.execute() + + def list_queue_order(self) -> list[str]: + return [str(x) for x in self._client.zrange(self._queue_key(), 0, -1)] + + def list_items(self, *, include_finished: bool = True) -> list[QueueItem]: + out: list[QueueItem] = [] + for item_id in self.list_queue_order(): + item = self.get_item(item_id) + if item is None: + continue + if not include_finished and item.status in { + QueueItemStatus.COMPLETED, + QueueItemStatus.FAILED, + QueueItemStatus.SKIPPED, + QueueItemStatus.ABORTED, + }: + continue + out.append(item) + return out + + def get_next_pending_item(self) -> QueueItem | None: + for item_id in self.list_queue_order(): + item = self.get_item(item_id) + if item is not None and item.status == QueueItemStatus.PENDING: + return item + return None + + def move_item(self, item_id: str, new_order_index: int) -> None: + if self.get_item(item_id) is None: + raise KeyError(f"Queue item not found: {item_id}") + self._client.zadd(self._queue_key(), {item_id: float(new_order_index)}) + self.update_item(item_id, {"order_index": new_order_index}) + + # ───────────────────────────────────────────── + # Runtime state + # ───────────────────────────────────────────── + + def get_runtime(self) -> RuntimeState: + raw = self._client.get(self._runtime_key()) + if raw is None: + return RuntimeState() + return RuntimeState.model_validate_json(raw) + + def set_runtime(self, runtime: RuntimeState) -> RuntimeState: + runtime.last_update = time.time() + self._client.set(self._runtime_key(), runtime.model_dump_json()) + return runtime + + def patch_runtime(self, patch: dict[str, Any]) -> RuntimeState: + runtime = self.get_runtime() + updated = runtime.model_copy(update=patch) + return self.set_runtime(updated) + + # ───────────────────────────────────────────── + # Control state + # ───────────────────────────────────────────── + + def get_control(self) -> ControlState: + raw = self._client.get(self._control_key()) + if raw is None: + return ControlState() + return ControlState.model_validate_json(raw) + + def request_control(self, patch: dict[str, Any], *, requested_by: str) -> ControlState: + control = self.get_control() + updated = control.model_copy(update={ + **patch, + "requested_by": requested_by, + "requested_at": time.time(), + }) + self._client.set(self._control_key(), updated.model_dump_json()) + + self.append_event(WorkflowEvent( + beamline=self._bl, + item_id=self.get_runtime().current_item_id, + event_type="control_requested", + actor=requested_by, + payload=patch, + )) + + return updated + + def clear_control(self) -> ControlState: + control = ControlState() + self._client.set(self._control_key(), control.model_dump_json()) + return control + + # ───────────────────────────────────────────── + # Event log + # ───────────────────────────────────────────── + + def append_event(self, event: WorkflowEvent) -> None: + payload = event.model_dump_json() + self._client.xadd(self._events_key(), {"json": payload}, maxlen=5000, approximate=True) + + def read_events(self, *, limit: int = 200) -> list[WorkflowEvent]: + rows = self._client.xrevrange(self._events_key(), count=limit) + out: list[WorkflowEvent] = [] + for _, fields in reversed(rows): + raw = fields.get("json") + if raw: + out.append(WorkflowEvent.model_validate_json(raw)) + return out + + # ───────────────────────────────────────────── + # Step status updates + # ───────────────────────────────────────────── + + def update_step( + self, + item_id: str, + step_index: int, + *, + status: str | None = None, + message: str | None = None, + error_detail: str | None = None, + ) -> QueueItem: + item = self.get_item(item_id) + if item is None: + raise KeyError(f"Queue item not found: {item_id}") + + if not (0 <= step_index < len(item.steps)): + raise IndexError(f"Step index out of range: {step_index}") + + step = item.steps[step_index] + + if status is not None: + step.status = status + if status == "running" and step.started_at is None: + step.started_at = time.time() + elif status in ("completed", "failed", "skipped"): + step.completed_at = time.time() + + if message is not None: + step.message = message + + if error_detail is not None: + step.error_detail = error_detail + + item.steps[step_index] = step + self._client.set(self._item_key(item_id), item.model_dump_json()) + + return item \ No newline at end of file diff --git a/src/aare/common/automation_workflow.py b/src/aare/common/automation_workflow.py new file mode 100644 index 00000000..98f26ac1 --- /dev/null +++ b/src/aare/common/automation_workflow.py @@ -0,0 +1,351 @@ +from __future__ import annotations + +from abc import ABC, abstractmethod +from aare.common.automation_models import ( + WorkflowStateKind, + WorkflowMode, + StepStatus, + StateResult, + WorkflowContext, + TransitionRule, + StateDefinition, + ) + + +STATE_REGISTRY: dict[WorkflowStateKind, StateDefinition] = { + WorkflowStateKind.MOUNT: StateDefinition( + kind=WorkflowStateKind.MOUNT, + description="Mount the sample", + transitions=( + TransitionRule( + to_state=WorkflowStateKind.LOOP_CENTRE, + allowed_modes=frozenset({ + WorkflowMode.FLEXIBLE_MANUAL, + WorkflowMode.GUIDED_MANUAL, + WorkflowMode.AUTOMATION, + }), + ), + ), + ), + WorkflowStateKind.LOOP_CENTRE: StateDefinition( + kind=WorkflowStateKind.LOOP_CENTRE, + description="Centre the loop", + transitions=( + TransitionRule( + to_state=WorkflowStateKind.RASTER, + allowed_modes=frozenset({ + WorkflowMode.FLEXIBLE_MANUAL, + WorkflowMode.GUIDED_MANUAL, + WorkflowMode.AUTOMATION, + }), + ), + TransitionRule( + to_state=WorkflowStateKind.DATA_COLLECTION, + allowed_modes=frozenset({ + WorkflowMode.FLEXIBLE_MANUAL, + WorkflowMode.GUIDED_MANUAL, + }), + optional=True, + ), + ), + ), + WorkflowStateKind.RASTER: StateDefinition( + kind=WorkflowStateKind.RASTER, + description="Run raster scan", + transitions=( + TransitionRule( + to_state=WorkflowStateKind.DATA_COLLECTION, + allowed_modes=frozenset({ + WorkflowMode.FLEXIBLE_MANUAL, + WorkflowMode.GUIDED_MANUAL, + WorkflowMode.AUTOMATION, + }), + ), + ), + ), + WorkflowStateKind.DATA_COLLECTION: StateDefinition( + kind=WorkflowStateKind.DATA_COLLECTION, + description="Collect diffraction data", + transitions=(), + ), +} + + +def get_state_definition(kind: WorkflowStateKind) -> StateDefinition: + try: + return STATE_REGISTRY[kind] + except KeyError as exc: + raise KeyError(f"Unknown workflow state: {kind}") from exc + + +def get_allowed_next_states( + kind: WorkflowStateKind, + mode: WorkflowMode | None = None, +) -> list[WorkflowStateKind]: + definition = get_state_definition(kind) + out: list[WorkflowStateKind] = [] + + for transition in definition.transitions: + if mode is None: + out.append(transition.to_state) + continue + + if not transition.allowed_modes or mode in transition.allowed_modes: + out.append(transition.to_state) + + return out + + +def can_transition( + from_state: WorkflowStateKind, + to_state: WorkflowStateKind, + mode: WorkflowMode | None = None, +) -> bool: + return to_state in get_allowed_next_states(from_state, mode=mode) + + +class StateHandler(ABC): + state_kind: WorkflowStateKind + + def __init__(self, registry: dict[WorkflowStateKind, StateDefinition] | None = None): + self._registry = registry or STATE_REGISTRY + + def definition(self) -> StateDefinition: + return get_state_definition(self.state_kind) + + def can_run(self, context: WorkflowContext) -> bool: + return context.current_state in (None, self.state_kind) + + @abstractmethod + def validate(self, context: WorkflowContext) -> None: + pass + + @abstractmethod + def execute(self, context: WorkflowContext) -> StateResult: + pass + + +class MountHandler(StateHandler): + state_kind = WorkflowStateKind.MOUNT + + def validate(self, context: WorkflowContext) -> None: + if context.abort_requested: + raise RuntimeError("Abort requested; cannot mount.") + + def execute(self, context: WorkflowContext) -> StateResult: + self.validate(context) + context.current_state = WorkflowStateKind.MOUNT + context.last_message = "Sample mounted" + return StateResult( + state=self.state_kind, + status=StepStatus.SUCCESS, + message="Sample mounted successfully.", + payload={"mounted": True}, + ) + + +class LoopCentreHandler(StateHandler): + state_kind = WorkflowStateKind.LOOP_CENTRE + + def validate(self, context: WorkflowContext) -> None: + if context.abort_requested: + raise RuntimeError("Abort requested; cannot loop-centre.") + + def execute(self, context: WorkflowContext) -> StateResult: + self.validate(context) + context.current_state = WorkflowStateKind.LOOP_CENTRE + context.last_message = "Loop centred" + return StateResult( + state=self.state_kind, + status=StepStatus.SUCCESS, + message="Loop centring completed.", + payload={"centred": True}, + ) + + +class RasterHandler(StateHandler): + state_kind = WorkflowStateKind.RASTER + + def validate(self, context: WorkflowContext) -> None: + if context.abort_requested: + raise RuntimeError("Abort requested; cannot raster.") + + def execute(self, context: WorkflowContext) -> StateResult: + self.validate(context) + context.current_state = WorkflowStateKind.RASTER + context.last_message = "Raster completed" + return StateResult( + state=self.state_kind, + status=StepStatus.SUCCESS, + message="Raster scan completed.", + payload={"best_spot_found": True}, + ) + + +class DataCollectionHandler(StateHandler): + state_kind = WorkflowStateKind.DATA_COLLECTION + + def validate(self, context: WorkflowContext) -> None: + if context.abort_requested: + raise RuntimeError("Abort requested; cannot collect data.") + + def execute(self, context: WorkflowContext) -> StateResult: + self.validate(context) + context.current_state = WorkflowStateKind.DATA_COLLECTION + context.last_message = "Data collected" + return StateResult( + state=self.state_kind, + status=StepStatus.SUCCESS, + message="Data collection completed.", + payload={"frames_collected": 1}, + ) + + +class WorkflowRunner: + """Simple in-memory runner (no persistence).""" + + def __init__( + self, + registry: dict[WorkflowStateKind, StateDefinition] | None = None, + handlers: dict[WorkflowStateKind, StateHandler] | None = None, + ): + self._registry = registry or STATE_REGISTRY + self._handlers = handlers or HANDLER_REGISTRY + + def get_handler(self, state: WorkflowStateKind) -> StateHandler: + try: + return self._handlers[state] + except KeyError as exc: + raise KeyError(f"No handler registered for state: {state}") from exc + + def can_move_to( + self, + current: WorkflowStateKind, + next_state: WorkflowStateKind, + mode: WorkflowMode, + ) -> bool: + return can_transition(current, next_state, mode) + + def run_state( + self, + context: WorkflowContext, + state: WorkflowStateKind, + ) -> StateResult: + if context.current_state is not None: + if not self.can_move_to(context.current_state, state, context.mode): + raise RuntimeError( + f"Transition not allowed: {context.current_state} -> {state}" + ) + + handler = self.get_handler(state) + result = handler.execute(context) + + context.current_state = state + context.current_step_index += 1 + context.last_message = result.message + + return result + + +class SimulatedMountHandler(StateHandler): + """Simulated mount handler for testing - doesn't actually mount.""" + state_kind = WorkflowStateKind.MOUNT + + def validate(self, context: WorkflowContext) -> None: + if context.abort_requested: + raise RuntimeError("Abort requested; cannot mount.") + + def execute(self, context: WorkflowContext) -> StateResult: + self.validate(context) + import time + time.sleep(0.5) # Simulate some work + context.current_state = WorkflowStateKind.MOUNT + context.last_message = "[SIMULATION] Sample would be mounted" + return StateResult( + state=self.state_kind, + status=StepStatus.SUCCESS, + message=f"[SIMULATION] Would mount sample_id={context.sample_id}", + payload={"mounted": True, "simulated": True}, + ) + + +class SimulatedLoopCentreHandler(StateHandler): + """Simulated loop centre handler for testing.""" + state_kind = WorkflowStateKind.LOOP_CENTRE + + def validate(self, context: WorkflowContext) -> None: + if context.abort_requested: + raise RuntimeError("Abort requested; cannot loop-centre.") + + def execute(self, context: WorkflowContext) -> StateResult: + self.validate(context) + import time + time.sleep(0.3) + context.current_state = WorkflowStateKind.LOOP_CENTRE + context.last_message = "[SIMULATION] Loop would be centred" + return StateResult( + state=self.state_kind, + status=StepStatus.SUCCESS, + message="[SIMULATION] Would run loop centering algorithm", + payload={"centred": True, "simulated": True}, + ) + + +class SimulatedRasterHandler(StateHandler): + """Simulated raster handler for testing.""" + state_kind = WorkflowStateKind.RASTER + + def validate(self, context: WorkflowContext) -> None: + if context.abort_requested: + raise RuntimeError("Abort requested; cannot raster.") + + def execute(self, context: WorkflowContext) -> StateResult: + self.validate(context) + import time + time.sleep(0.4) + context.current_state = WorkflowStateKind.RASTER + context.last_message = "[SIMULATION] Raster scan would be performed" + return StateResult( + state=self.state_kind, + status=StepStatus.SUCCESS, + message="[SIMULATION] Would run raster scan, find best diffraction spot", + payload={"best_spot_found": True, "simulated": True}, + ) + + +class SimulatedDataCollectionHandler(StateHandler): + """Simulated data collection handler for testing.""" + state_kind = WorkflowStateKind.DATA_COLLECTION + + def validate(self, context: WorkflowContext) -> None: + if context.abort_requested: + raise RuntimeError("Abort requested; cannot collect data.") + + def execute(self, context: WorkflowContext) -> StateResult: + self.validate(context) + import time + time.sleep(0.5) + context.current_state = WorkflowStateKind.DATA_COLLECTION + context.last_message = "[SIMULATION] Data collection would be performed" + return StateResult( + state=self.state_kind, + status=StepStatus.SUCCESS, + message="[SIMULATION] Would collect 1800 frames at 0.2° oscillation", + payload={"frames_collected": 1800, "simulated": True}, + ) + + +# Simulated handler registry for testing +SIMULATED_HANDLER_REGISTRY: dict[WorkflowStateKind, StateHandler] = { + WorkflowStateKind.MOUNT: SimulatedMountHandler(STATE_REGISTRY), + WorkflowStateKind.LOOP_CENTRE: SimulatedLoopCentreHandler(STATE_REGISTRY), + WorkflowStateKind.RASTER: SimulatedRasterHandler(STATE_REGISTRY), + WorkflowStateKind.DATA_COLLECTION: SimulatedDataCollectionHandler(STATE_REGISTRY), +} + +HANDLER_REGISTRY: dict[WorkflowStateKind, StateHandler] = { + WorkflowStateKind.MOUNT: MountHandler(STATE_REGISTRY), + WorkflowStateKind.LOOP_CENTRE: LoopCentreHandler(STATE_REGISTRY), + WorkflowStateKind.RASTER: RasterHandler(STATE_REGISTRY), + WorkflowStateKind.DATA_COLLECTION: DataCollectionHandler(STATE_REGISTRY), +} \ No newline at end of file diff --git a/src/aare/daq/auth.py b/src/aare/daq/auth.py index 3af7f247..d93edd06 100644 --- a/src/aare/daq/auth.py +++ b/src/aare/daq/auth.py @@ -5,8 +5,8 @@ from datetime import datetime, timedelta, UTC from typing import List import jwt -from fastapi import HTTPException, status -from fastapi.security import OAuth2PasswordRequestForm +from fastapi import Depends, HTTPException, status +from fastapi.security import OAuth2PasswordRequestForm, OAuth2PasswordBearer from pydantic import BaseModel from aare.daq.config import BeamlineConfig @@ -25,6 +25,8 @@ SESSION_EXPIRE_SECONDS = 60 * 10 STAFF_GROUP = "unx-MXgroup" SUPER_USERS = ["e10019", "e11206", "e18147"] +oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token") + class TokenData(BaseModel): sub: str # Username pgroups: List[str] @@ -54,7 +56,7 @@ def authenticate_user(cfg: BeamlineConfig, form_data: OAuth2PasswordRequestForm) return create_access_token(token) -def parse_token(token: str) -> TokenData: +def parse_token(token: str = Depends(oauth2_scheme)) -> TokenData: try: payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM]) token = TokenData(**payload) diff --git a/src/aare/daq/automation_api_router.py b/src/aare/daq/automation_api_router.py new file mode 100644 index 00000000..1236ffc0 --- /dev/null +++ b/src/aare/daq/automation_api_router.py @@ -0,0 +1,694 @@ +from __future__ import annotations + +import asyncio +import json +from typing import AsyncGenerator, TYPE_CHECKING + +from fastapi import APIRouter, Depends, HTTPException, status +from pydantic import BaseModel, Field +from starlette.responses import StreamingResponse + +from aare.common.automation_models import ( + QueueItem, + QueueItemStatus, + WorkflowContext, + WorkflowEvent, + WorkflowMode, + WorkflowStepRecord, + RuntimeState, + ControlState, + WorkflowStateKind, EventListResponse, StepActionResponse, ControlActionResponse, RuntimeResponse, + CreateQueueItemRequest, QueueListResponse, MoveItemRequest, AutomationStatusResponse, +) +from aare.common.automation_queue_manager import ( + WorkflowRedisManager, + build_default_steps, +) +from aare.common.automation_workflow import STATE_REGISTRY, HANDLER_REGISTRY +from aare.daq.automation_runner import PersistentWorkflowRunner +from aare.daq.auth import parse_token, check_jwt_rw, check_jwt_ro, oauth2_scheme +from aare.common.models import TokenData +from aare.daq.config import BeamlineConfig + +from aare.daq.automation_runner import AutomationLoop + +router = APIRouter(prefix="/workflow", tags=["workflow"]) + +# ───────────────────────────────────────────── +# Dependency: get managers +# ───────────────────────────────────────────── + +# These will be set up when the router is included +_redis_manager: WorkflowRedisManager | None = None +_runner: PersistentWorkflowRunner | None = None +_cfg: BeamlineConfig | None = None +_automation_loop: AutomationLoop | None = None + + +def set_workflow_dependencies( + redis_manager: WorkflowRedisManager, + runner: PersistentWorkflowRunner, + cfg: BeamlineConfig, +) -> None: + global _redis_manager, _runner, _cfg, _automation_loop + _redis_manager = redis_manager + _runner = runner + _cfg = cfg + _automation_loop = AutomationLoop(runner, redis_manager) + + +def get_automation_loop() -> AutomationLoop: + if _automation_loop is None: + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail="Automation loop not initialized", + ) + return _automation_loop + + +def get_cfg() -> BeamlineConfig: + if _cfg is None: + raise HTTPException( + status_code=503, + detail="Workflow system not initialized", + ) + return _cfg + + +def get_redis_manager() -> WorkflowRedisManager: + if _redis_manager is None: + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail="Workflow system not initialized", + ) + return _redis_manager + + +def get_runner() -> PersistentWorkflowRunner: + if _runner is None: + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail="Workflow runner not initialized", + ) + return _runner + + +# ───────────────────────────────────────────── +# Queue management endpoints +# ───────────────────────────────────────────── + + +@router.get("/queue", response_model=QueueListResponse) +async def list_queue( + include_finished: bool = False, + token: str = Depends(oauth2_scheme), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), +): + """List all items in the workflow queue.""" + data = parse_token(token) + check_jwt_ro(get_cfg(), data) + items = redis_mgr.list_items(include_finished=include_finished) + + # Filter by pgroup if not staff + if not data.staff: + items = [i for i in items if i.owner_pgroup in data.pgroups] + + return QueueListResponse(items=items, total=len(items)) + + +@router.post("/queue", response_model=QueueItem) +async def create_queue_item( + request: CreateQueueItemRequest, + token: str = Depends(oauth2_scheme), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), +): + """Add a new item to the workflow queue.""" + data = parse_token(token) + check_jwt_rw(get_cfg(), data) + + # Build steps + if request.steps: + steps = [ + WorkflowStepRecord(kind=s) + for s in request.steps + if s in [sk.value for sk in WorkflowStateKind] + ] + else: + steps = build_default_steps() + + item = QueueItem( + item_id="", + beamline=redis_mgr._bl, + sample_id=request.sample_id, + sample_name=request.sample_name, + owner_pgroup=data.pgroups[0] if data.pgroups else "", + created_by=data.sub, + priority=request.priority, + order_index=int(asyncio.get_event_loop().time() * 1000), + steps=steps, + recipe=request.recipe, + ) + + item = redis_mgr.create_item(item) + return item + + +@router.get("/queue/{item_id}", response_model=QueueItem) +async def get_queue_item( + item_id: str, + token: str = Depends(oauth2_scheme), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), +): + """Get a single queue item by ID.""" + data = parse_token(token) + check_jwt_ro(get_cfg(), data) + + item = redis_mgr.get_item(item_id) + if item is None: + raise HTTPException(status_code=404, detail="Item not found") + + # Check access + if not data.staff and item.owner_pgroup not in data.pgroups: + raise HTTPException(status_code=403, detail="Access denied") + + return item + + +@router.delete("/queue/{item_id}") +async def delete_queue_item( + item_id: str, + token: str = Depends(oauth2_scheme), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), +): + """Remove an item from the queue.""" + data = parse_token(token) + check_jwt_rw(get_cfg(), data) + + item = redis_mgr.get_item(item_id) + if item is None: + raise HTTPException(status_code=404, detail="Item not found") + + # Check access + if not data.staff and item.owner_pgroup not in data.pgroups: + raise HTTPException(status_code=403, detail="Access denied") + + # Don't allow deleting running items + if item.status == QueueItemStatus.RUNNING: + raise HTTPException(status_code=409, detail="Cannot delete running item") + + redis_mgr.delete_item(item_id) + return {"ok": True, "message": f"Item {item_id} deleted"} + + +@router.post("/queue/{item_id}/move") +async def move_queue_item( + item_id: str, + request: MoveItemRequest, + token: str = Depends(oauth2_scheme), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), +): + """Reorder an item in the queue.""" + data = parse_token(token) + check_jwt_rw(get_cfg(), data) + + item = redis_mgr.get_item(item_id) + if item is None: + raise HTTPException(status_code=404, detail="Item not found") + + if not data.staff and item.owner_pgroup not in data.pgroups: + raise HTTPException(status_code=403, detail="Access denied") + + redis_mgr.move_item(item_id, request.new_order_index) + return {"ok": True, "message": f"Item {item_id} moved"} + + +# ───────────────────────────────────────────── +# Runtime and control endpoints +# ───────────────────────────────────────────── + + +@router.get("/runtime", response_model=RuntimeResponse) +async def get_runtime( + token: str = Depends(oauth2_scheme), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), +): + """Get current runtime and control state.""" + data = parse_token(token) + check_jwt_ro(get_cfg(), data) + + return RuntimeResponse( + runtime=redis_mgr.get_runtime(), + control=redis_mgr.get_control(), + ) + + +@router.get("/control", response_model=ControlState) +async def get_control( + token: str = Depends(oauth2_scheme), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), +): + """Get current control state.""" + data = parse_token(token) + check_jwt_ro(get_cfg(), data) + return redis_mgr.get_control() + + +@router.post("/control/pause", response_model=ControlActionResponse) +async def request_pause( + token: str = Depends(oauth2_scheme), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), +): + """Request workflow pause after current step.""" + data = parse_token(token) + check_jwt_rw(get_cfg(), data) + + control = redis_mgr.request_control( + {"pause_requested": True, "resume_requested": False}, + requested_by=data.sub, + ) + + return ControlActionResponse( + ok=True, + control=control, + message="Pause requested", + ) + + +@router.post("/control/resume", response_model=ControlActionResponse) +async def request_resume( + token: str = Depends(oauth2_scheme), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), +): + """Resume paused workflow.""" + data = parse_token(token) + check_jwt_rw(get_cfg(), data) + + control = redis_mgr.request_control( + {"pause_requested": False, "resume_requested": True}, + requested_by=data.sub, + ) + + return ControlActionResponse( + ok=True, + control=control, + message="Resume requested", + ) + + +@router.post("/control/abort", response_model=ControlActionResponse) +async def request_abort( + token: str = Depends(oauth2_scheme), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), +): + """Abort current workflow execution.""" + data = parse_token(token) + check_jwt_rw(get_cfg(), data) + + control = redis_mgr.request_control( + {"abort_requested": True}, + requested_by=data.sub, + ) + + return ControlActionResponse( + ok=True, + control=control, + message="Abort requested", + ) + + +@router.post("/control/skip", response_model=ControlActionResponse) +async def request_skip( + token: str = Depends(oauth2_scheme), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), +): + """Skip current step.""" + data = parse_token(token) + check_jwt_rw(get_cfg(), data) + + control = redis_mgr.request_control( + {"skip_requested": True}, + requested_by=data.sub, + ) + + return ControlActionResponse( + ok=True, + control=control, + message="Skip requested", + ) + + +@router.post("/control/clear", response_model=ControlActionResponse) +async def clear_control( + token: str = Depends(oauth2_scheme), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), +): + """Clear all control requests.""" + data = parse_token(token) + check_jwt_rw(get_cfg(), data) + + control = redis_mgr.clear_control() + + return ControlActionResponse( + ok=True, + control=control, + message="Control state cleared", + ) + + +# ───────────────────────────────────────────── +# Execution endpoints +# ───────────────────────────────────────────── + + +@router.post("/start/{item_id}", response_model=StepActionResponse) +async def start_item( + item_id: str, + mode: WorkflowMode = WorkflowMode.GUIDED_MANUAL, + token: str = Depends(oauth2_scheme), + runner: PersistentWorkflowRunner = Depends(get_runner), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), +): + """Start processing a queue item.""" + data = parse_token(token) + check_jwt_rw(get_cfg(), data) + + item = redis_mgr.get_item(item_id) + if item is None: + raise HTTPException(status_code=404, detail="Item not found") + + if not data.staff and item.owner_pgroup not in data.pgroups: + raise HTTPException(status_code=403, detail="Access denied") + + if item.status == QueueItemStatus.RUNNING: + raise HTTPException(status_code=409, detail="Item already running") + + # Check if another item is running + runtime = redis_mgr.get_runtime() + if runtime.running: + raise HTTPException( + status_code=409, + detail=f"Another item is running: {runtime.current_item_id}", + ) + + # Clear any stale control requests + redis_mgr.clear_control() + + # Start the item + item = runner.start_item(item_id) + + return StepActionResponse( + ok=True, + item=item, + step=item.steps[0].kind if item.steps else None, + status="started", + message=f"Started processing {item.sample_name}", + ) + + +@router.post("/next", response_model=StepActionResponse) +async def run_next_step( + token: str = Depends(oauth2_scheme), + runner: PersistentWorkflowRunner = Depends(get_runner), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), +): + """ + Run the next step in guided mode. + + This is the main endpoint for guided manual operation. + User clicks "Next" and this executes one step. + """ + data = parse_token(token) + check_jwt_rw(get_cfg(), data) + + runtime = redis_mgr.get_runtime() + + if not runtime.running or not runtime.current_item_id: + raise HTTPException(status_code=409, detail="No item is currently running") + + item = redis_mgr.get_item(runtime.current_item_id) + if item is None: + raise HTTPException(status_code=404, detail="Running item not found") + + if not data.staff and item.owner_pgroup not in data.pgroups: + raise HTTPException(status_code=403, detail="Access denied") + + # Check if all steps completed + if item.current_step_index >= len(item.steps): + item = runner.complete_item(item.item_id, QueueItemStatus.COMPLETED) + return StepActionResponse( + ok=True, + item=item, + status="completed", + message="All steps completed", + ) + + # Check for abort + if runner.should_abort(): + item = runner.complete_item(item.item_id, QueueItemStatus.ABORTED) + redis_mgr.clear_control() + return StepActionResponse( + ok=True, + item=item, + status="aborted", + message="Workflow aborted by user", + ) + + # Check for skip + control = runner.check_control() + if control.skip_requested: + step_index = item.current_step_index + step_kind = item.steps[step_index].kind + + redis_mgr.update_step(item.item_id, step_index, status="skipped") + redis_mgr.update_item(item.item_id, {"current_step_index": step_index + 1}) + redis_mgr.request_control({"skip_requested": False}, requested_by=data.sub) + + redis_mgr.append_event(WorkflowEvent( + beamline=runner.beamline, + item_id=item.item_id, + step=step_kind, + event_type="step_skipped", + actor=data.sub, + message=f"Step {step_kind} skipped by user", + )) + + item = redis_mgr.get_item(item.item_id) + return StepActionResponse( + ok=True, + item=item, + step=step_kind, + status="skipped", + message=f"Skipped {step_kind}", + ) + + # Build context + context = WorkflowContext( + mode=WorkflowMode.GUIDED_MANUAL, + queue_id="default", + item_id=item.item_id, + sample_id=item.sample_id, + current_state=WorkflowStateKind(runtime.current_state) if runtime.current_state else None, + current_step_index=item.current_step_index, + ) + + # Run the step + try: + result = runner.run_current_step(context) + except Exception as e: + return StepActionResponse( + ok=False, + item=redis_mgr.get_item(item.item_id), + step=item.steps[item.current_step_index].kind, + status="failed", + message=str(e), + ) + + # Get updated item + item = redis_mgr.get_item(item.item_id) + + # Check if completed + if item.current_step_index >= len(item.steps): + item = runner.complete_item(item.item_id, QueueItemStatus.COMPLETED) + return StepActionResponse( + ok=True, + item=item, + step=result.state.value, + status="completed", + message="All steps completed", + ) + + return StepActionResponse( + ok=True, + item=item, + step=result.state.value, + status=result.status.value, + message=result.message, + ) + + +# ───────────────────────────────────────────── +# Events endpoints +# ───────────────────────────────────────────── + + +@router.get("/events", response_model=EventListResponse) +async def get_events( + limit: int = 100, + token: str = Depends(oauth2_scheme), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), +): + """Get recent workflow events.""" + data = parse_token(token) + check_jwt_ro(get_cfg(), data) + + events = redis_mgr.read_events(limit=limit) + return EventListResponse(events=events) + + +# ───────────────────────────────────────────── +# SSE stream +# ───────────────────────────────────────────── + + +async def workflow_event_stream( + redis_mgr: WorkflowRedisManager, +) -> AsyncGenerator[str, None]: + """ + Server-Sent Events stream for live workflow updates. + + Polls runtime state and streams changes. + """ + last_runtime_json = "" + last_control_json = "" + last_event_id = "0-0" + + try: + while True: + # Check runtime state + runtime = redis_mgr.get_runtime() + runtime_json = runtime.model_dump_json() + + if runtime_json != last_runtime_json: + last_runtime_json = runtime_json + yield f"event: runtime\ndata: {runtime_json}\n\n" + + # Check control state + control = redis_mgr.get_control() + control_json = control.model_dump_json() + + if control_json != last_control_json: + last_control_json = control_json + yield f"event: control\ndata: {control_json}\n\n" + + # Check for new events (using Redis streams) + try: + events_key = redis_mgr._events_key() + new_events = redis_mgr._client.xread( + {events_key: last_event_id}, + count=10, + block=0, + ) + + if new_events: + for _, messages in new_events: + for msg_id, fields in messages: + last_event_id = msg_id + raw = fields.get("json") + if raw: + yield f"event: workflow_event\ndata: {raw}\n\n" + except Exception: + pass # Redis stream read failed, continue polling + + await asyncio.sleep(0.2) + + except asyncio.CancelledError: + return + +# ───────────────────────────────────────────── +# Automation mode endpoints +# ───────────────────────────────────────────── + + +@router.get("/automation/status", response_model=AutomationStatusResponse) +async def get_automation_status( + token: str = Depends(oauth2_scheme), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), + loop: "AutomationLoop" = Depends(get_automation_loop), +): + """Get automation loop status.""" + data = parse_token(token) + check_jwt_ro(get_cfg(), data) + + return AutomationStatusResponse( + enabled=loop.is_enabled, + running=loop.is_running, + runtime=redis_mgr.get_runtime(), + control=redis_mgr.get_control(), + ) + + +@router.post("/automation/start", response_model=AutomationStatusResponse) +async def start_automation( + token: str = Depends(oauth2_scheme), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), + loop: "AutomationLoop" = Depends(get_automation_loop), +): + """Start automation mode - processes queue automatically.""" + data = parse_token(token) + check_jwt_rw(get_cfg(), data) + + loop.start() + + return AutomationStatusResponse( + enabled=loop.is_enabled, + running=loop.is_running, + runtime=redis_mgr.get_runtime(), + control=redis_mgr.get_control(), + ) + + +@router.post("/automation/stop", response_model=AutomationStatusResponse) +async def stop_automation( + token: str = Depends(oauth2_scheme), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), + loop: "AutomationLoop" = Depends(get_automation_loop), +): + """Stop automation mode - completes current step then stops.""" + data = parse_token(token) + check_jwt_rw(get_cfg(), data) + + loop.stop() + + return AutomationStatusResponse( + enabled=loop.is_enabled, + running=loop.is_running, + runtime=redis_mgr.get_runtime(), + control=redis_mgr.get_control(), + ) + +@router.get("/sse") +async def workflow_sse( + token: str = Depends(oauth2_scheme), + redis_mgr: WorkflowRedisManager = Depends(get_redis_manager), +): + """ + SSE endpoint for live workflow updates. + + Events: + - runtime: RuntimeState changes + - control: ControlState changes + - workflow_event: Individual workflow events + """ + data = parse_token(token) + check_jwt_ro(get_cfg(), data) + + return StreamingResponse( + workflow_event_stream(redis_mgr), + media_type="text/event-stream", + headers={ + "Cache-Control": "no-cache", + "Connection": "keep-alive", + "Access-Control-Allow-Origin": "*", + "Access-Control-Allow-Headers": "Cache-Control", + }, + ) \ No newline at end of file diff --git a/src/aare/daq/automation_runner.py b/src/aare/daq/automation_runner.py new file mode 100644 index 00000000..4626e601 --- /dev/null +++ b/src/aare/daq/automation_runner.py @@ -0,0 +1,427 @@ +import asyncio +from typing import Callable + +import redis + +from aare.common.automation_models import ( + WorkflowEvent, + QueueItemStatus, + QueueItem, + ControlState, + WorkflowStateKind, + WorkflowContext, + StateResult, + StateDefinition, + WorkflowMode, +) +from aare.common.automation_queue_manager import ( + WorkflowRedisManager, + build_default_steps +) + +from aare.common.automation_workflow import ( + can_transition, + StateHandler, + STATE_REGISTRY, + HANDLER_REGISTRY +) + + +class PersistentWorkflowRunner: + """ + Orchestrates workflow execution with Redis persistence. + + Combines: + - State handlers (from automation_workflow) + - Redis persistence (from automation_queue_manager) + """ + + def __init__( + self, + redis_manager: WorkflowRedisManager, + registry: dict[WorkflowStateKind, StateDefinition] | None = None, + handlers: dict[WorkflowStateKind, StateHandler] | None = None, + ): + self._redis = redis_manager + self._registry = registry or STATE_REGISTRY + self._handlers = handlers or HANDLER_REGISTRY + + @property + def beamline(self) -> str: + return self._redis._bl + + def get_handler(self, state: WorkflowStateKind) -> StateHandler: + try: + return self._handlers[state] + except KeyError as exc: + raise KeyError(f"No handler registered for state: {state}") from exc + + def start_item(self, item_id: str) -> QueueItem: + item = self._redis.get_item(item_id) + if item is None: + raise KeyError(f"Item not found: {item_id}") + + item = self._redis.update_item(item_id, { + "status": QueueItemStatus.RUNNING, + "current_step_index": 0, + }) + + self._redis.patch_runtime({ + "running": True, + "paused": False, + "current_item_id": item_id, + "current_step_index": 0, + "current_state": None, + "last_error": None, + }) + + self._redis.append_event(WorkflowEvent( + beamline=self.beamline, + item_id=item_id, + event_type="item_started", + message=f"Started processing {item.sample_name}", + )) + + return item + + def run_current_step(self, context: WorkflowContext) -> StateResult: + item = self._redis.get_item(context.item_id) + if item is None: + raise KeyError(f"Item not found: {context.item_id}") + + step_index = item.current_step_index + if step_index >= len(item.steps): + raise RuntimeError("No more steps to run") + + step_record = item.steps[step_index] + state_kind = WorkflowStateKind(step_record.kind) + + # Check transition is allowed + if context.current_state is not None: + if not can_transition(context.current_state, state_kind, context.mode): + raise RuntimeError( + f"Transition not allowed: {context.current_state} -> {state_kind}" + ) + + # Mark step as running + self._redis.update_step(context.item_id, step_index, status="running") + self._redis.patch_runtime({ + "current_state": state_kind.value, + "current_step_index": step_index, + }) + + self._redis.append_event(WorkflowEvent( + beamline=self.beamline, + item_id=context.item_id, + step=state_kind.value, + event_type="step_started", + message=f"Starting {state_kind.value}", + )) + + # Execute handler + handler = self.get_handler(state_kind) + + try: + result = handler.execute(context) + except Exception as e: + self._redis.update_step( + context.item_id, + step_index, + status="failed", + error_detail=str(e), + ) + self._redis.patch_runtime({"last_error": str(e)}) + self._redis.append_event(WorkflowEvent( + beamline=self.beamline, + item_id=context.item_id, + step=state_kind.value, + event_type="step_failed", + message=str(e), + )) + raise + + # Mark step as completed + self._redis.update_step( + context.item_id, + step_index, + status=result.status.value, + message=result.message, + ) + + # Advance step index + self._redis.update_item(context.item_id, { + "current_step_index": step_index + 1, + }) + + context.current_state = state_kind + context.current_step_index = step_index + 1 + + self._redis.append_event(WorkflowEvent( + beamline=self.beamline, + item_id=context.item_id, + step=state_kind.value, + event_type="step_completed", + message=result.message, + payload=result.payload, + )) + + return result + + def complete_item(self, item_id: str, status: QueueItemStatus) -> QueueItem: + item = self._redis.update_item(item_id, {"status": status}) + + self._redis.patch_runtime({ + "running": False, + "current_item_id": None, + "current_state": None, + }) + + self._redis.append_event(WorkflowEvent( + beamline=self.beamline, + item_id=item_id, + event_type="item_completed", + message=f"Item finished with status {status.value}", + )) + + return item + + def check_control(self) -> ControlState: + return self._redis.get_control() + + def should_pause(self) -> bool: + return self.check_control().pause_requested + + def should_abort(self) -> bool: + return self.check_control().abort_requested + + def should_skip(self) -> bool: + return self.check_control().skip_requested + + +class AutomationLoop: + """ + Background task that drives fully automated workflow execution. + + Polls control state and processes the queue automatically. + """ + + def __init__( + self, + runner: PersistentWorkflowRunner, + redis_manager: WorkflowRedisManager, + poll_interval: float = 0.5, + ): + self._runner = runner + self._redis = redis_manager + self._poll_interval = poll_interval + self._task: asyncio.Task | None = None + self._enabled = False + self._on_step_complete: Callable[[str, str, StateResult], None] | None = None + + @property + def is_running(self) -> bool: + return self._task is not None and not self._task.done() + + @property + def is_enabled(self) -> bool: + return self._enabled + + def set_step_callback(self, callback: Callable[[str, str, StateResult], None]) -> None: + """Set callback for step completion: callback(item_id, step_name, result)""" + self._on_step_complete = callback + + def start(self) -> None: + """Start the automation loop.""" + if self._task is not None and not self._task.done(): + return # Already running + + self._enabled = True + self._task = asyncio.create_task(self._run_loop()) + + self._redis.append_event(WorkflowEvent( + beamline=self._runner.beamline, + event_type="automation_started", + message="Automation mode enabled", + )) + + def stop(self) -> None: + """Stop the automation loop gracefully.""" + self._enabled = False + + self._redis.append_event(WorkflowEvent( + beamline=self._runner.beamline, + event_type="automation_stopped", + message="Automation mode disabled", + )) + + async def _run_loop(self) -> None: + """Main automation loop.""" + while self._enabled: + try: + await self._tick() + except Exception as e: + self._redis.patch_runtime({"last_error": str(e)}) + self._redis.append_event(WorkflowEvent( + beamline=self._runner.beamline, + event_type="automation_error", + message=f"Automation error: {e}", + )) + # Brief pause on error before retrying + await asyncio.sleep(2.0) + + await asyncio.sleep(self._poll_interval) + + async def _tick(self) -> None: + """Single iteration of the automation loop.""" + control = self._runner.check_control() + runtime = self._redis.get_runtime() + + # Handle pause state + if control.pause_requested: + if runtime.running and not runtime.paused: + self._redis.patch_runtime({"paused": True}) + self._redis.append_event(WorkflowEvent( + beamline=self._runner.beamline, + item_id=runtime.current_item_id, + event_type="workflow_paused", + message="Workflow paused by user request", + )) + return + + # Handle resume + if control.resume_requested and runtime.paused: + self._redis.patch_runtime({"paused": False}) + self._redis.request_control( + {"resume_requested": False}, + requested_by="automation_loop", + ) + self._redis.append_event(WorkflowEvent( + beamline=self._runner.beamline, + item_id=runtime.current_item_id, + event_type="workflow_resumed", + message="Workflow resumed", + )) + + # Don't process if paused + if runtime.paused: + return + + # Handle abort + if control.abort_requested and runtime.current_item_id: + self._runner.complete_item(runtime.current_item_id, QueueItemStatus.ABORTED) + self._redis.clear_control() + return + + # If nothing running, try to start next item + if not runtime.running or not runtime.current_item_id: + next_item = self._redis.get_next_pending_item() + if next_item is None: + return # Queue empty, nothing to do + + self._runner.start_item(next_item.item_id) + runtime = self._redis.get_runtime() + + # Process current item + item = self._redis.get_item(runtime.current_item_id) + if item is None: + self._redis.patch_runtime({"running": False, "current_item_id": None}) + return + + # Check if item is complete + if item.current_step_index >= len(item.steps): + self._runner.complete_item(item.item_id, QueueItemStatus.COMPLETED) + + # Check for next sample request + if control.next_sample_requested: + self._redis.request_control( + {"next_sample_requested": False}, + requested_by="automation_loop", + ) + return + + # Handle skip request + if control.skip_requested: + step_kind = item.steps[item.current_step_index].kind + self._redis.update_step(item.item_id, item.current_step_index, status="skipped") + self._redis.update_item(item.item_id, {"current_step_index": item.current_step_index + 1}) + self._redis.request_control({"skip_requested": False}, requested_by="automation_loop") + + self._redis.append_event(WorkflowEvent( + beamline=self._runner.beamline, + item_id=item.item_id, + step=step_kind, + event_type="step_skipped", + message=f"Step {step_kind} skipped", + )) + return + + # Build context and run step + context = WorkflowContext( + mode=WorkflowMode.AUTOMATION, + queue_id="default", + item_id=item.item_id, + sample_id=item.sample_id, + current_state=WorkflowStateKind(runtime.current_state) if runtime.current_state else None, + current_step_index=item.current_step_index, + ) + + # Run the step (this is blocking in the async context) + result = await asyncio.to_thread(self._runner.run_current_step, context) + + # Notify callback if set + if self._on_step_complete: + step_name = item.steps[item.current_step_index - 1].kind # -1 because index advanced + self._on_step_complete(item.item_id, step_name, result) + + + +if __name__ == "__main__": + client = redis.Redis(host="localhost", port=6379, db=0, decode_responses=True) + redis_mgr = WorkflowRedisManager(client, beamline="x10sa") + + # Create a queue item + item = QueueItem( + item_id="", + beamline="x10sa", + sample_id=123, + sample_name="lysozyme_01", + owner_pgroup="p12345", + created_by="user@example.com", + steps=build_default_steps(), + ) + item = redis_mgr.create_item(item) + + # Create runner + runner = PersistentWorkflowRunner( + redis_manager=redis_mgr, + registry=STATE_REGISTRY, + handlers=HANDLER_REGISTRY, + ) + + # Start processing + runner.start_item(item.item_id) + + # Build context + context = WorkflowContext( + mode=WorkflowMode.GUIDED_MANUAL, + queue_id="default", + item_id=item.item_id, + sample_id=item.sample_id, + ) + + # Run each step (in guided mode, user triggers each one) + while context.current_step_index < len(item.steps): + if runner.should_pause(): + print("Paused by user") + break + + if runner.should_abort(): + runner.complete_item(item.item_id, QueueItemStatus.ABORTED) + break + + result = runner.run_current_step(context) + print(f"Step {result.state.value}: {result.status.value}") + + # Mark complete if all steps done + if context.current_step_index >= len(item.steps): + runner.complete_item(item.item_id, QueueItemStatus.COMPLETED) \ No newline at end of file diff --git a/src/aare/daq/server.py b/src/aare/daq/server.py index 2665f44d..cf9a02ec 100644 --- a/src/aare/daq/server.py +++ b/src/aare/daq/server.py @@ -34,6 +34,12 @@ from aare.common.exception_handler import ( SampleException, UserRightsException, ) + +from aare.common.automation_queue_manager import WorkflowRedisManager +from aare.common.automation_workflow import STATE_REGISTRY, HANDLER_REGISTRY, SIMULATED_HANDLER_REGISTRY +from aare.daq.automation_runner import PersistentWorkflowRunner +from aare.daq.automation_api_router import router as workflow_router, set_workflow_dependencies + logger = setup_logger("aareDAQ") app = FastAPI() register_exception_handlers(app) @@ -45,6 +51,31 @@ bl = mx_beamline() cfg = BeamlineConfig(bl) daq = AareDAQ(cfg, bl) +# ───────────────────────────────────────────── +# Initialize workflow system +# ───────────────────────────────────────────── + +USE_SIMULATED_WORKFLOW = os.getenv("WORKFLOW_SIMULATION", "0") == "1" + +workflow_redis_manager = WorkflowRedisManager( + client=cfg._BeamlineConfig__client, # Reuse existing Redis connection + beamline=bl.value, +) + +workflow_runner = PersistentWorkflowRunner( + redis_manager=workflow_redis_manager, + registry=STATE_REGISTRY, + handlers=SIMULATED_HANDLER_REGISTRY if USE_SIMULATED_WORKFLOW else HANDLER_REGISTRY, +) + +if USE_SIMULATED_WORKFLOW: + logger.warning("⚠️ Workflow system running in SIMULATION mode - no actual DAQ operations") + +set_workflow_dependencies(workflow_redis_manager, workflow_runner, cfg) + +# Include the workflow router +app.include_router(workflow_router) + try: daq.sync_current_sample_from_tell(force=True) except Exception as e: diff --git a/src/aare/gui/main_window.py b/src/aare/gui/main_window.py index 9d05a250..67b75a75 100644 --- a/src/aare/gui/main_window.py +++ b/src/aare/gui/main_window.py @@ -55,6 +55,9 @@ from aare.gui.widgets.status_bar import StatusBar from aare.gui.widgets.video_image import VideoGraphicsView from aare.gui.panels.fluorescence_panel import FluorescencePanel +from aare.gui.panels.automation_panel import WorkflowPanel +from aare.gui.threads.workflow_sse_client import WorkflowSSEClient + logger = setup_logger("aareGUI") class MainWindow(QMainWindow): @@ -288,6 +291,17 @@ class MainWindow(QMainWindow): self.addDockWidget(Qt.DockWidgetArea.RightDockWidgetArea, self.smargon_trace_dock) self.smargon_trace_dock.hide() + # === Workflow Panel === + self.workflow_panel = WorkflowPanel() + self.workflow_dock = QDockWidget("Workflow", self) + self.workflow_dock.setObjectName("workflow_dock") + self.workflow_dock.setWidget(self.workflow_panel) + self.workflow_dock.setAllowedAreas( + Qt.DockWidgetArea.RightDockWidgetArea | Qt.DockWidgetArea.LeftDockWidgetArea + ) + self.addDockWidget(Qt.DockWidgetArea.RightDockWidgetArea, self.workflow_dock) + self.workflow_dock.hide() # Hidden by default + root_layout.addWidget(top_widget) self.setCentralWidget(root_widget) @@ -497,6 +511,46 @@ class MainWindow(QMainWindow): register_tutorials(self, self.tutorial_manager) + # Workflow SSE client + if self.__base_url is not None: + self.workflow_sse = WorkflowSSEClient(self.__base_url, self.__token, self) + self.workflow_sse.runtime_changed.connect(self.workflow_panel.update_runtime) + self.workflow_sse.control_changed.connect(self.workflow_panel.update_control) + self.workflow_sse.workflow_event.connect( + lambda e: self.workflow_panel.on_workflow_event(e.event_type, e.message) + ) + self.workflow_sse.connect() + else: + self.workflow_sse = None + + # Connect workflow panel signals to DAQ worker + self.workflow_panel.request_queue_refresh.connect(self.daq.workflow_load_queue) + self.workflow_panel.request_add_sample.connect(self.daq.workflow_add_sample) + self.workflow_panel.request_add_samples.connect(self.daq.workflow_add_samples) + self.workflow_panel.request_delete_item.connect(self.daq.workflow_delete_item) + self.workflow_panel.request_move_item.connect(self.daq.workflow_move_item) + self.workflow_panel.request_clear_queue.connect(self.daq.workflow_clear_queue) + self.workflow_panel.request_start_item.connect(self.daq.workflow_start_item) + self.workflow_panel.request_next_step.connect(self.daq.workflow_next_step) + self.workflow_panel.request_pause.connect(self.daq.workflow_pause) + self.workflow_panel.request_resume.connect(self.daq.workflow_resume) + self.workflow_panel.request_abort.connect(self.daq.workflow_abort) + self.workflow_panel.request_skip.connect(self.daq.workflow_skip) + self.workflow_panel.request_start_automation.connect(self.daq.workflow_start_automation) + self.workflow_panel.request_stop_automation.connect(self.daq.workflow_stop_automation) + + # DAQ worker -> workflow panel + self.daq.workflow_queue_loaded.connect(self.workflow_panel.update_queue) + self.daq.workflow_item_updated.connect(self.workflow_panel.update_current_item) + + # Forward sample list to workflow panel for "Add All" feature + self.daq.spreadsheet.connect( + lambda slist: self.workflow_panel.update_available_samples(slist.s) + ) + + # Initial load + QTimer.singleShot(1000, self.daq.workflow_load_queue) + @Slot(QPixmap) def _on_samcam_prediction_pixmap(self, pix: QPixmap) -> None: self._last_pred_image_ts = time.monotonic() @@ -596,6 +650,13 @@ class MainWindow(QMainWindow): ) view_menu.addAction(show_smargon_trace_action) + show_workflow_action = QAction("Show Workflow Panel", self) + show_workflow_action.setCheckable(True) + show_workflow_action.setChecked(False) + show_workflow_action.triggered.connect(lambda checked: self.workflow_dock.setVisible(checked)) + self.workflow_dock.visibilityChanged.connect(show_workflow_action.setChecked) + view_menu.addAction(show_workflow_action) + show_log_action = QAction("Show Log", self) show_log_action.setCheckable(True) show_log_action.setChecked(False) @@ -770,6 +831,9 @@ class MainWindow(QMainWindow): except Exception as e: logger.warning(f"Failed to stop _samcam_source_timer: {e}") + if hasattr(self, "workflow_sse") and self.workflow_sse is not None: + self.workflow_sse.disconnect() + for attr_name in ( "camera_thread", "prediction_thread", diff --git a/src/aare/gui/panels/automation_panel.py b/src/aare/gui/panels/automation_panel.py new file mode 100644 index 00000000..e53ea753 --- /dev/null +++ b/src/aare/gui/panels/automation_panel.py @@ -0,0 +1,640 @@ +""" +Workflow automation panel for queue management and step control. + +Displays: +- Queue items with status (supports drag-drop from sample list) +- Current step progress +- Control buttons (Next/Pause/Resume/Abort/Skip) +- Live status from SSE +""" + +from __future__ import annotations + +import json +from typing import Any + +from PySide6.QtCore import Qt, Signal, Slot, QTimer, QMimeData +from PySide6.QtGui import QColor, QDragEnterEvent, QDropEvent, QKeySequence, QShortcut +from PySide6.QtWidgets import ( + QWidget, + QVBoxLayout, + QHBoxLayout, + QLabel, + QPushButton, + QListWidget, + QListWidgetItem, + QProgressBar, + QGroupBox, + QFrame, + QSizePolicy, + QAbstractItemView, + QMenu, + QMessageBox, +) + +from aare.common.automation_models import ( + QueueItem, + QueueItemStatus, + RuntimeState, + ControlState, + WorkflowStepRecord, +) +from aare.common.models import SampleShortInfo, SampleShortInfoList +from aare.common.logger_config import setup_logger + +logger = setup_logger("aareGUI") + + +class StepProgressWidget(QWidget): + """Shows progress through workflow steps.""" + + def __init__(self, parent: QWidget | None = None): + super().__init__(parent) + self._steps: list[WorkflowStepRecord] = [] + self._current_index = 0 + self._setup_ui() + + def _setup_ui(self) -> None: + layout = QVBoxLayout(self) + layout.setContentsMargins(4, 4, 4, 4) + layout.setSpacing(2) + + self._step_labels: list[QLabel] = [] + + self._container = QWidget() + self._container_layout = QVBoxLayout(self._container) + self._container_layout.setContentsMargins(0, 0, 0, 0) + self._container_layout.setSpacing(2) + layout.addWidget(self._container) + + def set_steps(self, steps: list[WorkflowStepRecord], current_index: int) -> None: + """Update the step display.""" + self._steps = steps + self._current_index = current_index + + for lbl in self._step_labels: + lbl.deleteLater() + self._step_labels.clear() + + for i, step in enumerate(steps): + lbl = QLabel(f"{i + 1}. {step.kind}") + lbl.setStyleSheet(self._style_for_step(i, step.status)) + self._container_layout.addWidget(lbl) + self._step_labels.append(lbl) + + def update_step(self, index: int, status: str) -> None: + """Update a single step's status.""" + if 0 <= index < len(self._step_labels): + self._step_labels[index].setStyleSheet(self._style_for_step(index, status)) + + def _style_for_step(self, index: int, status: str) -> str: + """Get stylesheet for step based on status.""" + base = "padding: 4px; border-radius: 3px; " + + if status == "success": + return base + "background-color: #90EE90; color: #006400;" + elif status == "running": + return base + "background-color: #87CEEB; color: #00008B; font-weight: bold;" + elif status == "failed": + return base + "background-color: #FFB6C1; color: #8B0000;" + elif status == "skipped": + return base + "background-color: #D3D3D3; color: #696969; text-decoration: line-through;" + elif status == "paused": + return base + "background-color: #FFE4B5; color: #8B4513;" + else: # pending + return base + "background-color: #F0F0F0; color: #808080;" + + +class DraggableQueueListWidget(QListWidget): + """ + QListWidget that accepts drops from TellSamplePanel. + + Supports: + - Drag-drop samples from tell_sample_panel + - Internal reordering via drag + - Delete key to remove items + """ + + samples_dropped = Signal(list) # list[SampleShortInfo] + item_reordered = Signal(str, int) # item_id, new_index + delete_requested = Signal(list) # list[item_ids] + + def __init__(self, parent: QWidget | None = None): + super().__init__(parent) + + self.setAcceptDrops(True) + self.setDragEnabled(True) + self.setDragDropMode(QAbstractItemView.DragDropMode.DragDrop) + self.setDefaultDropAction(Qt.DropAction.MoveAction) + self.setSelectionMode(QAbstractItemView.SelectionMode.ExtendedSelection) + self.setContextMenuPolicy(Qt.ContextMenuPolicy.CustomContextMenu) + self.customContextMenuRequested.connect(self._show_context_menu) + + # Delete shortcut + self._delete_shortcut = QShortcut(QKeySequence.StandardKey.Delete, self) + self._delete_shortcut.activated.connect(self._on_delete_pressed) + + # Store item_id -> row mapping + self._item_ids: list[str] = [] + + def set_item_ids(self, item_ids: list[str]) -> None: + """Track item IDs for reordering.""" + self._item_ids = item_ids + + def dragEnterEvent(self, event: QDragEnterEvent) -> None: + """Accept drops from sample panels.""" + mime = event.mimeData() + if mime.hasText(): + # Check if it's sample data (JSON) + try: + text = mime.text() + if text.startswith("{") or text.startswith("["): + event.acceptProposedAction() + return + except Exception: + pass + # Accept internal moves + if event.source() == self: + event.acceptProposedAction() + return + event.ignore() + + def dragMoveEvent(self, event) -> None: + """Show drop indicator.""" + if event.mimeData().hasText() or event.source() == self: + event.acceptProposedAction() + else: + event.ignore() + + def dropEvent(self, event: QDropEvent) -> None: + """Handle drop - either samples from panel or internal reorder.""" + mime = event.mimeData() + + if mime.hasText(): + text = mime.text() + try: + # Try to parse as SampleShortInfoList + sample_list = SampleShortInfoList.model_validate_json(text) + if sample_list.s: + self.samples_dropped.emit(sample_list.s) + event.acceptProposedAction() + return + except Exception: + pass + + try: + # Try single sample + sample = SampleShortInfo.model_validate_json(text) + self.samples_dropped.emit([sample]) + event.acceptProposedAction() + return + except Exception: + pass + + # Internal reorder + if event.source() == self: + # Get drop position + drop_row = self.indexAt(event.position().toPoint()).row() + if drop_row < 0: + drop_row = self.count() + + # Get selected items + selected = self.selectedItems() + if selected and self._item_ids: + for item in selected: + row = self.row(item) + if 0 <= row < len(self._item_ids): + item_id = self._item_ids[row] + self.item_reordered.emit(item_id, drop_row) + + event.acceptProposedAction() + return + + event.ignore() + + def _on_delete_pressed(self) -> None: + """Handle delete key press.""" + selected = self.selectedItems() + if not selected: + return + + item_ids = [] + for item in selected: + row = self.row(item) + if 0 <= row < len(self._item_ids): + item_ids.append(self._item_ids[row]) + + if item_ids: + self.delete_requested.emit(item_ids) + + def _show_context_menu(self, pos) -> None: + """Show context menu for queue items.""" + item = self.itemAt(pos) + if not item: + return + + row = self.row(item) + if row < 0 or row >= len(self._item_ids): + return + + menu = QMenu(self) + + delete_action = menu.addAction("🗑 Remove from queue") + start_action = menu.addAction("▶ Start this item") + + menu.addSeparator() + move_top_action = menu.addAction("⬆ Move to top") + move_bottom_action = menu.addAction("⬇ Move to bottom") + + action = menu.exec_(self.mapToGlobal(pos)) + + if action == delete_action: + self.delete_requested.emit([self._item_ids[row]]) + elif action == start_action: + # Will need to emit a signal for this + pass + elif action == move_top_action: + self.item_reordered.emit(self._item_ids[row], 0) + elif action == move_bottom_action: + self.item_reordered.emit(self._item_ids[row], 999999) + + +class WorkflowPanel(QWidget): + """ + Main workflow panel combining queue display and controls. + + Supports: + - Drag-drop samples from TELL sample panel + - Queue management (delete, reorder, clear, add all) + - Step-by-step guided mode + - Full automation mode + """ + + # Signals for DAQ worker + request_queue_refresh = Signal() + request_add_sample = Signal(object) # SampleShortInfo + request_add_samples = Signal(list) # list[SampleShortInfo] + request_delete_item = Signal(str) # item_id + request_move_item = Signal(str, int) # item_id, new_order_index + request_clear_queue = Signal() + request_start_item = Signal(str) # item_id + request_next_step = Signal() + request_pause = Signal() + request_resume = Signal() + request_abort = Signal() + request_skip = Signal() + request_start_automation = Signal() + request_stop_automation = Signal() + + def __init__(self, parent: QWidget | None = None): + super().__init__(parent) + self._queue_items: list[QueueItem] = [] + self._all_samples: list[SampleShortInfo] = [] # Cache for "Add All" + self._runtime: RuntimeState | None = None + self._control: ControlState | None = None + self._automation_enabled = False + self._setup_ui() + + def _setup_ui(self) -> None: + layout = QVBoxLayout(self) + layout.setContentsMargins(8, 8, 8, 8) + layout.setSpacing(8) + + # === Status Section === + status_group = QGroupBox("Current Status") + status_layout = QVBoxLayout(status_group) + + self._status_label = QLabel("⏹️ Idle") + self._status_label.setStyleSheet("font-size: 14px; font-weight: bold;") + status_layout.addWidget(self._status_label) + + self._current_item_label = QLabel("No item running") + status_layout.addWidget(self._current_item_label) + + self._step_progress = StepProgressWidget() + status_layout.addWidget(self._step_progress) + + layout.addWidget(status_group) + + # === Control Buttons === + controls_group = QGroupBox("Controls") + controls_layout = QVBoxLayout(controls_group) + + # Mode toggle + mode_layout = QHBoxLayout() + self._mode_label = QLabel("Mode:") + self._guided_btn = QPushButton("Guided Manual") + self._guided_btn.setCheckable(True) + self._guided_btn.setChecked(True) + self._auto_btn = QPushButton("Automation") + self._auto_btn.setCheckable(True) + + self._guided_btn.clicked.connect(self._on_guided_mode) + self._auto_btn.clicked.connect(self._on_automation_mode) + + mode_layout.addWidget(self._mode_label) + mode_layout.addWidget(self._guided_btn) + mode_layout.addWidget(self._auto_btn) + mode_layout.addStretch() + controls_layout.addLayout(mode_layout) + + # Step controls + step_layout = QHBoxLayout() + + self._next_btn = QPushButton("Next Step") + self._next_btn.setStyleSheet("background-color: #4CAF50; color: white;") + self._next_btn.clicked.connect(self.request_next_step.emit) + + self._skip_btn = QPushButton("Skip") + self._skip_btn.setStyleSheet("background-color: #FF9800; color: white;") + self._skip_btn.clicked.connect(self.request_skip.emit) + + self._pause_btn = QPushButton("Pause") + self._pause_btn.setStyleSheet("background-color: #2196F3; color: white;") + self._pause_btn.clicked.connect(self._on_pause_resume) + + self._abort_btn = QPushButton("Abort") + self._abort_btn.setStyleSheet("background-color: #F44336; color: white;") + self._abort_btn.clicked.connect(self.request_abort.emit) + + step_layout.addWidget(self._next_btn) + step_layout.addWidget(self._skip_btn) + step_layout.addWidget(self._pause_btn) + step_layout.addWidget(self._abort_btn) + + controls_layout.addLayout(step_layout) + layout.addWidget(controls_group) + + # === Queue Section === + queue_group = QGroupBox("Queue (drag samples here)") + queue_layout = QVBoxLayout(queue_group) + + # Info label + self._drop_hint = QLabel("💡 Drag samples from the Sample List to add them") + self._drop_hint.setStyleSheet("color: #666; font-style: italic;") + queue_layout.addWidget(self._drop_hint) + + # Draggable queue list + self._queue_list = DraggableQueueListWidget() + self._queue_list.setMinimumHeight(150) + self._queue_list.itemDoubleClicked.connect(self._on_item_double_clicked) + self._queue_list.samples_dropped.connect(self._on_samples_dropped) + self._queue_list.delete_requested.connect(self._on_delete_requested) + self._queue_list.item_reordered.connect(self._on_item_reordered) + queue_layout.addWidget(self._queue_list) + + # Queue management buttons + queue_btn_layout = QHBoxLayout() + + self._add_all_btn = QPushButton("➕ Add All Samples") + self._add_all_btn.clicked.connect(self._on_add_all_clicked) + self._add_all_btn.setToolTip("Add all samples from the Sample List to the queue") + + self._remove_selected_btn = QPushButton("🗑 Remove Selected") + self._remove_selected_btn.clicked.connect(self._on_remove_selected) + + self._clear_btn = QPushButton("✖ Clear Queue") + self._clear_btn.clicked.connect(self._on_clear_queue) + + self._refresh_btn = QPushButton("🔄") + self._refresh_btn.setFixedWidth(40) + self._refresh_btn.setToolTip("Refresh queue") + self._refresh_btn.clicked.connect(self.request_queue_refresh.emit) + + queue_btn_layout.addWidget(self._add_all_btn) + queue_btn_layout.addWidget(self._remove_selected_btn) + queue_btn_layout.addWidget(self._clear_btn) + queue_btn_layout.addStretch() + queue_btn_layout.addWidget(self._refresh_btn) + queue_layout.addLayout(queue_btn_layout) + + layout.addWidget(queue_group) + + # Initial state + self._update_button_states() + + def _on_guided_mode(self) -> None: + """Switch to guided manual mode.""" + self._guided_btn.setChecked(True) + self._auto_btn.setChecked(False) + if self._automation_enabled: + self.request_stop_automation.emit() + self._update_button_states() + + def _on_automation_mode(self) -> None: + """Switch to automation mode.""" + self._auto_btn.setChecked(True) + self._guided_btn.setChecked(False) + if not self._automation_enabled: + self.request_start_automation.emit() + self._update_button_states() + + def _on_pause_resume(self) -> None: + """Toggle pause/resume.""" + if self._runtime and self._runtime.paused: + self.request_resume.emit() + else: + self.request_pause.emit() + + def _on_item_double_clicked(self, item: QListWidgetItem) -> None: + """Start processing double-clicked item.""" + idx = self._queue_list.row(item) + if 0 <= idx < len(self._queue_items): + queue_item = self._queue_items[idx] + if queue_item.status == QueueItemStatus.PENDING: + self.request_start_item.emit(queue_item.item_id) + + def _on_samples_dropped(self, samples: list[SampleShortInfo]) -> None: + """Handle samples dropped onto the queue.""" + logger.info(f"Adding {len(samples)} samples to workflow queue") + for sample in samples: + self.request_add_sample.emit(sample) + # Refresh after a short delay to let server process + QTimer.singleShot(500, self.request_queue_refresh.emit) + + def _on_delete_requested(self, item_ids: list[str]) -> None: + """Handle delete request from queue list.""" + for item_id in item_ids: + self.request_delete_item.emit(item_id) + QTimer.singleShot(300, self.request_queue_refresh.emit) + + def _on_item_reordered(self, item_id: str, new_index: int) -> None: + """Handle item reorder.""" + self.request_move_item.emit(item_id, new_index) + QTimer.singleShot(300, self.request_queue_refresh.emit) + + def _on_remove_selected(self) -> None: + """Remove selected items from queue.""" + selected = self._queue_list.selectedItems() + if not selected: + return + + item_ids = self._queue_list._item_ids + for item in selected: + row = self._queue_list.row(item) + if 0 <= row < len(item_ids): + self.request_delete_item.emit(item_ids[row]) + + QTimer.singleShot(300, self.request_queue_refresh.emit) + + def _on_clear_queue(self) -> None: + """Clear all items from queue.""" + if not self._queue_items: + return + + reply = QMessageBox.question( + self, + "Clear Queue", + f"Remove all {len(self._queue_items)} items from the queue?", + QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No, + ) + + if reply == QMessageBox.StandardButton.Yes: + self.request_clear_queue.emit() + QTimer.singleShot(500, self.request_queue_refresh.emit) + + def _on_add_all_clicked(self) -> None: + """Add all available samples to the queue.""" + if not self._all_samples: + QMessageBox.information( + self, + "No Samples", + "No samples available to add. Load samples in the Sample List first.", + ) + return + + reply = QMessageBox.question( + self, + "Add All Samples", + f"Add all {len(self._all_samples)} samples to the workflow queue?", + QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No, + ) + + if reply == QMessageBox.StandardButton.Yes: + self.request_add_samples.emit(self._all_samples) + QTimer.singleShot(500, self.request_queue_refresh.emit) + + def _update_button_states(self) -> None: + """Update button enabled/disabled states based on current state.""" + is_running = self._runtime is not None and self._runtime.running + is_paused = self._runtime is not None and self._runtime.paused + is_automation = self._auto_btn.isChecked() + + # In automation mode, hide the Next button + self._next_btn.setVisible(not is_automation) + self._next_btn.setEnabled(is_running and not is_paused) + + self._skip_btn.setEnabled(is_running) + self._abort_btn.setEnabled(is_running) + + if is_paused: + self._pause_btn.setText("Resume") + self._pause_btn.setStyleSheet("background-color: #4CAF50; color: white;") + else: + self._pause_btn.setText("Pause") + self._pause_btn.setStyleSheet("background-color: #2196F3; color: white;") + + self._pause_btn.setEnabled(is_running) + + # Update drop hint visibility + has_items = len(self._queue_items) > 0 + self._drop_hint.setVisible(not has_items) + + # === Public Update Methods === + + @Slot(list) + def update_queue(self, items: list[QueueItem]) -> None: + """Update queue display.""" + self._queue_items = items + self._queue_list.clear() + + item_ids = [] + for item in items: + status_emoji = { + QueueItemStatus.PENDING: "⏳", + QueueItemStatus.RUNNING: "▶️", + QueueItemStatus.COMPLETED: "✅", + QueueItemStatus.FAILED: "❌", + QueueItemStatus.ABORTED: "🛑", + QueueItemStatus.SKIPPED: "⏭️", + }.get(item.status, "❓") + + display_text = f"{status_emoji} {item.sample_name or item.item_id}" + list_item = QListWidgetItem(display_text) + + if item.status == QueueItemStatus.RUNNING: + list_item.setBackground(QColor("#E6F3FF")) + elif item.status == QueueItemStatus.COMPLETED: + list_item.setBackground(QColor("#E6FFE6")) + elif item.status == QueueItemStatus.FAILED: + list_item.setBackground(QColor("#FFE6E6")) + + self._queue_list.addItem(list_item) + item_ids.append(item.item_id) + + self._queue_list.set_item_ids(item_ids) + self._update_button_states() + + @Slot(object) + def update_runtime(self, runtime: RuntimeState) -> None: + """Update from runtime state.""" + self._runtime = runtime + + if runtime.running: + if runtime.paused: + self._status_label.setText("⏸️ Paused") + self._status_label.setStyleSheet( + "font-size: 14px; font-weight: bold; color: #FF9800;" + ) + else: + self._status_label.setText("▶️ Running") + self._status_label.setStyleSheet( + "font-size: 14px; font-weight: bold; color: #4CAF50;" + ) + + self._current_item_label.setText( + f"Item: {runtime.current_item_id or 'Unknown'}" + ) + else: + self._status_label.setText("⏹️ Idle") + self._status_label.setStyleSheet( + "font-size: 14px; font-weight: bold; color: #757575;" + ) + self._current_item_label.setText("No item running") + + self._update_button_states() + + @Slot(object) + def update_control(self, control: ControlState) -> None: + """Update from control state.""" + self._control = control + self._update_button_states() + + @Slot(object) + def update_current_item(self, item: QueueItem) -> None: + """Update step progress for current item.""" + self._step_progress.set_steps(item.steps, item.current_step_index) + + for i, step in enumerate(item.steps): + self._step_progress.update_step(i, step.status) + + @Slot(bool) + def update_automation_enabled(self, enabled: bool) -> None: + """Update automation mode state.""" + self._automation_enabled = enabled + self._auto_btn.setChecked(enabled) + self._guided_btn.setChecked(not enabled) + self._update_button_states() + + @Slot(list) + def update_available_samples(self, samples: list[SampleShortInfo]) -> None: + """Update the cached sample list for 'Add All' functionality.""" + self._all_samples = samples + + @Slot(str, str) + def on_workflow_event(self, event_type: str, message: str) -> None: + """Handle workflow events from SSE stream.""" + logger.debug(f"Workflow event: {event_type} - {message}") + + if event_type in ("item_started", "item_completed", "step_completed", "step_skipped"): + self.request_queue_refresh.emit() \ No newline at end of file diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py index 7984bf57..df49f5f7 100644 --- a/src/aare/gui/threads/daq_worker.py +++ b/src/aare/gui/threads/daq_worker.py @@ -51,6 +51,14 @@ class DAQWorker(QObject): last_error_payload_changed = Signal(dict) last_error_payloads_changed = Signal(list) + # Workflow signals + workflow_queue_loaded = Signal(list) # list[QueueItem] + workflow_runtime_changed = Signal(object) # RuntimeState + workflow_control_changed = Signal(object) # ControlState + workflow_item_updated = Signal(object) # QueueItem + workflow_automation_status = Signal(bool) # enabled + workflow_event = Signal(object) # WorkflowEvent + def __init__(self, base_url: str | None, token: str, parent=None): super().__init__(parent) self._active_status_error_key = None @@ -1252,4 +1260,162 @@ class DAQWorker(QObject): query.append(f"message={quote(message)}") suffix = f"?{'&'.join(query)}" if query else "" - self.generic_post(f"samcam/send_screenshot_db{suffix}") \ No newline at end of file + self.generic_post(f"samcam/send_screenshot_db{suffix}") + + # ───────────────────────────────────────────── + # Workflow API methods + # ───────────────────────────────────────────── + + @Slot() + def workflow_load_queue(self): + """Load workflow queue.""" + if self.__base_url is None: + return + + request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/queue")) + request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode()) + reply = self.__net_manager.get(request) + reply.finished.connect(lambda: self._handle_workflow_queue_response(reply)) + + def _handle_workflow_queue_response(self, reply: QNetworkReply): + try: + response_data = self.handle_response(reply) + data = json.loads(response_data) + from aare.common.automation_models import QueueItem + items = [QueueItem.model_validate(i) for i in data.get("items", [])] + self.workflow_queue_loaded.emit(items) + except Exception as e: + logger.error(f"Failed to load workflow queue: {e}") + + @Slot(object) # SampleShortInfo + def workflow_add_sample(self, sample: "SampleShortInfo"): + """Add a sample to the workflow queue.""" + if self.__base_url is None: + logger.info(f"POST /workflow/queue: {sample.sample_name}") + return + + from aare.common.automation_models import CreateQueueItemRequest + request_data = CreateQueueItemRequest( + sample_id=sample.db_id, + sample_name=sample.sample_name, + priority=int(sample.priority or 100), + ) + + request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/queue")) + request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode()) + request.setRawHeader(b"Content-Type", b"application/json") + reply = self.__net_manager.post(request, QByteArray(request_data.model_dump_json().encode())) + reply.finished.connect(lambda: self.handle_req_response(reply)) + + @Slot(list) # list[SampleShortInfo] + def workflow_add_samples(self, samples: list): + """Add multiple samples to the workflow queue.""" + for sample in samples: + self.workflow_add_sample(sample) + + @Slot(str) + def workflow_delete_item(self, item_id: str): + """Delete an item from the workflow queue.""" + if self.__base_url is None: + logger.info(f"DELETE /workflow/queue/{item_id}") + return + + self.generic_delete(f"workflow/queue/{item_id}") + + @Slot(str, int) + def workflow_move_item(self, item_id: str, new_order_index: int): + """Move/reorder an item in the workflow queue.""" + if self.__base_url is None: + logger.info(f"POST /workflow/queue/{item_id}/move order={new_order_index}") + return + + body = json.dumps({"new_order_index": new_order_index}) + self.generic_post(f"workflow/queue/{item_id}/move", body) + + @Slot() + def workflow_clear_queue(self): + """Clear all pending items from the queue.""" + # Load queue first, then delete all pending items + if self.__base_url is None: + return + + request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/queue")) + request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode()) + reply = self.__net_manager.get(request) + reply.finished.connect(lambda: self._handle_clear_queue_response(reply)) + + def _handle_clear_queue_response(self, reply: QNetworkReply): + try: + response_data = self.handle_response(reply) + data = json.loads(response_data) + from aare.common.automation_models import QueueItem, QueueItemStatus + items = [QueueItem.model_validate(i) for i in data.get("items", [])] + + # Delete all non-running items + for item in items: + if item.status != QueueItemStatus.RUNNING: + self.workflow_delete_item(item.item_id) + except Exception as e: + logger.error(f"Failed to clear workflow queue: {e}") + + @Slot(str) + def workflow_start_item(self, item_id: str): + """Start processing a queue item.""" + self.generic_post(f"workflow/start/{item_id}") + QTimer.singleShot(500, self.workflow_load_queue) + + @Slot() + def workflow_next_step(self): + """Run next step in guided mode.""" + if self.__base_url is None: + return + + request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/next")) + request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode()) + request.setRawHeader(b"Content-Type", b"application/json") + reply = self.__net_manager.post(request, QByteArray(b"")) + reply.finished.connect(lambda: self._handle_workflow_step_response(reply)) + + def _handle_workflow_step_response(self, reply: QNetworkReply): + try: + response_data = self.handle_response(reply) + data = json.loads(response_data) + if data.get("item"): + from aare.common.automation_models import QueueItem + item = QueueItem.model_validate(data["item"]) + self.workflow_item_updated.emit(item) + # Also refresh queue + self.workflow_load_queue() + except Exception as e: + logger.error(f"Workflow step error: {e}") + self.http_error.emit(str(e)) + + @Slot() + def workflow_pause(self): + """Request pause.""" + self.generic_post("workflow/control/pause") + + @Slot() + def workflow_resume(self): + """Request resume.""" + self.generic_post("workflow/control/resume") + + @Slot() + def workflow_abort(self): + """Request abort.""" + self.generic_post("workflow/control/abort") + + @Slot() + def workflow_skip(self): + """Request skip.""" + self.generic_post("workflow/control/skip") + + @Slot() + def workflow_start_automation(self): + """Start automation mode.""" + self.generic_post("workflow/automation/start") + + @Slot() + def workflow_stop_automation(self): + """Stop automation mode.""" + self.generic_post("workflow/automation/stop") \ No newline at end of file diff --git a/src/aare/gui/threads/workflow_sse_client.py b/src/aare/gui/threads/workflow_sse_client.py new file mode 100644 index 00000000..f1f4d857 --- /dev/null +++ b/src/aare/gui/threads/workflow_sse_client.py @@ -0,0 +1,158 @@ +""" +SSE client for workflow events. + +Connects to the /workflow/sse endpoint and emits signals for: +- Runtime state changes +- Control state changes +- Individual workflow events +""" + +from __future__ import annotations + +import json + +from PySide6.QtCore import QObject, Signal, Slot, QTimer +from PySide6.QtNetwork import QNetworkAccessManager, QNetworkRequest, QNetworkReply +from PySide6.QtCore import QUrl, QByteArray + +from aare.common.automation_models import RuntimeState, ControlState, WorkflowEvent +from aare.common.logger_config import setup_logger + +logger = setup_logger("aareGUI") + + +class WorkflowSSEClient(QObject): + """ + SSE client that subscribes to workflow events. + + Emits signals when state changes are received. + """ + + # Signals + runtime_changed = Signal(object) # RuntimeState + control_changed = Signal(object) # ControlState + workflow_event = Signal(object) # WorkflowEvent + connected = Signal() + disconnected = Signal() + error = Signal(str) + + def __init__( + self, + base_url: str, + token: str, + parent: QObject | None = None, + ): + super().__init__(parent) + + self._base_url = base_url + self._token = token + self._manager = QNetworkAccessManager(self) + self._reply: QNetworkReply | None = None + self._buffer = "" + + # Reconnection + self._reconnect_timer = QTimer(self) + self._reconnect_timer.setInterval(5000) # 5 seconds + self._reconnect_timer.timeout.connect(self.connect) + self._should_reconnect = False + + def connect(self) -> None: + """Start SSE connection.""" + if self._reply is not None: + return # Already connected + + self._should_reconnect = True + + url = QUrl(f"{self._base_url}/workflow/sse") + # Add token as query param for SSE (can't use headers easily) + url.setQuery(f"token={self._token}") + + request = QNetworkRequest(url) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) + request.setRawHeader(b"Accept", b"text/event-stream") + request.setRawHeader(b"Cache-Control", b"no-cache") + + self._reply = self._manager.get(request) + self._reply.readyRead.connect(self._on_data_ready) + self._reply.finished.connect(self._on_finished) + self._reply.errorOccurred.connect(self._on_error) + + self._reconnect_timer.stop() + logger.debug("Workflow SSE: connecting...") + + def disconnect(self) -> None: + """Stop SSE connection.""" + self._should_reconnect = False + self._reconnect_timer.stop() + + if self._reply is not None: + self._reply.abort() + self._reply.deleteLater() + self._reply = None + + self.disconnected.emit() + + @Slot() + def _on_data_ready(self) -> None: + """Handle incoming SSE data.""" + if self._reply is None: + return + + data = self._reply.readAll().data().decode("utf-8") + self._buffer += data + + # Process complete events (separated by double newlines) + while "\n\n" in self._buffer: + event_data, self._buffer = self._buffer.split("\n\n", 1) + self._parse_event(event_data) + + def _parse_event(self, event_data: str) -> None: + """Parse a single SSE event.""" + event_type = "message" + data_lines = [] + + for line in event_data.split("\n"): + if line.startswith("event:"): + event_type = line[6:].strip() + elif line.startswith("data:"): + data_lines.append(line[5:].strip()) + + if not data_lines: + return + + data_str = "\n".join(data_lines) + + try: + if event_type == "runtime": + runtime = RuntimeState.model_validate_json(data_str) + self.runtime_changed.emit(runtime) + elif event_type == "control": + control = ControlState.model_validate_json(data_str) + self.control_changed.emit(control) + elif event_type == "workflow_event": + event = WorkflowEvent.model_validate_json(data_str) + self.workflow_event.emit(event) + except Exception as e: + logger.warning(f"Workflow SSE: failed to parse {event_type}: {e}") + + @Slot() + def _on_finished(self) -> None: + """Handle connection finished.""" + if self._reply is not None: + self._reply.deleteLater() + self._reply = None + + self._buffer = "" + self.disconnected.emit() + + # Reconnect if desired + if self._should_reconnect: + logger.debug("Workflow SSE: disconnected, will reconnect...") + self._reconnect_timer.start() + + @Slot(QNetworkReply.NetworkError) + def _on_error(self, error: QNetworkReply.NetworkError) -> None: + """Handle connection error.""" + error_msg = self._reply.errorString() if self._reply else str(error) + logger.warning(f"Workflow SSE error: {error_msg}") + self.error.emit(error_msg) -- 2.54.0 From 3c1cf09417e0d3ace3f42b4600baf4de748ec540 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Thu, 26 Mar 2026 11:20:31 +0100 Subject: [PATCH 52/56] aaredaq: continuing work on baton exchange procedure --- src/aare/common/auth_models.py | 51 ++++ src/aare/daq/auth.py | 264 +++++++++++++++++++ src/aare/daq/config.py | 129 ++++++++- src/aare/daq/server.py | 137 +++++++++- src/aare/gui/main_window.py | 143 ++++++++++ src/aare/gui/threads/daq_worker.py | 160 ++++++++++- src/aare/gui/widgets/alert_banner.py | 75 ++++++ src/aare/gui/widgets/baton_request_dialog.py | 220 ++++++++++++++++ src/aare/gui/widgets/status_bar.py | 261 +++++++++++++++--- 9 files changed, 1402 insertions(+), 38 deletions(-) create mode 100644 src/aare/common/auth_models.py create mode 100644 src/aare/gui/widgets/baton_request_dialog.py diff --git a/src/aare/common/auth_models.py b/src/aare/common/auth_models.py new file mode 100644 index 00000000..6897751c --- /dev/null +++ b/src/aare/common/auth_models.py @@ -0,0 +1,51 @@ +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 # True if YOU are being asked to give up baton \ No newline at end of file diff --git a/src/aare/daq/auth.py b/src/aare/daq/auth.py index 3af7f247..06d24829 100644 --- a/src/aare/daq/auth.py +++ b/src/aare/daq/auth.py @@ -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 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,6 +25,7 @@ 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"] @@ -113,3 +118,262 @@ def check_jwt_staff(cfg: BeamlineConfig, data: TokenData) -> None: def force_current_sesion(cfg: BeamlineConfig, data: TokenData) -> None: cfg.force_set_active_session(data.session, 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: + """Get full baton status for a user.""" + pending = cfg.pending_baton_request + + if pending is not None and pending.status == BatonRequestStatus.PENDING: + elapsed = time.time() - pending.created_at + if elapsed >= pending.timeout_seconds: + resolve_baton_timeout_if_needed(cfg) + pending = cfg.pending_baton_request + + holder = cfg.baton_holder + queued = cfg.queued_baton_transfer + + you_are_holder = holder is not None and holder.session == data.session + you_have_pending_request = ( + pending is not None and + pending.requester_session == data.session and + pending.status == BatonRequestStatus.PENDING + ) + incoming_request = ( + pending is not None and + holder is not None and + holder.session == data.session and + pending.status == BatonRequestStatus.PENDING + ) + + return BatonStatus( + holder=holder, + pending_request=pending, + queued_transfer=queued, + you_are_holder=you_are_holder, + you_have_pending_request=you_have_pending_request, + incoming_request=incoming_request + ) + + +def request_baton(cfg: BeamlineConfig, data: TokenData) -> dict: + """ + Request the baton (control) of the beamline. + """ + 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 + + if holder and holder.is_staff and not data.staff: + return { + "error": True, + "message": "Cannot request baton from staff. Please ask them directly." + } + + 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 user'}" + } + +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"} \ No newline at end of file diff --git a/src/aare/daq/config.py b/src/aare/daq/config.py index 61196563..ba8f8328 100644 --- a/src/aare/daq/config.py +++ b/src/aare/daq/config.py @@ -18,6 +18,13 @@ 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 @@ -152,6 +159,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 +169,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") @@ -624,4 +751,4 @@ class BeamlineConfig: self.__client.set(f"{self.__bl}:failed_mount_count", count) def increment_failed_mount_count(self) -> int: - return int(self.__client.incr(f"{self.__bl}:failed_mount_count")) + return int(self.__client.incr(f"{self.__bl}:failed_mount_count")) \ No newline at end of file diff --git a/src/aare/daq/server.py b/src/aare/daq/server.py index 45f57d4f..94bcfe09 100644 --- a/src/aare/daq/server.py +++ b/src/aare/daq/server.py @@ -7,6 +7,8 @@ import json import cv2 import urllib3 import uvicorn + +from aare.common.auth_models import BatonStatus 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 @@ -659,11 +661,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 +706,129 @@ 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 + """ + data = auth.parse_token(token) + auth.check_jwt_ro(cfg, data) + + # holder = cfg.baton_holder + # if holder is not None and holder.is_staff and not data.staff: + # return { + # "error": True, + # "message": "Cannot request baton from staff. Please ask them directly." + # } + + 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) + auth.check_jwt_ro(cfg, data) + 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) + + resolved = auth.resolve_baton_timeout_if_needed(cfg) + if resolved is not None: + return resolved.model_dump() + + pending = cfg.pending_baton_request + if pending is None: + return {"no_pending": True} + + if pending.requester_session != data.session: + return {"not_your_request": True} + + 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) + + +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: diff --git a/src/aare/gui/main_window.py b/src/aare/gui/main_window.py index dcc40ce0..3e2b8522 100644 --- a/src/aare/gui/main_window.py +++ b/src/aare/gui/main_window.py @@ -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 @@ -42,6 +44,7 @@ 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 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 @@ -71,6 +74,8 @@ class MainWindow(QMainWindow): self._controls_help_dialog = None self._cleanup_done = False + self._waiting_for_baton_response: bool = False + # Tutorial manager (define tutorials after widgets exist) self.tutorial_manager = TutorialManager(self) @@ -283,6 +288,12 @@ 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._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.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) @@ -394,9 +405,18 @@ 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.daq.baton_status_changed.connect(self.status_bar.update_baton_status) + self.daq.baton_status_changed.connect(self.status_bar.update_baton_status) + 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) @@ -404,6 +424,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) @@ -657,6 +678,122 @@ class MainWindow(QMainWindow): self.__mounting = False self.video_tab.setCurrentIndex(0) + + # ========== BATON DIALOG HANDLING ========== + + @Slot(dict) + def _on_baton_incoming_request(self, info: dict): + """Show dialog when someone requests our baton.""" + requester = info.get("requester", "Another user") + timeout = info.get("timeout", 30) + + # Don't show multiple dialogs + 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._on_baton_dialog_accepted) + self._baton_request_dialog.refused_signal.connect(self._on_baton_dialog_refused) + self._baton_request_dialog.show() + + @Slot() + def _on_baton_dialog_accepted(self): + """User clicked Accept in baton dialog.""" + self.daq.respond_to_baton_request(True) + self._baton_request_dialog = None + + @Slot() + def _on_baton_dialog_refused(self): + """User clicked Refuse in baton dialog.""" + self.daq.respond_to_baton_request(False) + self._baton_request_dialog = None + + @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) + logger.info("Baton acquired") + elif result.get("pending"): + self._waiting_for_baton_response = True + timeout = result.get("timeout_seconds", 30) + holder = result.get("message", "Waiting for response...") + 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") + elif result.get("already_holder"): + self._waiting_for_baton_response = False + # Don't show anything - user already has control + logger.debug("Already baton holder") + 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')}") + + @Slot(dict) + def _on_baton_response_result(self, result: dict): + """Handle result after we responded to someone else's request.""" + if result.get("accepted"): + self.alert_banner.show_message("Control transferred", False) + elif result.get("refused"): + self.alert_banner.show_message("Request declined", False) + + @Slot(dict) + def _on_baton_request_received(self, info: dict): + """Incoming baton requests are now owned by StatusBar.""" + requester = info.get("requester", "Another user") + timeout = info.get("timeout", 30) + logger.info(f"Incoming baton request from {requester} ({timeout}s)") + + @Slot(BatonStatus) + def _on_baton_status_changed(self, status: BatonStatus): + """Handle baton status updates from SSE stream.""" + if self._waiting_for_baton_response: + pending = status.pending_request + if status.you_are_holder: + self._waiting_for_baton_response = False + self.alert_banner.show_message("Baton acquired!", False) + elif pending and pending.status == BatonRequestStatus.REFUSED: + self._waiting_for_baton_response = False + holder_name = pending.holder_username or "Current user" + self.alert_banner.show_message(f"Request declined by {holder_name}", True) + elif pending is None: + self._waiting_for_baton_response = False + self.alert_banner.clear_message() + elif status.you_have_pending_request: + import time + elapsed = time.time() - pending.created_at + remaining = max(0, int(pending.timeout_seconds - elapsed)) + if remaining > 0: + self.alert_banner.show_waiting( + f"Requesting control from {pending.holder_username or 'current user'}", + remaining + ) + else: + self.alert_banner.show_waiting("Processing timeout...", 0) + + @Slot(dict) + def _on_baton_timeout_checked(self, result: dict): + """Refresh waiting UI when the backend confirms timeout state.""" + 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) + elif result.get("granted"): + self._waiting_for_baton_response = False + self.alert_banner.show_message("Baton acquired!", False) + elif result.get("queued"): + self._waiting_for_baton_response = True + self.alert_banner.show_waiting("Control transfer queued - waiting for beamline") + def _restore_window_state(self) -> None: settings = QSettings() geometry = settings.value("main_window/geometry") @@ -675,6 +812,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: diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py index 16f77b8d..53c4c212 100644 --- a/src/aare/gui/threads/daq_worker.py +++ b/src/aare/gui/threads/daq_worker.py @@ -8,6 +8,7 @@ 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.auth_models import BatonStatus from aare.common.coordinate import SmargonCoordinate, Coordinate from aare.common.error_codes import export_error_codes from aare.common.models import DAQStatusModel, SampleShortInfoList, SampleShortInfo, SampleCameraSettings, \ @@ -45,6 +46,12 @@ class DAQWorker(QObject): last_error_payload_changed = Signal(dict) last_error_payloads_changed = Signal(list) + 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.__token = token @@ -81,10 +88,19 @@ class DAQWorker(QObject): self._server_connected: bool | None = None self._last_server_error: str | 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() + self._baton_timeout_timer.start() def get_last_error_payload(self) -> dict: return dict(self._last_error_payload or {}) @@ -547,6 +563,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): @@ -1022,7 +1039,6 @@ class DAQWorker(QObject): out[str(k)] = str(v) return out - @Slot() @Slot() def get_error_codes(self) -> None: """ @@ -1089,4 +1105,144 @@ 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}") \ No newline at end of file + 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) + + # Detect incoming request + 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") or result.get("queued") or result.get("pending"): + self.check_baton_timeout() + 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() + 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"): + self.send_status_request() + 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}") \ No newline at end of file diff --git a/src/aare/gui/widgets/alert_banner.py b/src/aare/gui/widgets/alert_banner.py index fe10f4c7..c1ffe20b 100644 --- a/src/aare/gui/widgets/alert_banner.py +++ b/src/aare/gui/widgets/alert_banner.py @@ -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) @@ -34,6 +41,8 @@ class AlertBanner(QFrame): @Slot(str, bool) def show_message(self, msg: str, is_error: bool = True): + """Show error (red) or success (green) message.""" + self._stop_countdown() self._clear_timer.stop() if not msg: @@ -77,6 +86,72 @@ class AlertBanner(QFrame): 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() + # Don't auto-clear - let the baton status update handle that + 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._clear_timer.stop() diff --git a/src/aare/gui/widgets/baton_request_dialog.py b/src/aare/gui/widgets/baton_request_dialog.py new file mode 100644 index 00000000..90aab3b9 --- /dev/null +++ b/src/aare/gui/widgets/baton_request_dialog.py @@ -0,0 +1,220 @@ +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"{self._requester} 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( + "If the beamline is busy, transfer will occur after " + "the current operation completes." + ) + 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) \ No newline at end of file diff --git a/src/aare/gui/widgets/status_bar.py b/src/aare/gui/widgets/status_bar.py index b59ae947..7dec639a 100644 --- a/src/aare/gui/widgets/status_bar.py +++ b/src/aare/gui/widgets/status_bar.py @@ -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,17 @@ class StatusBar(QStatusBar): force_session = Signal() end_session = Signal() - close_shutter = Signal() - open_shutter = Signal() + request_baton = Signal() # New signal for baton request + cancel_baton_request = Signal() # Cancel pending request + release_baton = Signal() # Voluntarily release + 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 +47,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,22 +214,133 @@ 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) + self._baton_status = status + self._has_pending_request = status.you_have_pending_request if status else False + self._update_session_display() + + incoming = bool(status and status.incoming_request) + if incoming and not prev_incoming: + self._show_incoming_baton_request_dialog(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 _show_incoming_baton_request_dialog(self, status: BatonStatus) -> None: + if self._baton_request_dialog is not None and self._baton_request_dialog.isVisible(): + return + + 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, + }) + + self._baton_request_dialog = BatonRequestDialog( + requester=requester, + timeout_seconds=timeout, + parent=self + ) + self._baton_request_dialog.accepted_signal.connect(self._on_baton_dialog_accepted) + self._baton_request_dialog.refused_signal.connect(self._on_baton_dialog_refused) + self._baton_request_dialog.show() + + 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 + action_grab = menu.addAction("Request (Staff has control)") + action_grab.setEnabled(False) + 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): @@ -302,39 +427,109 @@ class StatusBar(QStatusBar): menu.exec() def _on_grab_clicked(self): - self.grab_session_clicked() + """Handle grab/request click - shows p-group dialog first if needed.""" + available_pgroups = self.__allowed_pgroups or [] - def after_grab(): - - if not self.__status: - logger.info("status is None when session is grabbed") - return - - check_state = self.__status.state in ( - BeamlineStateEnum.SampleAlignment, - BeamlineStateEnum.SampleExchange, - BeamlineStateEnum.DewarTransfer, - BeamlineStateEnum.Maintenance, + if not self.__is_staff and not available_pgroups: + QMessageBox.warning( + self, + "No P-Groups", + "You don't have any p-groups available. Cannot request control." ) + return - 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 self.__is_staff and not available_pgroups: + self.request_baton.emit() + return - if allowed: - self.show_change_dialog() + if len(available_pgroups) == 1: + self.set_pgroup.emit(available_pgroups[0]) + self.request_baton.emit() + return + self._show_pgroup_selection_for_baton(available_pgroups) + + def _show_pgroup_selection_for_baton(self, available_pgroups: list[str]): + """Show p-group selection dialog before requesting baton.""" + if self._pgroup_dialog_for_baton is not None and self._pgroup_dialog_for_baton.isVisible(): + return + + # Get current p-group if available + curr_pgroup = None + if self.__status and self.__status.session: + curr_pgroup = self.__status.session.current_pgroup + + # If staff, allow all p-groups + if self.__is_staff: + def _on_staff_loaded(lst: list): + try: + merged = {str(p).strip() for p in (lst or []) if p is not None and str(p).strip()} + # Add user's own p-groups too + for pg in available_pgroups: + if pg and str(pg).strip(): + merged.add(str(pg).strip()) + if not merged: + merged = set(available_pgroups) + self._show_pgroup_dialog_and_request_baton(curr_pgroup, sorted(merged)) + finally: + try: + self.staff_pgroups_loaded.disconnect(_on_staff_loaded) + except Exception as e: + logger.debug(f"Error disconnecting staff_pgroups_loaded: {e}") + + self.staff_pgroups_loaded.connect(_on_staff_loaded) + self.__list_staff_pgroups() + else: + # Non-staff: only show their own p-groups + self._show_pgroup_dialog_and_request_baton(curr_pgroup, available_pgroups) + + def _show_pgroup_dialog_and_request_baton(self, curr_pgroup: str | None, pgroups: list[str]): + """Show the dialog and request baton if accepted.""" + self._pgroup_dialog_for_baton = PGroupDialog( + curr_pgroup=curr_pgroup, + 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 and selected_pgroup.strip(): + # Validate the selection + if not self.__is_staff and selected_pgroup not in (self.__allowed_pgroups or []): + QMessageBox.warning( + self, + "Invalid P-Group", + f"P-group '{selected_pgroup}' is not in your allowed list.\n" + f"Please select from: {', '.join(self.__allowed_pgroups or [])}" + ) + else: + # Set the p-group and request baton + self.set_pgroup.emit(selected_pgroup) + self.request_baton.emit() 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}" + QMessageBox.warning( + self, + "No P-Group Selected", + "You must select a p-group to request control." ) - QTimer.singleShot(300, after_grab) + self._pgroup_dialog_for_baton = None + + def _on_release_clicked(self): + """Handle release click.""" + self.release_baton.emit() + + def _on_cancel_request_clicked(self): + """Handle cancel request click.""" + self.cancel_baton_request.emit() + + 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() -- 2.54.0 From 5a47411727acd72f949955c3f58e8761f44a21d9 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Thu, 26 Mar 2026 14:16:25 +0100 Subject: [PATCH 53/56] aaredaq: continuing work on baton exchange procedure --- src/aare/daq/auth.py | 3 + src/aare/daq/server.py | 10 -- src/aare/gui/main_window.py | 124 +++++++++++--------- src/aare/gui/threads/daq_worker.py | 18 ++- src/aare/gui/widgets/alert_banner.py | 5 +- src/aare/gui/widgets/pgroup_dialog.py | 100 +++++++++++++--- src/aare/gui/widgets/status_bar.py | 161 ++++++++++++-------------- 7 files changed, 243 insertions(+), 178 deletions(-) diff --git a/src/aare/daq/auth.py b/src/aare/daq/auth.py index 06d24829..5e946fd0 100644 --- a/src/aare/daq/auth.py +++ b/src/aare/daq/auth.py @@ -210,6 +210,9 @@ def get_baton_status(cfg: BeamlineConfig, data: TokenData) -> BatonStatus: def request_baton(cfg: BeamlineConfig, data: TokenData) -> dict: """ Request the baton (control) of the beamline. + + This should not depend on active p-group membership at the server route level. + Baton ownership rules are handled here. """ resolve_baton_timeout_if_needed(cfg) diff --git a/src/aare/daq/server.py b/src/aare/daq/server.py index 94bcfe09..c0877fdf 100644 --- a/src/aare/daq/server.py +++ b/src/aare/daq/server.py @@ -727,15 +727,6 @@ async def baton_request(token: str = Depends(oauth2_scheme)) -> dict: - Non-staff cannot request from staff """ data = auth.parse_token(token) - auth.check_jwt_ro(cfg, data) - - # holder = cfg.baton_holder - # if holder is not None and holder.is_staff and not data.staff: - # return { - # "error": True, - # "message": "Cannot request baton from staff. Please ask them directly." - # } - return auth.request_baton(cfg, data) @app.post("/baton/respond") @@ -747,7 +738,6 @@ async def baton_respond(accept: bool, token: str = Depends(oauth2_scheme)) -> di - accept=false: refuses the request """ data = auth.parse_token(token) - auth.check_jwt_ro(cfg, data) return auth.respond_to_baton_request(cfg, data, accept) diff --git a/src/aare/gui/main_window.py b/src/aare/gui/main_window.py index 3e2b8522..a237f3f7 100644 --- a/src/aare/gui/main_window.py +++ b/src/aare/gui/main_window.py @@ -75,6 +75,7 @@ class MainWindow(QMainWindow): self._cleanup_done = False self._waiting_for_baton_response: bool = False + self._baton_request_dialog: BatonRequestDialog | None = None # Tutorial manager (define tutorials after widgets exist) self.tutorial_manager = TutorialManager(self) @@ -289,11 +290,15 @@ class MainWindow(QMainWindow): self.daq = DAQWorker(base_url=self.__base_url, token=self.__token) - self.daq.baton_status_changed.connect(self._on_baton_status_changed) + self.daq.baton_status_changed.connect(self.status_bar.update_baton_status) 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) @@ -413,9 +418,12 @@ class MainWindow(QMainWindow): 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.daq.baton_status_changed.connect(self.status_bar.update_baton_status) - self.daq.baton_status_changed.connect(self.status_bar.update_baton_status) + 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) @@ -682,57 +690,63 @@ class MainWindow(QMainWindow): # ========== BATON DIALOG HANDLING ========== @Slot(dict) - def _on_baton_incoming_request(self, info: dict): - """Show dialog when someone requests our baton.""" - requester = info.get("requester", "Another user") - timeout = info.get("timeout", 30) + def _show_baton_request_dialog(self, payload: dict): + requester = str(payload.get("requester") or "Another user") + timeout = int(payload.get("timeout") or 30) - # Don't show multiple dialogs 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 + parent=self, ) - self._baton_request_dialog.accepted_signal.connect(self._on_baton_dialog_accepted) - self._baton_request_dialog.refused_signal.connect(self._on_baton_dialog_refused) + 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 _on_baton_dialog_accepted(self): - """User clicked Accept in baton dialog.""" + def _accept_baton_request(self): self.daq.respond_to_baton_request(True) - self._baton_request_dialog = None @Slot() - def _on_baton_dialog_refused(self): - """User clicked Refuse in baton dialog.""" + def _refuse_baton_request(self): self.daq.respond_to_baton_request(False) - self._baton_request_dialog = None @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) + self.alert_banner.show_message("Baton acquired!", False, auto_clear_ms=10000) logger.info("Baton acquired") + self._close_baton_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...") 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") + elif result.get("already_holder"): self._waiting_for_baton_response = False - # Don't show anything - user already has control logger.debug("Already baton holder") + elif result.get("error"): self._waiting_for_baton_response = False self.alert_banner.show_message(result.get("message", "Request failed"), True) @@ -741,58 +755,54 @@ class MainWindow(QMainWindow): @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.alert_banner.show_message("Control transferred", False) + 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.alert_banner.show_message("Request declined", False) - - @Slot(dict) - def _on_baton_request_received(self, info: dict): - """Incoming baton requests are now owned by StatusBar.""" - requester = info.get("requester", "Another user") - timeout = info.get("timeout", 30) - logger.info(f"Incoming baton request from {requester} ({timeout}s)") - - @Slot(BatonStatus) - def _on_baton_status_changed(self, status: BatonStatus): - """Handle baton status updates from SSE stream.""" - if self._waiting_for_baton_response: - pending = status.pending_request - if status.you_are_holder: - self._waiting_for_baton_response = False - self.alert_banner.show_message("Baton acquired!", False) - elif pending and pending.status == BatonRequestStatus.REFUSED: - self._waiting_for_baton_response = False - holder_name = pending.holder_username or "Current user" - self.alert_banner.show_message(f"Request declined by {holder_name}", True) - elif pending is None: - self._waiting_for_baton_response = False - self.alert_banner.clear_message() - elif status.you_have_pending_request: - import time - elapsed = time.time() - pending.created_at - remaining = max(0, int(pending.timeout_seconds - elapsed)) - if remaining > 0: - self.alert_banner.show_waiting( - f"Requesting control from {pending.holder_username or 'current user'}", - remaining - ) - else: - self.alert_banner.show_waiting("Processing timeout...", 0) + 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) + elif result.get("granted"): self._waiting_for_baton_response = False - self.alert_banner.show_message("Baton acquired!", False) + self.alert_banner.show_message("Baton acquired!", False, auto_clear_ms=10000) + self._close_baton_dialog() + elif result.get("queued"): self._waiting_for_baton_response = True self.alert_banner.show_waiting("Control transfer queued - waiting for beamline") + self._close_baton_dialog() + + 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() + + else: + logger.debug(f"replied with {result}") + self.alert_banner.clear_message() + + def _close_baton_dialog(self) -> None: + if self._baton_request_dialog is not None: + try: + self._baton_request_dialog.close() + finally: + self._baton_request_dialog = None def _restore_window_state(self) -> None: settings = QSettings() diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py index 53c4c212..7643aaf9 100644 --- a/src/aare/gui/threads/daq_worker.py +++ b/src/aare/gui/threads/daq_worker.py @@ -1169,8 +1169,16 @@ class DAQWorker(QObject): result = json.loads(response_data) if response_data else {} self.baton_request_result.emit(result) - if result.get("granted") or result.get("queued") or result.get("pending"): - self.check_baton_timeout() + 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)) @@ -1196,10 +1204,13 @@ class DAQWorker(QObject): 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.""" @@ -1227,8 +1238,9 @@ class DAQWorker(QObject): result = json.loads(response_data) if response_data else {} self.baton_timeout_checked.emit(result) - if result.get("granted") or result.get("queued"): + 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)) diff --git a/src/aare/gui/widgets/alert_banner.py b/src/aare/gui/widgets/alert_banner.py index c1ffe20b..6389a4e3 100644 --- a/src/aare/gui/widgets/alert_banner.py +++ b/src/aare/gui/widgets/alert_banner.py @@ -40,7 +40,7 @@ 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() @@ -81,7 +81,8 @@ 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._label.setText(decorated) self.setVisible(True) diff --git a/src/aare/gui/widgets/pgroup_dialog.py b/src/aare/gui/widgets/pgroup_dialog.py index 156ddb3a..88387dcb 100644 --- a/src/aare/gui/widgets/pgroup_dialog.py +++ b/src/aare/gui/widgets/pgroup_dialog.py @@ -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: {message}") + 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() \ No newline at end of file + return (self.combo.currentText() or "").strip() \ No newline at end of file diff --git a/src/aare/gui/widgets/status_bar.py b/src/aare/gui/widgets/status_bar.py index 7dec639a..6bf7d833 100644 --- a/src/aare/gui/widgets/status_bar.py +++ b/src/aare/gui/widgets/status_bar.py @@ -24,9 +24,11 @@ class StatusBar(QStatusBar): force_session = Signal() end_session = Signal() - request_baton = Signal() # New signal for baton request - cancel_baton_request = Signal() # Cancel pending request - release_baton = Signal() # Voluntarily release + 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) @@ -35,7 +37,6 @@ class StatusBar(QStatusBar): close_shutter = Signal() open_shutter = Signal() - def __init__(self, token: TokenData, parent=None): super().__init__(parent) self.__decoded_token = token @@ -224,17 +225,14 @@ class StatusBar(QStatusBar): incoming = bool(status and status.incoming_request) if incoming and not prev_incoming: - self._show_incoming_baton_request_dialog(status) + 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 _show_incoming_baton_request_dialog(self, status: BatonStatus) -> None: - if self._baton_request_dialog is not None and self._baton_request_dialog.isVisible(): - return - + def _emit_incoming_baton_request(self, status: BatonStatus) -> None: requester = "Another user" timeout = 30 if status.pending_request is not None: @@ -246,14 +244,15 @@ class StatusBar(QStatusBar): "timeout": timeout, }) - self._baton_request_dialog = BatonRequestDialog( - requester=requester, - timeout_seconds=timeout, - parent=self - ) - self._baton_request_dialog.accepted_signal.connect(self._on_baton_dialog_accepted) - self._baton_request_dialog.refused_signal.connect(self._on_baton_dialog_refused) - self._baton_request_dialog.show() + @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.""" @@ -426,95 +425,75 @@ class StatusBar(QStatusBar): menu.exec() - def _on_grab_clicked(self): - """Handle grab/request click - shows p-group dialog first if needed.""" - available_pgroups = self.__allowed_pgroups or [] + def _latest_pgroup(self, pgroups: list[str]) -> str | None: + if not pgroups: + return None - if not self.__is_staff and not available_pgroups: - QMessageBox.warning( - self, - "No P-Groups", - "You don't have any p-groups available. Cannot request control." - ) + 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 self.__is_staff and not available_pgroups: - self.request_baton.emit() + if len(pgroups) == 1: + self.set_pgroup.emit(pgroups[0]) return - if len(available_pgroups) == 1: - self.set_pgroup.emit(available_pgroups[0]) - self.request_baton.emit() - return - - self._show_pgroup_selection_for_baton(available_pgroups) - - def _show_pgroup_selection_for_baton(self, available_pgroups: list[str]): - """Show p-group selection dialog before requesting baton.""" - if self._pgroup_dialog_for_baton is not None and self._pgroup_dialog_for_baton.isVisible(): - return - - # Get current p-group if available - curr_pgroup = None + default_pgroup = self._latest_pgroup(pgroups) + curr = None if self.__status and self.__status.session: - curr_pgroup = self.__status.session.current_pgroup + curr = self.__status.session.current_pgroup or default_pgroup - # If staff, allow all p-groups - if self.__is_staff: - def _on_staff_loaded(lst: list): - try: - merged = {str(p).strip() for p in (lst or []) if p is not None and str(p).strip()} - # Add user's own p-groups too - for pg in available_pgroups: - if pg and str(pg).strip(): - merged.add(str(pg).strip()) - if not merged: - merged = set(available_pgroups) - self._show_pgroup_dialog_and_request_baton(curr_pgroup, sorted(merged)) - finally: - try: - self.staff_pgroups_loaded.disconnect(_on_staff_loaded) - except Exception as e: - logger.debug(f"Error disconnecting staff_pgroups_loaded: {e}") - - self.staff_pgroups_loaded.connect(_on_staff_loaded) - self.__list_staff_pgroups() - else: - # Non-staff: only show their own p-groups - self._show_pgroup_dialog_and_request_baton(curr_pgroup, available_pgroups) - - def _show_pgroup_dialog_and_request_baton(self, curr_pgroup: str | None, pgroups: list[str]): - """Show the dialog and request baton if accepted.""" self._pgroup_dialog_for_baton = PGroupDialog( - curr_pgroup=curr_pgroup, + curr_pgroup=curr, pgroups=pgroups, - parent=self + parent=self, ) if self._pgroup_dialog_for_baton.exec() == QDialog.DialogCode.Accepted: selected_pgroup = self._pgroup_dialog_for_baton.get_input() - if selected_pgroup and selected_pgroup.strip(): - # Validate the selection - if not self.__is_staff and selected_pgroup not in (self.__allowed_pgroups or []): - QMessageBox.warning( - self, - "Invalid P-Group", - f"P-group '{selected_pgroup}' is not in your allowed list.\n" - f"Please select from: {', '.join(self.__allowed_pgroups or [])}" - ) - else: - # Set the p-group and request baton - self.set_pgroup.emit(selected_pgroup) - self.request_baton.emit() - else: - QMessageBox.warning( - self, - "No P-Group Selected", - "You must select a p-group to request control." - ) + 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): + """Handle grab/request click - baton first, p-group after grant.""" + self.request_baton.emit() + def _on_release_clicked(self): """Handle release click.""" self.release_baton.emit() @@ -554,6 +533,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() @@ -565,6 +545,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) -- 2.54.0 From 303173ae82d0c356b12712a7093f4f8ad2b10a40 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Thu, 26 Mar 2026 17:29:17 +0100 Subject: [PATCH 54/56] Automation 2.0: WIP added endpoints to backend, frontend connected and panel made, now debugging --- src/aare/common/automation_queue_manager.py | 5 + src/aare/common/automation_workflow.py | 8 +- src/aare/daq/automation_api_router.py | 71 +++++- src/aare/daq/automation_runner.py | 136 ++++++---- src/aare/gui/gui.py | 4 +- src/aare/gui/main_window.py | 1 + src/aare/gui/panels/automation_panel.py | 260 +++++++++++++------- src/aare/gui/threads/daq_worker.py | 23 +- src/aare/gui/threads/workflow_sse_client.py | 97 +++++--- 9 files changed, 415 insertions(+), 190 deletions(-) diff --git a/src/aare/common/automation_queue_manager.py b/src/aare/common/automation_queue_manager.py index 54e4d2e2..4b63acdd 100644 --- a/src/aare/common/automation_queue_manager.py +++ b/src/aare/common/automation_queue_manager.py @@ -15,6 +15,7 @@ from aare.common.automation_models import ( WorkflowStateKind, ) +CURRENT_QUEUE_LIMIT = 576 def build_default_steps() -> list[WorkflowStepRecord]: return [ @@ -60,6 +61,10 @@ class WorkflowRedisManager: 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)}) diff --git a/src/aare/common/automation_workflow.py b/src/aare/common/automation_workflow.py index 98f26ac1..dacd48d0 100644 --- a/src/aare/common/automation_workflow.py +++ b/src/aare/common/automation_workflow.py @@ -258,7 +258,7 @@ class SimulatedMountHandler(StateHandler): def execute(self, context: WorkflowContext) -> StateResult: self.validate(context) import time - time.sleep(0.5) # Simulate some work + time.sleep(10) # Simulate some work context.current_state = WorkflowStateKind.MOUNT context.last_message = "[SIMULATION] Sample would be mounted" return StateResult( @@ -280,7 +280,7 @@ class SimulatedLoopCentreHandler(StateHandler): def execute(self, context: WorkflowContext) -> StateResult: self.validate(context) import time - time.sleep(0.3) + time.sleep(10) context.current_state = WorkflowStateKind.LOOP_CENTRE context.last_message = "[SIMULATION] Loop would be centred" return StateResult( @@ -302,7 +302,7 @@ class SimulatedRasterHandler(StateHandler): def execute(self, context: WorkflowContext) -> StateResult: self.validate(context) import time - time.sleep(0.4) + time.sleep(10) context.current_state = WorkflowStateKind.RASTER context.last_message = "[SIMULATION] Raster scan would be performed" return StateResult( @@ -324,7 +324,7 @@ class SimulatedDataCollectionHandler(StateHandler): def execute(self, context: WorkflowContext) -> StateResult: self.validate(context) import time - time.sleep(0.5) + time.sleep(10) context.current_state = WorkflowStateKind.DATA_COLLECTION context.last_message = "[SIMULATION] Data collection would be performed" return StateResult( diff --git a/src/aare/daq/automation_api_router.py b/src/aare/daq/automation_api_router.py index 1236ffc0..bc4f67eb 100644 --- a/src/aare/daq/automation_api_router.py +++ b/src/aare/daq/automation_api_router.py @@ -149,7 +149,11 @@ async def create_queue_item( recipe=request.recipe, ) - item = redis_mgr.create_item(item) + try: + item = redis_mgr.create_item(item) + except ValueError as e: + raise HTTPException(status_code=400, detail=str(e)) + return item @@ -199,6 +203,68 @@ async def delete_queue_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( @@ -304,8 +370,9 @@ async def request_abort( data = parse_token(token) check_jwt_rw(get_cfg(), data) + # Clear pause_requested when aborting control = redis_mgr.request_control( - {"abort_requested": True}, + {"abort_requested": True, "pause_requested": False, "resume_requested": False}, requested_by=data.sub, ) diff --git a/src/aare/daq/automation_runner.py b/src/aare/daq/automation_runner.py index 4626e601..8b672235 100644 --- a/src/aare/daq/automation_runner.py +++ b/src/aare/daq/automation_runner.py @@ -29,12 +29,12 @@ from aare.common.automation_workflow import ( class PersistentWorkflowRunner: """ - Orchestrates workflow execution with Redis persistence. + Orchestrates workflow execution with Redis persistence. - Combines: - - State handlers (from automation_workflow) - - Redis persistence (from automation_queue_manager) - """ + Combines: + - State handlers (from automation_workflow) + - Redis persistence (from automation_queue_manager) + """ def __init__( self, @@ -183,6 +183,10 @@ class PersistentWorkflowRunner: 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: @@ -203,17 +207,22 @@ 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.5, + 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 @@ -231,7 +240,7 @@ class AutomationLoop: self._on_step_complete = callback def start(self) -> None: - """Start the automation loop.""" + """Start the automation loop. Must be explicitly called.""" if self._task is not None and not self._task.done(): return # Already running @@ -248,6 +257,10 @@ class AutomationLoop: """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", @@ -258,7 +271,57 @@ class AutomationLoop: """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( @@ -266,51 +329,14 @@ class AutomationLoop: event_type="automation_error", message=f"Automation error: {e}", )) - # Brief pause on error before retrying await asyncio.sleep(2.0) await asyncio.sleep(self._poll_interval) async def _tick(self) -> None: - """Single iteration of the automation loop.""" - control = self._runner.check_control() + """Single iteration of the automation loop - process one step.""" runtime = self._redis.get_runtime() - - # Handle pause state - if control.pause_requested: - if runtime.running and not runtime.paused: - self._redis.patch_runtime({"paused": True}) - self._redis.append_event(WorkflowEvent( - beamline=self._runner.beamline, - item_id=runtime.current_item_id, - event_type="workflow_paused", - message="Workflow paused by user request", - )) - return - - # Handle resume - if control.resume_requested and runtime.paused: - self._redis.patch_runtime({"paused": False}) - self._redis.request_control( - {"resume_requested": False}, - requested_by="automation_loop", - ) - self._redis.append_event(WorkflowEvent( - beamline=self._runner.beamline, - item_id=runtime.current_item_id, - event_type="workflow_resumed", - message="Workflow resumed", - )) - - # Don't process if paused - if runtime.paused: - return - - # Handle abort - if control.abort_requested and runtime.current_item_id: - self._runner.complete_item(runtime.current_item_id, QueueItemStatus.ABORTED) - self._redis.clear_control() - return + control = self._runner.check_control() # If nothing running, try to start next item if not runtime.running or not runtime.current_item_id: @@ -319,8 +345,15 @@ class AutomationLoop: 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: @@ -330,8 +363,7 @@ class AutomationLoop: # Check if item is complete if item.current_step_index >= len(item.steps): self._runner.complete_item(item.item_id, QueueItemStatus.COMPLETED) - - # Check for next sample request + # Clear next_sample_requested if set if control.next_sample_requested: self._redis.request_control( {"next_sample_requested": False}, @@ -339,7 +371,7 @@ class AutomationLoop: ) return - # Handle skip request + # 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") @@ -366,15 +398,17 @@ class AutomationLoop: ) # 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: - step_name = item.steps[item.current_step_index - 1].kind # -1 because index advanced self._on_step_complete(item.item_id, step_name, result) - if __name__ == "__main__": client = redis.Redis(host="localhost", port=6379, db=0, decode_responses=True) redis_mgr = WorkflowRedisManager(client, beamline="x10sa") diff --git a/src/aare/gui/gui.py b/src/aare/gui/gui.py index 284ec734..d5957486 100644 --- a/src/aare/gui/gui.py +++ b/src/aare/gui/gui.py @@ -37,8 +37,8 @@ if __name__ == "__main__": 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_pred_zmq_addr = "tcp://x10sa-spark-01:9091" #"tcp://sls-gpu-003:9089"#"" + default_zmq_addr = "tcp://sls-gpu-003:9089"#"tcp://x10sa-spark-01:9091" #"tcp://x10sa-spark-01:9091" #"tcp://x10sa-pserv-01:9089" # + default_pred_zmq_addr = "tcp://sls-gpu-003:9089"#"""tcp://x10sa-spark-01:9091" # default_beamline_cam_addr = "axis-accc8eb02488.psi.ch" default_gonio_cam_addr = "axis-accc8ea5e463.psi.ch" default_gonio_camera_id = 1 diff --git a/src/aare/gui/main_window.py b/src/aare/gui/main_window.py index 67b75a75..be28070b 100644 --- a/src/aare/gui/main_window.py +++ b/src/aare/gui/main_window.py @@ -538,6 +538,7 @@ class MainWindow(QMainWindow): 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) diff --git a/src/aare/gui/panels/automation_panel.py b/src/aare/gui/panels/automation_panel.py index e53ea753..3fd82e5c 100644 --- a/src/aare/gui/panels/automation_panel.py +++ b/src/aare/gui/panels/automation_panel.py @@ -4,7 +4,7 @@ Workflow automation panel for queue management and step control. Displays: - Queue items with status (supports drag-drop from sample list) - Current step progress -- Control buttons (Next/Pause/Resume/Abort/Skip) +- Control buttons (Start Guided/Start Automation/Pause/Resume/Abort/Skip) - Live status from SSE """ @@ -87,6 +87,13 @@ class StepProgressWidget(QWidget): 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; " @@ -118,6 +125,7 @@ class DraggableQueueListWidget(QListWidget): 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) @@ -145,7 +153,6 @@ class DraggableQueueListWidget(QListWidget): """Accept drops from sample panels.""" mime = event.mimeData() if mime.hasText(): - # Check if it's sample data (JSON) try: text = mime.text() if text.startswith("{") or text.startswith("["): @@ -153,7 +160,6 @@ class DraggableQueueListWidget(QListWidget): return except Exception: pass - # Accept internal moves if event.source() == self: event.acceptProposedAction() return @@ -173,7 +179,6 @@ class DraggableQueueListWidget(QListWidget): if mime.hasText(): text = mime.text() try: - # Try to parse as SampleShortInfoList sample_list = SampleShortInfoList.model_validate_json(text) if sample_list.s: self.samples_dropped.emit(sample_list.s) @@ -183,7 +188,6 @@ class DraggableQueueListWidget(QListWidget): pass try: - # Try single sample sample = SampleShortInfo.model_validate_json(text) self.samples_dropped.emit([sample]) event.acceptProposedAction() @@ -191,14 +195,11 @@ class DraggableQueueListWidget(QListWidget): except Exception: pass - # Internal reorder if event.source() == self: - # Get drop position drop_row = self.indexAt(event.position().toPoint()).row() if drop_row < 0: drop_row = self.count() - # Get selected items selected = self.selectedItems() if selected and self._item_ids: for item in selected: @@ -251,8 +252,7 @@ class DraggableQueueListWidget(QListWidget): if action == delete_action: self.delete_requested.emit([self._item_ids[row]]) elif action == start_action: - # Will need to emit a signal for this - pass + 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: @@ -266,8 +266,8 @@ class WorkflowPanel(QWidget): Supports: - Drag-drop samples from TELL sample panel - Queue management (delete, reorder, clear, add all) - - Step-by-step guided mode - - Full automation mode + - Step-by-step guided mode (explicit start) + - Full automation mode (explicit start) """ # Signals for DAQ worker @@ -282,17 +282,20 @@ class WorkflowPanel(QWidget): request_pause = Signal() request_resume = Signal() request_abort = Signal() - request_skip = 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] = [] # Cache for "Add All" + 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: @@ -316,84 +319,104 @@ class WorkflowPanel(QWidget): 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("Controls") + controls_group = QGroupBox("Step Controls") controls_layout = QVBoxLayout(controls_group) - # Mode toggle - mode_layout = QHBoxLayout() - self._mode_label = QLabel("Mode:") - self._guided_btn = QPushButton("Guided Manual") - self._guided_btn.setCheckable(True) - self._guided_btn.setChecked(True) - self._auto_btn = QPushButton("Automation") - self._auto_btn.setCheckable(True) - - self._guided_btn.clicked.connect(self._on_guided_mode) - self._auto_btn.clicked.connect(self._on_automation_mode) - - mode_layout.addWidget(self._mode_label) - mode_layout.addWidget(self._guided_btn) - mode_layout.addWidget(self._auto_btn) - mode_layout.addStretch() - controls_layout.addLayout(mode_layout) - - # Step controls + # Step 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") + 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._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) + 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._pause_btn) - step_layout.addWidget(self._abort_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) - # Info label self._drop_hint = QLabel("💡 Drag samples from the Sample List to add them") self._drop_hint.setStyleSheet("color: #666; font-style: italic;") queue_layout.addWidget(self._drop_hint) - # Draggable queue list self._queue_list = DraggableQueueListWidget() self._queue_list.setMinimumHeight(150) - self._queue_list.itemDoubleClicked.connect(self._on_item_double_clicked) self._queue_list.samples_dropped.connect(self._on_samples_dropped) self._queue_list.delete_requested.connect(self._on_delete_requested) self._queue_list.item_reordered.connect(self._on_item_reordered) + 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 Samples") + 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 to the queue") + self._add_all_btn.setToolTip("Add all samples from the Sample List") - self._remove_selected_btn = QPushButton("🗑 Remove Selected") + self._remove_selected_btn = QPushButton("🗑 Remove") self._remove_selected_btn.clicked.connect(self._on_remove_selected) - self._clear_btn = QPushButton("✖ Clear Queue") + self._clear_btn = QPushButton("✖ Clear") self._clear_btn.clicked.connect(self._on_clear_queue) self._refresh_btn = QPushButton("🔄") @@ -413,20 +436,47 @@ class WorkflowPanel(QWidget): # Initial state self._update_button_states() - def _on_guided_mode(self) -> None: - """Switch to guided manual mode.""" - self._guided_btn.setChecked(True) - self._auto_btn.setChecked(False) - if self._automation_enabled: - self.request_stop_automation.emit() + 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_automation_mode(self) -> None: - """Switch to automation mode.""" - self._auto_btn.setChecked(True) - self._guided_btn.setChecked(False) - if not self._automation_enabled: - self.request_start_automation.emit() + 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: @@ -436,20 +486,11 @@ class WorkflowPanel(QWidget): else: self.request_pause.emit() - def _on_item_double_clicked(self, item: QListWidgetItem) -> None: - """Start processing double-clicked item.""" - idx = self._queue_list.row(item) - if 0 <= idx < len(self._queue_items): - queue_item = self._queue_items[idx] - if queue_item.status == QueueItemStatus.PENDING: - self.request_start_item.emit(queue_item.item_id) - def _on_samples_dropped(self, samples: list[SampleShortInfo]) -> None: """Handle samples dropped onto the queue.""" logger.info(f"Adding {len(samples)} samples to workflow queue") for sample in samples: self.request_add_sample.emit(sample) - # Refresh after a short delay to let server process QTimer.singleShot(500, self.request_queue_refresh.emit) def _on_delete_requested(self, item_ids: list[str]) -> None: @@ -485,14 +526,36 @@ class WorkflowPanel(QWidget): reply = QMessageBox.question( self, "Clear Queue", - f"Remove all {len(self._queue_items)} items from the 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: @@ -518,25 +581,34 @@ class WorkflowPanel(QWidget): """Update button enabled/disabled states based on current state.""" is_running = self._runtime is not None and self._runtime.running is_paused = self._runtime is not None and self._runtime.paused - is_automation = self._auto_btn.isChecked() + has_pending = any(i.status == QueueItemStatus.PENDING for i in self._queue_items) - # In automation mode, hide the Next button - self._next_btn.setVisible(not is_automation) + 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.setText("▶ Resume") self._pause_btn.setStyleSheet("background-color: #4CAF50; color: white;") else: - self._pause_btn.setText("Pause") + self._pause_btn.setText("⏸ Pause") self._pause_btn.setStyleSheet("background-color: #2196F3; color: white;") - self._pause_btn.setEnabled(is_running) - # Update drop hint visibility + # Drop hint has_items = len(self._queue_items) > 0 self._drop_hint.setVisible(not has_items) @@ -587,7 +659,8 @@ class WorkflowPanel(QWidget): "font-size: 14px; font-weight: bold; color: #FF9800;" ) else: - self._status_label.setText("▶️ Running") + 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;" ) @@ -601,6 +674,14 @@ class WorkflowPanel(QWidget): "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() @@ -622,8 +703,6 @@ class WorkflowPanel(QWidget): def update_automation_enabled(self, enabled: bool) -> None: """Update automation mode state.""" self._automation_enabled = enabled - self._auto_btn.setChecked(enabled) - self._guided_btn.setChecked(not enabled) self._update_button_states() @Slot(list) @@ -636,5 +715,16 @@ class WorkflowPanel(QWidget): """Handle workflow events from SSE stream.""" logger.debug(f"Workflow event: {event_type} - {message}") - if event_type in ("item_started", "item_completed", "step_completed", "step_skipped"): - self.request_queue_refresh.emit() \ No newline at end of file + 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() \ No newline at end of file diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py index df49f5f7..dffd18e0 100644 --- a/src/aare/gui/threads/daq_worker.py +++ b/src/aare/gui/threads/daq_worker.py @@ -1281,9 +1281,15 @@ class DAQWorker(QObject): try: response_data = self.handle_response(reply) data = json.loads(response_data) - from aare.common.automation_models import QueueItem + 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}") @@ -1334,15 +1340,20 @@ class DAQWorker(QObject): @Slot() def workflow_clear_queue(self): - """Clear all pending items from the queue.""" - # Load queue first, then delete all pending items + """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")) + request = QNetworkRequest(QUrl(f"{self.__base_url}/workflow/queue/clear")) request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode()) - reply = self.__net_manager.get(request) - reply.finished.connect(lambda: self._handle_clear_queue_response(reply)) + 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: diff --git a/src/aare/gui/threads/workflow_sse_client.py b/src/aare/gui/threads/workflow_sse_client.py index f1f4d857..ec264106 100644 --- a/src/aare/gui/threads/workflow_sse_client.py +++ b/src/aare/gui/threads/workflow_sse_client.py @@ -3,7 +3,7 @@ SSE client for workflow events. Connects to the /workflow/sse endpoint and emits signals for: - Runtime state changes -- Control state changes +- Control state changes - Individual workflow events """ @@ -24,10 +24,10 @@ 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 @@ -35,7 +35,7 @@ class WorkflowSSEClient(QObject): connected = Signal() disconnected = Signal() error = Signal(str) - + def __init__( self, base_url: str, @@ -43,85 +43,92 @@ class WorkflowSSEClient(QObject): parent: QObject | None = None, ): super().__init__(parent) - + self._base_url = base_url self._token = token self._manager = QNetworkAccessManager(self) self._reply: QNetworkReply | None = None self._buffer = "" - + # Reconnection self._reconnect_timer = QTimer(self) self._reconnect_timer.setInterval(5000) # 5 seconds self._reconnect_timer.timeout.connect(self.connect) self._should_reconnect = False - + def connect(self) -> None: """Start SSE connection.""" if self._reply is not None: return # Already connected - + self._should_reconnect = True - + url = QUrl(f"{self._base_url}/workflow/sse") - # Add token as query param for SSE (can't use headers easily) - url.setQuery(f"token={self._token}") - + request = QNetworkRequest(url) request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) request.setRawHeader(b"Accept", b"text/event-stream") request.setRawHeader(b"Cache-Control", b"no-cache") - + self._reply = self._manager.get(request) self._reply.readyRead.connect(self._on_data_ready) self._reply.finished.connect(self._on_finished) self._reply.errorOccurred.connect(self._on_error) - + self._reconnect_timer.stop() logger.debug("Workflow SSE: connecting...") - + def disconnect(self) -> None: """Stop SSE connection.""" self._should_reconnect = False self._reconnect_timer.stop() - + if self._reply is not None: - self._reply.abort() - self._reply.deleteLater() + 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 - - 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) - + + 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) @@ -134,25 +141,35 @@ class WorkflowSSEClient(QObject): self.workflow_event.emit(event) except Exception as e: logger.warning(f"Workflow SSE: failed to parse {event_type}: {e}") - + @Slot() def _on_finished(self) -> None: """Handle connection finished.""" if self._reply is not None: - self._reply.deleteLater() + 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 = self._reply.errorString() if self._reply else str(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) + self.error.emit(error_msg) \ No newline at end of file -- 2.54.0 From 5949f14599fcf261e7cab42d741927fb3ece5192 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Thu, 26 Mar 2026 17:31:13 +0100 Subject: [PATCH 55/56] aaredaq: continuing work on baton exchange procedure --- src/aare/common/auth_models.py | 3 +- src/aare/daq/auth.py | 89 ++++++------- src/aare/daq/config.py | 21 ++- src/aare/daq/server.py | 22 +++- src/aare/gui/main_window.py | 79 +++++++++++- src/aare/gui/threads/daq_worker.py | 9 +- src/aare/gui/widgets/alert_banner.py | 4 +- src/aare/gui/widgets/baton_request_dialog.py | 128 +++++++++++++++++++ src/aare/gui/widgets/status_bar.py | 18 ++- 9 files changed, 311 insertions(+), 62 deletions(-) diff --git a/src/aare/common/auth_models.py b/src/aare/common/auth_models.py index 6897751c..9930e7cd 100644 --- a/src/aare/common/auth_models.py +++ b/src/aare/common/auth_models.py @@ -48,4 +48,5 @@ class BatonStatus(BaseModel): queued_transfer: BatonTransferQueue | None = None you_are_holder: bool = False you_have_pending_request: bool = False - incoming_request: bool = False # True if YOU are being asked to give up baton \ No newline at end of file + incoming_request: bool = False + allow_non_staff_request: bool = False \ No newline at end of file diff --git a/src/aare/daq/auth.py b/src/aare/daq/auth.py index 5e946fd0..ac9b6601 100644 --- a/src/aare/daq/auth.py +++ b/src/aare/daq/auth.py @@ -117,7 +117,13 @@ 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: """ @@ -172,48 +178,43 @@ def resolve_baton_timeout_if_needed(cfg: BeamlineConfig) -> BatonStatus | None: return _finalize_expired_baton_request(cfg, pending) def get_baton_status(cfg: BeamlineConfig, data: TokenData) -> BatonStatus: - """Get full baton status for a user.""" + """ + 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 - if pending is not None and pending.status == BatonRequestStatus.PENDING: - elapsed = time.time() - pending.created_at - if elapsed >= pending.timeout_seconds: - resolve_baton_timeout_if_needed(cfg) - pending = cfg.pending_baton_request - - holder = cfg.baton_holder - queued = cfg.queued_baton_transfer - - you_are_holder = holder is not None and holder.session == data.session - you_have_pending_request = ( - pending is not None and - pending.requester_session == data.session and - pending.status == BatonRequestStatus.PENDING + is_requester = bool( + pending + and pending.status in (BatonRequestStatus.PENDING, BatonRequestStatus.REFUSED) + and pending.requester_session == data.session ) - incoming_request = ( - pending is not None and - holder is not None and - holder.session == data.session and - pending.status == BatonRequestStatus.PENDING + + 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=pending, - queued_transfer=queued, - you_are_holder=you_are_holder, - you_have_pending_request=you_have_pending_request, - incoming_request=incoming_request + 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: - """ - Request the baton (control) of the beamline. - - This should not depend on active p-group membership at the server route level. - Baton ownership rules are handled here. - """ resolve_baton_timeout_if_needed(cfg) session_state = cfg.session_state(data.session) @@ -224,7 +225,7 @@ def request_baton(cfg: BeamlineConfig, data: TokenData) -> dict: to_username=data.sub, to_is_staff=data.staff, to_pgroup=cfg.pgroup, - expiry_sec=SESSION_EXPIRE_SECONDS + expiry_sec=SESSION_EXPIRE_SECONDS, ) return {"granted": True, "message": "Baton acquired (beamline was vacant)"} @@ -233,11 +234,11 @@ def request_baton(cfg: BeamlineConfig, data: TokenData) -> dict: return {"already_holder": True, "message": "You already hold the baton"} holder = cfg.baton_holder - - if holder and holder.is_staff and not data.staff: + 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": "Cannot request baton from staff. Please ask them directly." + "message": "Requesting baton from staff is disabled by backend policy.", } if data.staff: @@ -248,11 +249,11 @@ def request_baton(cfg: BeamlineConfig, data: TokenData) -> dict: target_is_staff=data.staff, target_pgroup=cfg.pgroup, queued_at=time.time(), - reason="beamline_busy_staff_override" + reason="beamline_busy_staff_override", ) return { "queued": True, - "message": "Staff override queued - will transfer when beamline is available" + "message": "Staff override queued - will transfer when beamline is available", } cfg.execute_baton_transfer( @@ -260,7 +261,7 @@ def request_baton(cfg: BeamlineConfig, data: TokenData) -> dict: to_username=data.sub, to_is_staff=data.staff, to_pgroup=cfg.pgroup, - expiry_sec=SESSION_EXPIRE_SECONDS + expiry_sec=SESSION_EXPIRE_SECONDS, ) return {"granted": True, "override": True, "message": "Staff override - baton acquired"} @@ -275,11 +276,11 @@ def request_baton(cfg: BeamlineConfig, data: TokenData) -> dict: "pending": True, "existing": True, "remaining_seconds": max(0, remaining), - "message": f"Request already pending ({remaining:.0f}s remaining)" + "message": f"Request already pending ({remaining:.0f}s remaining)", } return { "error": True, - "message": "Another user already has a pending request" + "message": "Another user already has a pending request", } request = BatonRequest( @@ -291,7 +292,7 @@ def request_baton(cfg: BeamlineConfig, data: TokenData) -> dict: holder_session=holder.session if holder else None, created_at=time.time(), timeout_seconds=BATON_REQUEST_TIMEOUT_SECONDS, - status=BatonRequestStatus.PENDING + status=BatonRequestStatus.PENDING, ) cfg.set_pending_baton_request(request, timeout_sec=BATON_REQUEST_TIMEOUT_SECONDS) @@ -299,7 +300,7 @@ def request_baton(cfg: BeamlineConfig, data: TokenData) -> dict: "pending": True, "request_id": request.request_id, "timeout_seconds": BATON_REQUEST_TIMEOUT_SECONDS, - "message": f"Request sent to {holder.username if holder else 'current user'}" + "message": f"Request sent to {holder.username if holder else 'current holder'}", } def respond_to_baton_request(cfg: BeamlineConfig, data: TokenData, accept: bool) -> dict: diff --git a/src/aare/daq/config.py b/src/aare/daq/config.py index ba8f8328..827915d3 100644 --- a/src/aare/daq/config.py +++ b/src/aare/daq/config.py @@ -84,6 +84,20 @@ class BeamlineConfig: # 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")) @@ -751,4 +765,9 @@ class BeamlineConfig: self.__client.set(f"{self.__bl}:failed_mount_count", count) def increment_failed_mount_count(self) -> int: - return int(self.__client.incr(f"{self.__bl}:failed_mount_count")) \ No newline at end of file + 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 \ No newline at end of file diff --git a/src/aare/daq/server.py b/src/aare/daq/server.py index c0877fdf..9393c8c5 100644 --- a/src/aare/daq/server.py +++ b/src/aare/daq/server.py @@ -8,7 +8,7 @@ import cv2 import urllib3 import uvicorn -from aare.common.auth_models import BatonStatus +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 @@ -726,6 +726,7 @@ async def baton_request(token: str = Depends(oauth2_scheme)) -> dict: - 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) @@ -763,17 +764,24 @@ async def baton_check_timeout(token: str = Depends(oauth2_scheme)) -> dict: """ data = auth.parse_token(token) - resolved = auth.resolve_baton_timeout_if_needed(cfg) - if resolved is not None: - return resolved.model_dump() + 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 { @@ -783,6 +791,12 @@ async def baton_check_timeout(token: str = Depends(oauth2_scheme)) -> dict: 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.""" diff --git a/src/aare/gui/main_window.py b/src/aare/gui/main_window.py index a237f3f7..dcdfeca6 100644 --- a/src/aare/gui/main_window.py +++ b/src/aare/gui/main_window.py @@ -44,7 +44,7 @@ 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 +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 @@ -76,6 +76,7 @@ class MainWindow(QMainWindow): 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_manager = TutorialManager(self) @@ -291,6 +292,7 @@ class MainWindow(QMainWindow): 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) @@ -716,6 +718,26 @@ class MainWindow(QMainWindow): 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.""" @@ -723,7 +745,9 @@ class MainWindow(QMainWindow): self._waiting_for_baton_response = False self.alert_banner.show_message("Baton acquired!", False, auto_clear_ms=10000) logger.info("Baton acquired") - self._close_baton_dialog() + + # 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: @@ -735,6 +759,16 @@ class MainWindow(QMainWindow): 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") @@ -743,14 +777,23 @@ class MainWindow(QMainWindow): 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): @@ -777,33 +820,57 @@ class MainWindow(QMainWindow): 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) - self._close_baton_dialog() + + # 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") - self._close_baton_dialog() + + 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_dialog() + 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 self._baton_request_dialog is not 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") diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py index 7643aaf9..9c3a1de9 100644 --- a/src/aare/gui/threads/daq_worker.py +++ b/src/aare/gui/threads/daq_worker.py @@ -100,7 +100,6 @@ class DAQWorker(QObject): if self.__base_url is not None: self.start_face_detection_stream() self.start_baton_stream() - self._baton_timeout_timer.start() def get_last_error_payload(self) -> dict: return dict(self._last_error_payload or {}) @@ -1136,7 +1135,13 @@ class DAQWorker(QObject): if payload: status = BatonStatus.model_validate_json(payload) - # Detect incoming request + 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)): diff --git a/src/aare/gui/widgets/alert_banner.py b/src/aare/gui/widgets/alert_banner.py index 6389a4e3..3267a7f2 100644 --- a/src/aare/gui/widgets/alert_banner.py +++ b/src/aare/gui/widgets/alert_banner.py @@ -144,7 +144,9 @@ class AlertBanner(QFrame): self._countdown_remaining -= 1 if self._countdown_remaining <= 0: self._stop_countdown() - # Don't auto-clear - let the baton status update handle that + self.clear_message() + return + self._update_waiting_text() def _stop_countdown(self): diff --git a/src/aare/gui/widgets/baton_request_dialog.py b/src/aare/gui/widgets/baton_request_dialog.py index 90aab3b9..b23bca2e 100644 --- a/src/aare/gui/widgets/baton_request_dialog.py +++ b/src/aare/gui/widgets/baton_request_dialog.py @@ -216,5 +216,133 @@ class BatonRequestDialog(QDialog): """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 {self._target_user} 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) \ No newline at end of file diff --git a/src/aare/gui/widgets/status_bar.py b/src/aare/gui/widgets/status_bar.py index 6bf7d833..a6b546a9 100644 --- a/src/aare/gui/widgets/status_bar.py +++ b/src/aare/gui/widgets/status_bar.py @@ -219,10 +219,19 @@ class StatusBar(QStatusBar): 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) @@ -311,8 +320,11 @@ class StatusBar(QStatusBar): action_grab.triggered.connect(self._on_grab_clicked) elif holder_is_staff: # Non-staff cannot request from staff - action_grab = menu.addAction("Request (Staff has control)") - action_grab.setEnabled(False) + 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") @@ -343,7 +355,7 @@ class StatusBar(QStatusBar): 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}") -- 2.54.0 From 53296a166eb5536bd201b287bf0c2847dc648989 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Fri, 27 Mar 2026 11:55:22 +0100 Subject: [PATCH 56/56] Bug fixes following big merge and updated alert banenr --- src/aare/gui/gui.py | 2 +- src/aare/gui/main_window.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/src/aare/gui/gui.py b/src/aare/gui/gui.py index 9b4ffb63..712380ef 100644 --- a/src/aare/gui/gui.py +++ b/src/aare/gui/gui.py @@ -36,7 +36,7 @@ if __name__ == "__main__": default_gonio_cam_addr = "axis-accc8ed2972e.psi.ch" default_gonio_camera_id = 3 case MXBeamline.X10SA: - default_url = "http://mx-x10sa-queue-01.psi.ch:5210" #"http://127.0.0.1:5210" + 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" diff --git a/src/aare/gui/main_window.py b/src/aare/gui/main_window.py index 5c1e63bc..9ec75126 100644 --- a/src/aare/gui/main_window.py +++ b/src/aare/gui/main_window.py @@ -124,10 +124,10 @@ class MainWindow(QMainWindow): root_layout.setContentsMargins(0, 0, 0, 0) root_layout.setSpacing(0) - self.alert_banner = AlertBanner(parent=root_widget, error_timeout_ms=30000, recover_timeout_ms=5000) + self.alert_banner = AlertBanner(parent=root_widget) root_layout.addWidget(self.alert_banner) - self.alert_banner_secondary = AlertBanner(parent=root_widget, error_timeout_ms=10000, recover_timeout_ms=5000) + self.alert_banner_secondary = AlertBanner(parent=root_widget) root_layout.addWidget(self.alert_banner_secondary) top_widget = QWidget(parent=root_widget) -- 2.54.0