From 13bddd43c7c36624f6e93e2718d82439d39809ff Mon Sep 17 00:00:00 2001 From: David Perl Date: Mon, 6 Jul 2026 11:48:54 +0200 Subject: [PATCH] refactor: use aarecommon and remove common --- pyproject.toml | 2 +- scripts/demo_scan_ingest.py | 25 +- src/aare/common/__init__.py | 0 src/aare/common/aarelc_infer.py | 135 --- src/aare/common/aerotech_models.py | 155 --- src/aare/common/auth_models.py | 60 - src/aare/common/autofocus_tools.py | 60 - src/aare/common/automation_models.py | 78 -- src/aare/common/beamline.py | 122 -- src/aare/common/config/logging_dev.yaml | 45 - src/aare/common/config/logging_prod.yaml | 39 - src/aare/common/config/simulated.yaml | 7 - src/aare/common/config/x06da.yaml | 64 - src/aare/common/config/x06sa.yaml | 41 - src/aare/common/config/x10sa.yaml | 51 - src/aare/common/coordinate.py | 122 -- src/aare/common/diffraction_geometry.py | 82 -- src/aare/common/error_codes.py | 189 --- src/aare/common/exception_handler.py | 517 -------- src/aare/common/find_xtal.py | 626 ---------- src/aare/common/logger_config.py | 180 --- src/aare/common/logger_events.py | 158 --- src/aare/common/models.py | 689 ----------- src/aare/common/raster_grid.py | 140 --- src/aare/common/recurrence_watcher.py | 136 --- src/aare/common/rotation_scan.py | 25 - src/aare/common/sample_geometry.py | 97 -- src/aare/common/simulate_raster.py | 217 ---- src/aare/common/tell_models.py | 62 - src/aare/daq/aaredb.py | 306 +++-- src/aare/daq/auth.py | 139 ++- src/aare/daq/config.py | 262 ++-- src/aare/daq/daq.py | 1077 ++++++++++------- src/aare/daq/devices.py | 180 +-- src/aare/daq/mlbox.py | 280 +++-- .../daq/operations/common/ml_bounding_box.py | 104 +- src/aare/daq/operations/common/runtime.py | 6 +- .../operations/common/simulate_scan_result.py | 5 +- .../daq/operations/face_detection/service.py | 44 +- .../daq/operations/face_detection/utils.py | 126 +- .../daq/operations/loop_centering/analyzer.py | 16 +- .../daq/operations/loop_centering/models.py | 5 +- .../daq/operations/loop_centering/service.py | 48 +- src/aare/daq/operations/mounting/models.py | 7 +- src/aare/daq/operations/mounting/service.py | 28 +- src/aare/daq/operations/raster/models.py | 5 +- src/aare/daq/operations/raster/service.py | 218 +++- src/aare/daq/operations/rotation/service.py | 38 +- src/aare/daq/operations/screenshot/service.py | 22 +- src/aare/daq/server.py | 422 +++++-- src/aare/daq/server_exception_handler.py | 60 +- src/aare/daq/spreadsheetupdater.py | 51 +- src/aare/daq/tell_state_machine.py | 32 +- src/aare/daq/tellupdater.py | 70 +- src/aare/daq/workflows.py | 92 +- src/aare/devices/aerotech.py | 161 ++- src/aare/devices/bec_worker.py | 245 ++-- .../devices/experimental_hutch_shutter.py | 10 +- src/aare/devices/filter_transmission.py | 7 +- src/aare/devices/fluorimeter.py | 47 +- src/aare/devices/jfjoch.py | 161 ++- src/aare/devices/pss_state.py | 4 +- src/aare/devices/smargon.py | 36 +- src/aare/devices/tell_backend.py | 78 +- src/aare/devices/tell_client.py | 148 ++- src/aare/devices/zmq_client.py | 9 +- src/aare/gui/auth.py | 33 +- src/aare/gui/gui.py | 142 ++- src/aare/gui/main_window.py | 911 +++++++++----- src/aare/gui/models/bookmark.py | 3 +- src/aare/gui/models/sample_queue_model.py | 19 +- src/aare/gui/models/user_sample_model.py | 47 +- src/aare/gui/panels/LogPanel.py | 2 +- src/aare/gui/panels/abr_tweak_panel.py | 22 +- src/aare/gui/panels/automation_panel.py | 77 +- src/aare/gui/panels/beam_center_panel.py | 4 +- src/aare/gui/panels/beam_mark_panel.py | 8 +- src/aare/gui/panels/beam_size_panel.py | 4 +- .../gui/panels/beamline_recovery_panel.py | 28 +- src/aare/gui/panels/beamline_state_panel.py | 153 ++- .../gui/panels/compact_automation_panel.py | 14 +- .../gui/panels/data_collection_settings.py | 35 +- src/aare/gui/panels/developer_help_dialog.py | 6 +- src/aare/gui/panels/face_detection_panel.py | 30 +- src/aare/gui/panels/file_path_panel.py | 93 +- .../panels/fluorescence_data_collection.py | 18 +- src/aare/gui/panels/fluorescence_panel.py | 50 +- src/aare/gui/panels/illumination_panel.py | 12 +- src/aare/gui/panels/local_contact_panel.py | 214 +++- src/aare/gui/panels/manual_sample_panel.py | 17 +- src/aare/gui/panels/monochromator_panel.py | 4 +- src/aare/gui/panels/omega_panel.py | 4 +- src/aare/gui/panels/portrait_mode.py | 98 +- .../gui/panels/prediction_metrics_panel.py | 104 +- src/aare/gui/panels/raster_data_collection.py | 115 +- src/aare/gui/panels/reference_tools_panel.py | 77 +- .../gui/panels/rotation_data_collection.py | 157 ++- src/aare/gui/panels/samcam_panel.py | 82 +- src/aare/gui/panels/sample_queue_panel.py | 132 +- src/aare/gui/panels/scan_settings_panel.py | 75 +- src/aare/gui/panels/smargon_panel.py | 21 +- src/aare/gui/panels/smart_rotation_panel.py | 199 ++- src/aare/gui/panels/status_panel.py | 12 +- src/aare/gui/panels/target_stability_panel.py | 343 ++++-- src/aare/gui/panels/tell_sample_panel.py | 62 +- src/aare/gui/panels/zoom_panel.py | 13 +- .../gui/scan_logic/raster_grid_manager.py | 304 +++-- .../gui/scan_logic/rotation_scan_manager.py | 7 +- src/aare/gui/scan_logic/sample_mount_logic.py | 5 +- src/aare/gui/threads/camera_thread.py | 42 +- src/aare/gui/threads/daq_worker.py | 540 ++++++--- src/aare/gui/threads/jfjoch_viewer.py | 6 +- src/aare/gui/threads/prediction_subscriber.py | 41 +- src/aare/gui/widgets/alert_banner.py | 23 +- src/aare/gui/widgets/automation_progress.py | 17 +- src/aare/gui/widgets/busy_overlay.py | 19 +- src/aare/gui/widgets/camera_image.py | 286 +++-- .../widgets/local_contact_status_widget.py | 84 +- src/aare/gui/widgets/login.py | 49 +- src/aare/gui/widgets/message_box.py | 45 +- src/aare/gui/widgets/status_bar.py | 178 ++- tests/conftest.py | 37 +- tests/unit/common/test_aare_exception.py | 130 +- tests/unit/common/test_aerotech_models.py | 25 +- tests/unit/common/test_autofocus_tools.py | 35 +- tests/unit/common/test_beamline.py | 9 +- tests/unit/common/test_coordinate.py | 53 +- .../common/test_data_collection_parameters.py | 5 +- .../unit/common/test_diffraction_geometry.py | 21 +- tests/unit/common/test_error_codes.py | 102 +- tests/unit/common/test_exception_handler.py | 74 +- tests/unit/common/test_find_xtal.py | 66 +- tests/unit/common/test_logger_events.py | 18 +- tests/unit/common/test_mlbox_model.py | 5 +- tests/unit/common/test_models_extra.py | 21 +- tests/unit/common/test_raster_grid_common.py | 7 +- tests/unit/common/test_recurrence_watcher.py | 10 +- .../common/test_zoom_model_camera_settings.py | 6 +- .../test_face_detection_service.py | 113 +- .../test_loop_centering_analyzer.py | 72 +- .../test_loop_centering_service.py | 25 +- .../mounting/test_mounting_service.py | 40 +- .../screenshot/test_screenshot_service.py | 20 +- .../daq/operations/test_ml_raster_plan.py | 82 +- .../unit/daq/test_aare_daq_loop_centering.py | 22 +- tests/unit/daq/test_aaredb.py | 232 +++- tests/unit/daq/test_auth.py | 231 +++- .../test_automation_progress_state_manager.py | 5 +- tests/unit/daq/test_face_detection.py | 26 +- tests/unit/daq/test_gui_timeout.py | 22 +- tests/unit/daq/test_mlbox.py | 12 +- tests/unit/daq/test_mount.py | 78 +- tests/unit/daq/test_raster_logic.py | 13 +- tests/unit/daq/test_server.py | 126 +- .../unit/daq/test_server_exception_handler.py | 115 +- tests/unit/daq/test_spreadsheetupdater.py | 62 +- tests/unit/daq/test_tell_state_updater.py | 5 +- tests/unit/daq/test_workflows.py | 47 +- tests/unit/devices/test_aerotech.py | 53 +- .../test_experimental_hutch_shutter.py | 27 +- .../unit/devices/test_filter_transmission.py | 36 +- tests/unit/devices/test_fluorimeter.py | 36 +- tests/unit/devices/test_jfjoch.py | 89 +- tests/unit/devices/test_tell_client.py | 17 +- .../gui/test_automation_progress_parser.py | 25 +- tests/unit/gui/test_camera_thread.py | 18 +- .../unit/gui/test_data_collection_settings.py | 36 +- tests/unit/gui/test_gui_main.py | 4 +- tests/unit/gui/test_main_window.py | 228 ++-- tests/unit/gui/test_models.py | 60 +- tests/unit/gui/test_panels.py | 62 +- tests/unit/gui/test_threads_logic.py | 59 +- 172 files changed, 8267 insertions(+), 8214 deletions(-) delete mode 100644 src/aare/common/__init__.py delete mode 100644 src/aare/common/aarelc_infer.py delete mode 100644 src/aare/common/aerotech_models.py delete mode 100644 src/aare/common/auth_models.py delete mode 100644 src/aare/common/autofocus_tools.py delete mode 100644 src/aare/common/automation_models.py delete mode 100644 src/aare/common/beamline.py delete mode 100644 src/aare/common/config/logging_dev.yaml delete mode 100644 src/aare/common/config/logging_prod.yaml delete mode 100644 src/aare/common/config/simulated.yaml delete mode 100644 src/aare/common/config/x06da.yaml delete mode 100644 src/aare/common/config/x06sa.yaml delete mode 100644 src/aare/common/config/x10sa.yaml delete mode 100644 src/aare/common/coordinate.py delete mode 100644 src/aare/common/diffraction_geometry.py delete mode 100644 src/aare/common/error_codes.py delete mode 100644 src/aare/common/exception_handler.py delete mode 100644 src/aare/common/find_xtal.py delete mode 100644 src/aare/common/logger_config.py delete mode 100644 src/aare/common/logger_events.py delete mode 100644 src/aare/common/models.py delete mode 100644 src/aare/common/raster_grid.py delete mode 100644 src/aare/common/recurrence_watcher.py delete mode 100644 src/aare/common/rotation_scan.py delete mode 100644 src/aare/common/sample_geometry.py delete mode 100644 src/aare/common/simulate_raster.py delete mode 100644 src/aare/common/tell_models.py diff --git a/pyproject.toml b/pyproject.toml index 5f433814..a48bcaf1 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -6,7 +6,7 @@ readme = "README.md" requires-python = ">=3.11" dependencies = [ "uv", - "aarecommon>=0.1", + "aarecommon>=0.1.1", "pydantic==2.11.4", "numpy==2.2.5", "jfjoch_client==1.0.0rc146", diff --git a/scripts/demo_scan_ingest.py b/scripts/demo_scan_ingest.py index 7ad4ff9e..2ad114cd 100644 --- a/scripts/demo_scan_ingest.py +++ b/scripts/demo_scan_ingest.py @@ -1,16 +1,23 @@ -from aare.daq.daq import AareDAQ -from aare.devices.jfjoch import JFJochWrapper -from aare.common.beamline import MXBeamline, mx_beamline -from aare.daq.aaredb import (AareWrapper) -from aare.daq.config import BeamlineConfig -from aare.common.diffraction_geometry import DiffractionGeometry -from aare.daq.devices import BeamlineDevices +from aarecommon.config.beamline import mx_beamline +from aarecommon.math.diffraction_geometry import DiffractionGeometry +from aarecommon.models.beamline import MXBeamline -bl=mx_beamline() +from aare.daq.aaredb import AareWrapper +from aare.daq.config import BeamlineConfig +from aare.daq.daq import AareDAQ +from aare.daq.devices import BeamlineDevices +from aare.devices.jfjoch import JFJochWrapper + +bl = mx_beamline() j = JFJochWrapper(bl) a = AareWrapper(bl) c = BeamlineConfig(bl) d = BeamlineDevices(bl) daq = AareDAQ(bl=bl, cfg=c) result = j.wait_till_done(60) -a.ingest_scan(sample=c.current_sample, result=result, geom= daq.sample_geometry,beam_mark_pxl=c.get_beam_mark(d.zoom)) \ No newline at end of file +a.ingest_scan( + sample=c.current_sample, + result=result, + geom=daq.sample_geometry, + beam_mark_pxl=c.get_beam_mark(d.zoom), +) diff --git a/src/aare/common/__init__.py b/src/aare/common/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/src/aare/common/aarelc_infer.py b/src/aare/common/aarelc_infer.py deleted file mode 100644 index 454aa9da..00000000 --- a/src/aare/common/aarelc_infer.py +++ /dev/null @@ -1,135 +0,0 @@ -import cv2 -import numpy as np -from aarelcinfer_client import Client, AuthenticatedClient -from aarelcinfer_client.api import config, predictions, beam -from aarelcinfer_client.models import RuntimeConfigPatchModel, LatestPredictionModel -import io -from PIL import Image - -#from aare.common.logger_config import setup_logger -from aare.common.beamline import MXBeamline, mx_beamline, cfg_get - - -class AareLCInferWrapper: - def __init__( - self, - bl: MXBeamline, - secret: str = "1s3ng@rd", - ): - if bl == MXBeamline.X10SA or bl == MXBeamline.X06DA: - host = cfg_get("daq.hardware.aarelc_url") - if host is None: - raise Exception("AareLCInferWrapper: AareLC URL not configured") - elif bl == MXBeamline.X06SA: - raise NotImplementedError(f"AareLCInferWrapper not implemented for {bl}") - elif bl == MXBeamline.SIMULATED: - raise NotImplementedError(f"AareLCInferWrapper not implemented for {bl}") - else: - raise Exception(f"Unknown beamline {bl}") - - self.client = AuthenticatedClient(base_url=host, api_key=secret) - self.client.headers["X-API-Key"] = secret - self._host = host - self._api_config = config - self._api_predictions = predictions - self._api_beam = beam - - def get_config(self) -> config.ConfigSnapshotResponse: - return self._api_config.get_config(self.client) - - def update_config(self, patch: RuntimeConfigPatchModel) -> config.ConfigUpdateResponse: - return self._api_config.update_config(self.client, patch) - - def get_latest_prediction(self) -> LatestPredictionModel: - return self._api_predictions.get_latest_prediction(self.client) - - def get_latest_frame_png(self): - return self._api_predictions.get_latest_frame_png(self.client) - - def get_latest_prediction_bundle(self): - return self._api_predictions.get_latest_prediction_bundle(self.client) - - def send_samcam_details(self, beam_mark, beam_dimensions): - return self._api_beam.set_beam_mark(beam_mark, beam_dimensions) - - -if __name__ == "__main__": - wrapper = AareLCInferWrapper(bl=mx_beamline()) - - try: - print("Fetching config...") - config = wrapper.get_config() - print(f"Config: {config}") - print("Updating config...") - patch = RuntimeConfigPatchModel( - conf=0.4 - # infer_scale: float | None = Field(default=None, gt=0.0, le=1.0) - # skip: int | None = Field(default=None, ge=0) - # max_fps: float | None = Field(default=None, ge=0.0) - # device: str | None = None - # publish_enabled: bool | None = None - # compute_target_point: bool | None = None - # focus_enabled: bool | None = None - # focus_epics_enabled: bool | None = None - # focus_pv: str | None = None - # focus_every: int | None = Field(default=None, ge=1) - # focus_scale: float | None = Field(default=None, gt=0.0, le=1.0) - # focus_pv_min_period_ms: float | None = Field(default=None, ge=0.0) - # tracker: Literal["none", "bytetrack", "botsort"] | None = None - # pt: str | None = None - # engine: str | None = None - # task: Literal["auto", "detect", "segment"] | None = None - ) - response = wrapper.update_config(patch) - print("Config updated!") - print(f"Model reloaded: {getattr(response, 'model_reloaded', False)}") - except Exception as e: - print(f"Error updating config: {e}") - - try: - print("Fetching bundle...") - bundle = wrapper.get_latest_prediction_bundle() - - jpeg_bytes = bundle.image_jpeg - - # The client seems to already return the validated model as 'metadata' - prediction = bundle.metadata - - # Safety check: if it's still a dict for some reason, validate it; otherwise use as is - if isinstance(prediction, dict): - prediction = LatestPredictionModel.model_validate(prediction) - - print(f"Received bundle. Image: {len(jpeg_bytes)} bytes. Detections: {len(prediction.boxes)}") - - # Decode Image - img = Image.open(io.BytesIO(jpeg_bytes)).convert("RGB") - img = cv2.cvtColor(np.array(img), cv2.COLOR_RGB2BGR) - - # Draw Overlay using the model attributes - print(prediction.target_point.x, prediction.target_point.y) - print("boxes: ", prediction.boxes) - for det in prediction.boxes: - x1, y1 = int(det.x1), int(det.y1) - x2, y2 = int(det.x2), int(det.y2) - - cv2.rectangle(img, (x1, y1), (x2, y2), (0, 255, 0), 2) - cv2.putText( - img, f"{det.label} {det.conf:.2f}", (x1, max(20, y1 - 8)), - cv2.FONT_HERSHEY_SIMPLEX, 0.6, (0, 255, 0), 2 - ) - - if det.poly: - # Convert to numpy and shift coordinates from relative to absolute - pts = np.array(det.poly, dtype=np.int32) - pts[:, 0] += x1 # Shift X - pts[:, 1] += y1 # Shift Y - - pts = pts.reshape((-1, 1, 2)) - cv2.polylines(img, [pts], isClosed=True, color=(0, 0, 255), thickness=2) - - # cv2.imshow("Latest Prediction Bundle", img) - # cv2.waitKey(0) - # cv2.destroyAllWindows() - - except Exception as e: - print(f"Error processing prediction bundle: {e}") \ No newline at end of file diff --git a/src/aare/common/aerotech_models.py b/src/aare/common/aerotech_models.py deleted file mode 100644 index 786ea1c3..00000000 --- a/src/aare/common/aerotech_models.py +++ /dev/null @@ -1,155 +0,0 @@ -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 diff --git a/src/aare/common/auth_models.py b/src/aare/common/auth_models.py deleted file mode 100644 index 445d137b..00000000 --- a/src/aare/common/auth_models.py +++ /dev/null @@ -1,60 +0,0 @@ -from enum import Enum -from pydantic import BaseModel -from datetime import datetime - -class BatonRequestStatus(Enum): - PENDING = "pending" - ACCEPTED = "accepted" - REFUSED = "refused" - TIMEOUT = "timeout" - CANCELLED = "cancelled" - - -class BatonHolderInfo(BaseModel): - """Information about the current baton holder.""" - username: str - session: int - is_staff: bool - pgroup: str | None = None - - -class BatonRequest(BaseModel): - """A request from one user to take the baton from another.""" - request_id: str - requester_username: str - requester_session: int - requester_is_staff: bool - holder_username: str | None = None - holder_session: int | None = None - created_at: float # Unix timestamp - timeout_seconds: int = 30 - status: BatonRequestStatus = BatonRequestStatus.PENDING - - -class BatonTransferQueue(BaseModel): - """Queued baton transfer waiting for beamline to be available.""" - target_session: int - target_username: str - target_is_staff: bool - target_pgroup: str | None = None - queued_at: float # Unix timestamp - reason: str = "beamline_busy" - - -class BatonStatus(BaseModel): - """Full baton status for GUI display.""" - holder: BatonHolderInfo | None = None - pending_request: BatonRequest | None = None - queued_transfer: BatonTransferQueue | None = None - you_are_holder: bool = False - you_have_pending_request: bool = False - incoming_request: bool = False - allow_non_staff_request: bool = False - - -def get_user(): - import os, getpass - try: - return os.getlogin() - except OSError: - return os.environ.get("USER") or getpass.getuser() \ No newline at end of file diff --git a/src/aare/common/autofocus_tools.py b/src/aare/common/autofocus_tools.py deleted file mode 100644 index 06b05b72..00000000 --- a/src/aare/common/autofocus_tools.py +++ /dev/null @@ -1,60 +0,0 @@ -import cv2 -import numpy as np - -def focus_measure_edges(gray: np.ndarray, mask: np.ndarray | None = None, verbose: bool = False) -> float: - # mild denoise (optional but usually stabilizes the curve) - gray = cv2.GaussianBlur(gray, (3, 3), 0) - - gx = cv2.Scharr(gray, cv2.CV_64F, 1, 0) - gy = cv2.Scharr(gray, cv2.CV_64F, 0, 1) - g2 = gx * gx + gy * gy - - roi = g2[mask] if mask is not None else g2.reshape(-1) - if roi.size == 0: - return 0.0 - - # threshold relative to median -> knocks out noise floor - t = float(np.median(roi) * 3.0) - strong = roi[roi > t] - - if strong.size == 0: - return 0.0 - if verbose: - mask_sum = mask.sum() if mask is not None else roi.size - print(f"mask pixels: {mask_sum}, " - f"focus={strong.mean():.2f}" - f"strong_size={strong.size}") - - return float(strong.mean()) # higher = sharper - -def focus_measure_blob_size(gray: np.ndarray, mask: np.ndarray | None = None) -> float: - """ - Measures sharpness for a single bright blob. - Higher = sharper (smaller blob). - """ - g = gray.astype(np.float64) - - if mask is not None: - g = np.where(mask, g, 0.0) - - # Background subtraction is crucial for blob metrics - # Use a large-ish blur as background estimate (tune ksize to your scale) - bg = cv2.GaussianBlur(g, (0, 0), sigmaX=10.0, sigmaY=10.0) - s = g - bg - s[s < 0] = 0.0 - - total = float(s.sum()) - if total <= 0: - return 0.0 - - h, w = s.shape - y, x = np.mgrid[0:h, 0:w] - - cx = float((s * x).sum() / total) - cy = float((s * y).sum() / total) - - # intensity-weighted second central moment (variance) - var = float((s * ((x - cx) ** 2 + (y - cy) ** 2)).sum() / total) - - # smaller var => sharper, so invert - return float(1.0 / (var + 1e-9)) \ No newline at end of file diff --git a/src/aare/common/automation_models.py b/src/aare/common/automation_models.py deleted file mode 100644 index f87e047a..00000000 --- a/src/aare/common/automation_models.py +++ /dev/null @@ -1,78 +0,0 @@ -from __future__ import annotations - -from dataclasses import dataclass, field -from datetime import datetime, timezone -from enum import Enum -from typing import Any, Literal - - -class WorkflowStateKind(str, Enum): - MOUNT = "mount" - LOOP_CENTRE = "loop_centre" - RASTER = "raster" - DATA_COLLECTION = "data_collection" - FINAL = "Paused/Finished" - - -class StepStatus(str, Enum): - PENDING = "pending" - RUNNING = "running" - SUCCESS = "success" - FAILED = "failed" - SKIPPED = "skipped" - PAUSED = "paused" - - -@dataclass -class StepState: - step: WorkflowStateKind - status: StepStatus = StepStatus.PENDING - message: str = "" - started_at: float | None = None - completed_at: float | None = None - error_code: str | None = None - - -@dataclass -class AutomationProgress: - current_step: str | None - steps: list[StepState] = field(default_factory=list) - events: list[LogEvent] = field(default_factory=list) - finished: bool = False - success: bool | None = None - samples_in_queue: int = 0 - avg_time_per_sample: float = 0.0 - current_sample_name: str = "" - - def append_event( - self, - *, - level: Literal["INFO", "WARNING", "ERROR"], - code: str, - message: str, - exception_class: str | None = None, - sample_id: int | None = None, - context: dict[str, Any] | None = None, - ) -> None: - self.events.append( - LogEvent( - ts=datetime.now(timezone.utc), - level=level, - code=code, - exception_class=exception_class, - message=message, - sample_id=sample_id, - context=dict(context or {}), - ) - ) - - -@dataclass -class LogEvent: - ts: datetime - level: Literal["INFO", "WARNING", "ERROR"] - code: str - exception_class: str | None - message: str - sample_id: int | None - context: dict[str, Any] = field(default_factory=dict) \ No newline at end of file diff --git a/src/aare/common/beamline.py b/src/aare/common/beamline.py deleted file mode 100644 index 03e377e7..00000000 --- a/src/aare/common/beamline.py +++ /dev/null @@ -1,122 +0,0 @@ -import os -from enum import Enum -from pathlib import Path -import yaml -from typing import Any, Dict - - -class MXBeamline(Enum): - X06SA = "X06SA" - X10SA = "X10SA" - X06DA = "X06DA" - SIMULATED = "SIMULATED" - - -class BeamlineYAMLConfig: - def __init__(self): - self.beamline: MXBeamline = mx_beamline() - self.config = self._load_config() - - def _load_config(self) -> Dict[str, Any]: - config_dir = Path(__file__).parent / "config" - yaml_file = config_dir / f"{self.beamline.value.lower()}.yaml" - - if not yaml_file.exists(): - raise FileNotFoundError(f"Config file not found: {yaml_file}") - - with open(yaml_file, "r") as f: - return yaml.safe_load(f) - - def get(self, key: str, default: Any = None) -> Any: - return self.config.get(key, default) - - -def mx_beamline() -> MXBeamline: - name = os.getenv("BEAMLINE") - if name is None: - return MXBeamline.SIMULATED - name=name.strip().upper() - return MXBeamline[name] if name in MXBeamline.__members__ else MXBeamline.SIMULATED - -def get_jfjoch_url(bl: MXBeamline) -> str: - """Centralized URL resolution for JFJoch services.""" - match bl: - case MXBeamline.X10SA: - return cfg_get("daq.hardware.jfjoch_url", "http://sls-gpu-002:8080") - case MXBeamline.X06DA: - return cfg_get("daq.hardware.jfjoch_url", "http://sls-gpu-001:8080") - case MXBeamline.SIMULATED: - return cfg_get("daq.hardware.jfjoch_url", "http://localhost:8080") - case MXBeamline.X06SA: - raise NotImplementedError("X06SA beamline not supported yet") - case _: - raise ValueError(f"unknown beamline {bl}") - -def get_beamline_config() -> BeamlineYAMLConfig: - beamline_config = BeamlineYAMLConfig() - return beamline_config - -def cfg_get(path: str, default: Any = None) -> Any: - """ - Hierarchical getter for YAML entries using dotted paths, e.g.: - cfg_get("endpoints.smargon_base") - cfg_get("epics.pv_prefix") - """ - beamline_config = get_beamline_config() - node = beamline_config.config - for part in path.split("."): - if not isinstance(node, dict) or part not in node: - return default - node = node[part] - return node - -def jfjoch_url() -> str | None: - return cfg_get("shared.jfjoch.jfjoch_url") - - -def smargon_url() -> str | None: - return cfg_get("daq.hardware.smargon_url") - - -def aerotech_url() -> str | None: - return cfg_get("daq.hardware.aerotech_url") - - -def tell_url() -> str | None: - return cfg_get("daq.hardware.tell_url") - - -def bec_host() -> str | None: - return cfg_get("daq.hardware.bec_url") - - -def redis_host() -> str | None: - return cfg_get("daq.hardware.redis_url") - - -def gui_sample_camera_zmq() -> str | None: - return cfg_get("gui.cameras.sample_camera_zmq_url") - - -def gui_prediction_zmq() -> str | None: - return cfg_get("gui.cameras.prediction_zmq_url") - - -def gui_beamline_camera_addr() -> str | None: - return cfg_get("gui.cameras.beamline_camera_url") - - -def gui_gonio_camera_addr() -> str | None: - return cfg_get("gui.cameras.gonio_camera_url") - - -def gui_gonio_camera_id() -> int | None: - return cfg_get("gui.cameras.gonio_camera_id") - - -def daq_base_url() -> str | None: - return cfg_get("gui.daq.daq_url") - -# Global instance for easy import -if __name__ == "__main__": - beamline_config = BeamlineYAMLConfig() \ No newline at end of file diff --git a/src/aare/common/config/logging_dev.yaml b/src/aare/common/config/logging_dev.yaml deleted file mode 100644 index 571696b9..00000000 --- a/src/aare/common/config/logging_dev.yaml +++ /dev/null @@ -1,45 +0,0 @@ -version: 1 -disable_existing_loggers: False - -formatters: - detailed: - format: "%(asctime)s - %(name)s - %(levelname)s - %(message)s" - -handlers: - console: - class: logging.StreamHandler - level: DEBUG - formatter: detailed - stream: ext://sys.stdout - - app_file: - class: logging.handlers.RotatingFileHandler - level: DEBUG - formatter: detailed - filename: app.log - maxBytes: 1440000 #not to overflow a diskette - backupCount: 10 - encoding: utf8 - - error_file: - class: logging.handlers.RotatingFileHandler - level: ERROR - formatter: detailed - filename: errors.log - maxBytes: 1440000 - backupCount: 5 - encoding: utf8 - -loggers: - aareDAQ: - level: DEBUG - handlers: [console, app_file, error_file] - propagate: yes - aareGUI: - level: DEBUG - handlers: [ console, app_file, error_file ] - propagate: yes - -root: - level: DEBUG - handlers: [console] \ No newline at end of file diff --git a/src/aare/common/config/logging_prod.yaml b/src/aare/common/config/logging_prod.yaml deleted file mode 100644 index d9e42fd7..00000000 --- a/src/aare/common/config/logging_prod.yaml +++ /dev/null @@ -1,39 +0,0 @@ -version: 1 -disable_existing_loggers: False - -formatters: - simple: - format: "%(asctime)s - %(levelname)s - %(message)s" - -handlers: - app_file: - class: logging.handlers.RotatingFileHandler - level: INFO - formatter: simple - filename: logs/app.log - maxBytes: 1000000 - backupCount: 10 - encoding: utf8 - - error_file: - class: logging.handlers.RotatingFileHandler - level: ERROR - formatter: simple - filename: logs/errors.log - maxBytes: 500000 - backupCount: 5 - encoding: utf8 - -loggers: - aareDAQ: - level: DEBUG - handlers: [console, app_file, error_file] - propagate: no - aareGUI: - level: DEBUG - handlers: [ console, app_file, error_file ] - propagate: n - -root: - level: WARNING - handlers: [] diff --git a/src/aare/common/config/simulated.yaml b/src/aare/common/config/simulated.yaml deleted file mode 100644 index ed36a091..00000000 --- a/src/aare/common/config/simulated.yaml +++ /dev/null @@ -1,7 +0,0 @@ -beamline: "SIMULATED" -smargon_base: "simulated" -aerotech_base: "simulated" -zmq_camera_url: "simulated" -jfjoch_url: "simulated" -tell_url: "simulated" -redis_host: "localhost" \ No newline at end of file diff --git a/src/aare/common/config/x06da.yaml b/src/aare/common/config/x06da.yaml deleted file mode 100644 index 43aaa2e9..00000000 --- a/src/aare/common/config/x06da.yaml +++ /dev/null @@ -1,64 +0,0 @@ -beamline_id: "X06DA" - -gui: - cameras: - sample_camera_zmq_url: "tcp://sls-gpu-003:9093" - prediction_zmq_url: "tcp://sls-gpu-003:9093" - beamline_camera_url: "x06da-axis-1.psi.ch" - secondary_beamline_camera_url: "axis-accc8e9a2995.psi.ch" - gonio_camera_url: "axis-server-es.psi.ch" - gonio_camera_id: "1" - - daq: - daq_url: "https://mx-x06da-queue-01.psi.ch" - cert_path: "/sls/x06da/misc/.cert/6d.crt" - -shared: - jfjoch: - jfjoch_url: "http://sls-gpu-001:8080" - -daq: - hardware: - smargon_url: "http://x06da-smargopolo.psi.ch:3000" - smargon_frontend_url: "http://x06da-smargopolo.psi.ch:8080" - aerotech_url: "http://mx-x06da-queue-01.psi.ch:5234" # Adjust if needed - tell_url: "http://x06da-tell.psi.ch:22222" - bec_url: "x06da-bec-001.psi.ch" - redis_url: "x06da-redis.psi.ch" - aarelc_url: "http://sls-gpu-003:9094" - dtz_safe_position: null - lens_magnification: 10 #or 5 currently - default_detector_distance_minimum: 86 - default_detector_distance_maximum: 900 - zoom_min: 1 # camera zoom travel limits (used to clamp auto-center zoom-to-fit) - zoom_max: 1000 - sam_cam: - #TODO wire this - white_balance_ratio = 1.233 - - db: - aaredb_url: "https://mx-aaredb-dmz-01.psi.ch/dispatcher" - - data_collection_settings: - default_raster_scan_settings: - exp_time_s: 0.02 - transmission: 1.0 - dtz: 150.0 - - default_rotation_settings: - exp_time_s: 0.02 - transmission: 1.0 - dtz: 150.0 - start_omega_deg: 0.0 - increment_omega_deg: 0.2 - steps: 1800 - - detector_limit_modifier: 2.0 - maximum_flux: 4e11 - - auto_raster: - grid_padding_fraction_x: 0.15 # pad the 1st grid scan by this fraction of its size per side in x (min 1 cell) - grid_padding_fraction_y: 0.15 # ... and the TOP in y; shifts smargon_top_left outward (cells before cell 0) - grid_padding_fraction_y_bottom: 0.30 # pad the BOTTOM of the grid (far end of n_y) more; defaults to grid_padding_fraction_y - include_crystal: false # extend the grid to cover crystals outside the loop box - line_scan_y_padding_fraction: 0.15 # pad the 2nd-stage vertical line scan height by this per side (10-20%) diff --git a/src/aare/common/config/x06sa.yaml b/src/aare/common/config/x06sa.yaml deleted file mode 100644 index 72d52b13..00000000 --- a/src/aare/common/config/x06sa.yaml +++ /dev/null @@ -1,41 +0,0 @@ -beamline_id: "X06SA" -gui: - cameras: - sample_camera_zmq_url: "tcp://x06sa-pserv-01:9089" - prediction_zmq_url: "" - beamline_camera_url: "" - gonio_camera_url: "" - gonio_camera_id: "" - - daq: - daq_url: "https://mx-x06sa-queue-01.psi.ch" - cert_path: "/sls/x06sa/misc/.cert/6s.crt" - -shared: - jfjoch: - jfjoch_url: "" - -daq: - hardware: - smargon_url: "http://x06sa-smargopolo.psi.ch:3000" - aerotech_url: "http://x06sa-queue-01.psi.ch:5234" # Adjust if needed - tell_url: "" - bec_url: "x06sa-bec-001.psi.ch" - redis_url: "" - - db: - aaredb_url: "https://mx-aaredb-dmz-01.psi.ch/dispatcher" - - data_collection_settings: - default_raster_scan_settings: - exp_time_s: 0.02 - transmission: 1.0 - dtz: 150.0 - - default_rotation_settings: - exp_time_s: 0.02 - transmission: 1.0 - dtz: 150.0 - start_omega_deg: 0.0 - increment_omega_deg: 0.2 - steps: 1800 diff --git a/src/aare/common/config/x10sa.yaml b/src/aare/common/config/x10sa.yaml deleted file mode 100644 index 60e5c983..00000000 --- a/src/aare/common/config/x10sa.yaml +++ /dev/null @@ -1,51 +0,0 @@ -beamline_id: "X10SA" - -gui: - cameras: - sample_camera_zmq_url: "tcp://sls-gpu-003:9089" #"tcp://x10sa-spark-01:9091" # - prediction_zmq_url: "tcp://sls-gpu-003:9089" #"tcp://x10sa-spark-01:9091" # - beamline_camera_url: "axis-accc8eb02488.psi.ch" - gonio_camera_url: "axis-accc8ea5e463.psi.ch" - gonio_camera_id: "1" - - daq: - daq_url: "https://mx-x10sa-queue-01.psi.ch" - cert_path: "/sls/x10sa/misc/.cert/10s.crt" - -shared: - jfjoch: - jfjoch_url: "http://sls-gpu-002:8080" - -daq: - hardware: - smargon_url: "http://x10sa-smargopolo.psi.ch:3000" - smargon_frontend_url: "http://x10sa-smargopolo.psi.ch:8080" - aerotech_url: "http://mx-x10sa-queue-01.psi.ch:5234" # Adjust if needed - tell_url: "http://PC17488:22222" - bec_url: "x10sa-bec-001.psi.ch" - redis_url: "x10sa-redis.psi.ch" - aarelc_url: "http://sls-gpu-003:9090" - dtz_safe_position: 300 - lens_magnification: 10 #or 5 currently - default_detector_distance_minimum: 150 - default_detector_distance_maximum: 900 - - db: - aaredb_url: "https://mx-aaredb-dmz-01.psi.ch/dispatcher" - - data_collection_settings: - default_raster_scan_settings: - exp_time_s: 0.01 - transmission: 1.0 - dtz: 200 - - default_rotation_settings: - exp_time_s: 0.01 - transmission: 1.0 - dtz: 200 - start_omega_deg: 0.0 - increment_omega_deg: 0.2 - steps: 1800 - - detector_distance_limit_modifier: 2.0 - maximum_flux: 1e12 diff --git a/src/aare/common/coordinate.py b/src/aare/common/coordinate.py deleted file mode 100644 index 76da6e25..00000000 --- a/src/aare/common/coordinate.py +++ /dev/null @@ -1,122 +0,0 @@ -from typing import Optional - -import numpy as np -from pydantic import BaseModel - - -class Coordinate(BaseModel): - x: float = 0.0 - y: float = 0.0 - z: float = 0.0 - - # Overload + - def __add__(self, other: "Coordinate") -> "Coordinate": - if isinstance(other, Coordinate): - return Coordinate( - x=self.x + other.x, y=self.y + other.y, z=self.z + other.z - ) - return NotImplemented - - # Overload - - def __sub__(self, other: "Coordinate") -> "Coordinate": - if isinstance(other, Coordinate): - return Coordinate( - x=self.x - other.x, y=self.y - other.y, z=self.z - other.z - ) - return NotImplemented - - # Overload * - def __mul__(self, other): - if isinstance(other, (int, float)): # Scalar multiplication - return Coordinate(x=self.x * other, y=self.y * other, z=self.z * other) - elif isinstance(other, Coordinate): # Dot product - return self.x * other.x + self.y * other.y + self.z * other.z - return NotImplemented - - # Overload / - def __truediv__(self, scalar: float) -> "Coordinate": - if isinstance(scalar, (int, float)) and scalar != 0: # Avoid division by zero - return Coordinate(x=self.x / scalar, y=self.y / scalar, z=self.z / scalar) - elif scalar == 0: - raise ValueError("Cannot divide by zero") - return NotImplemented - - def normalize(self) -> "Coordinate": - magnitude = np.sqrt(self.x**2 + self.y**2 + self.z**2) - if magnitude == 0: - raise ValueError("Cannot normalize a zero-magnitude vector.") - return self / magnitude - - def rotate(self, angle_deg: float, axis: str) -> "Coordinate": - """ - Rotate the coordinate around a specified axis by a given angle in degrees. - - :param angle_deg: The angle of rotation in degrees. - :param axis: The axis to rotate around ('x', 'y', or 'z'). - :return: A new Coordinate after rotation. - """ - angle_rad = np.radians(angle_deg) # Convert angle to radians - c = np.cos(angle_rad) - s = np.sin(angle_rad) - - if axis == "x": - # Rotate around x-axis (affects y, z) - y_new = self.y * c - self.z * s - z_new = self.y * s + self.z * c - return Coordinate(x=self.x, y=y_new, z=z_new) - elif axis == "y": - # Rotate around y-axis (affects x, z) - x_new = self.x * c + self.z * s - z_new = -self.x * s + self.z * c - return Coordinate(x=x_new, y=self.y, z=z_new) - elif axis == "z": - # Rotate around z-axis (affects x, y) - x_new = self.x * c - self.y * s - y_new = self.x * s + self.y * c - return Coordinate(x=x_new, y=y_new, z=self.z) - else: - raise ValueError("Invalid axis. Choose 'x', 'y', or 'z'.") - - -class SmargonCoordinate(BaseModel): - sh_mm: Optional[Coordinate] = None # SH coordinate in Smargon - phi_deg: Optional[float] = None - chi_deg: Optional[float] = None - - def eq(self, other: "SmargonCoordinate", tol: float) -> bool: - return ( - abs(self.sh_mm.x - other.sh_mm.x) < tol - and abs(self.sh_mm.y - other.sh_mm.y) < tol - and abs(self.sh_mm.z - other.sh_mm.z) < tol - and abs(self.phi_deg - other.phi_deg) < tol - and abs(self.chi_deg - other.chi_deg) < tol - ) - - def __eq__(self, other: object) -> bool: - if isinstance(other, SmargonCoordinate): - return self.eq(other, tol=0.1) - return NotImplemented - - -def positive_coords(value: Coordinate) -> Coordinate: - if value.x <= 0 or value.y <= 0: - raise ValueError("Coordinates must be positive") - return value - - -class AerotechCoordinate(BaseModel): - at_mm: Optional[Coordinate] = None - omega_deg: Optional[float] = None - - def eq(self, other: "AerotechCoordinate", tol: float) -> bool: - return ( - abs(self.at_mm.x - other.at_mm.x) < tol - and abs(self.at_mm.y - other.at_mm.y) < tol - and abs(self.at_mm.z - other.at_mm.z) < tol - and abs(self.omega_deg - other.omega_deg) < tol - ) - - def __eq__(self, other: object) -> bool: - if isinstance(other, AerotechCoordinate): - return self.eq(other, tol=0.01) - return NotImplemented \ No newline at end of file diff --git a/src/aare/common/diffraction_geometry.py b/src/aare/common/diffraction_geometry.py deleted file mode 100644 index a850227e..00000000 --- a/src/aare/common/diffraction_geometry.py +++ /dev/null @@ -1,82 +0,0 @@ -import math -from typing import Annotated, Tuple - -from pydantic import Field, BaseModel - -class DiffractionGeometry(BaseModel): - energy_keV: Annotated[float, Field(gt=1.0, lt=100.0)] - dtz_mm: Annotated[float, Field(gt=10.0, lt=5000.0)] - pixel_size_mm: Annotated[float, Field(ge=0.05, le=0.5)] - beam_center_pxl: Tuple[float, float] - detector_size_pxl: Tuple[int, int] - detector_description: str - detector_serial_number: str - poni_rot1_rad: float - poni_rot2_rad: float - - @property - def detector_max_radius_pxl(self): - x0 = self.detector_size_pxl[0] - self.beam_center_pxl[0] - x1 = self.beam_center_pxl[0] - y0 = self.detector_size_pxl[1] - self.beam_center_pxl[1] - y1 = self.beam_center_pxl[1] - return max(abs(x0), abs(x1), abs(y0), abs(y1)) - - @property - def detector_radius_mm(self): - return self.detector_max_radius_pxl * self.pixel_size_mm - - @property - def wavelength_angstrom(self): - return 12.398 / self.energy_keV - - @property - def max_resolution_angstrom(self): - return self.resolution_angstrom(self.dtz_mm) - - def resolution_angstrom(self, exp_dtz_mm: float) -> float: - if exp_dtz_mm <= 0: - raise ValueError(f"Detector distance must be positive {exp_dtz_mm}") - - theta = math.atan(self.detector_radius_mm / exp_dtz_mm)*0.5 - return self.wavelength_angstrom / (2 * math.sin(theta)) - - def calc_dtz_mm(self, exp_resolution_angstrom: float) -> float: - if exp_resolution_angstrom <= 0: - raise ValueError("Resolution must be positive") - x = self.wavelength_angstrom / (2 * exp_resolution_angstrom) - if x >= 1.0 or x <= -1.0: - return 0.0 - theta = math.asin(x) - return self.detector_radius_mm / math.tan(2*theta) - -if __name__ == "__main__": - from aare.devices.jfjoch import JFJochWrapper - from aare.common.beamline import mx_beamline - from aare.daq.config import BeamlineConfig - bl = mx_beamline() - client = JFJochWrapper(bl) - cfg = BeamlineConfig(bl) - print(cfg) - det_cfg = client.detector() - print(det_cfg) - print(f"pixel_size: {det_cfg.pixel_size_mm:.4f} mm") - print(f"Detector height: {det_cfg.height:.2f} pixel") - print(f"Detector width: {det_cfg.width:.2f} pixel") - print(f"beam centre = ({cfg.beam_center[0]:.2f}, {cfg.beam_center[1]:.2f})") - geom = DiffractionGeometry( - energy_keV=12.4, - dtz_mm=200.0, - pixel_size_mm=det_cfg.pixel_size_mm, - beam_center_pxl=(det_cfg.width/2, det_cfg.height/2), - detector_size_pxl=(det_cfg.width, det_cfg.height), - detector_description=det_cfg.description, - detector_serial_number=det_cfg.serial_number, - poni_rot1_rad=-0.001396263, - poni_rot2_rad=-0.003839724, - ) - print(f"geom.max_resolution_angstrom: {geom.max_resolution_angstrom:.2f} Angstrom") - print(f"geom.calc_dtz_mm(3.9): {geom.calc_dtz_mm(3.9):.2f} mm") - - print(f"geom.resolution_angstrom(1244.42): {geom.resolution_angstrom(1244.42):.2f} Angstrom") - print(f"geom.detector_radius_mm: {geom.detector_radius_mm:.2f} mm") diff --git a/src/aare/common/error_codes.py b/src/aare/common/error_codes.py deleted file mode 100644 index cc103a42..00000000 --- a/src/aare/common/error_codes.py +++ /dev/null @@ -1,189 +0,0 @@ -from __future__ import annotations - -import re -from enum import StrEnum - - -class AuthErrorCode(StrEnum): - """ - Stable, machine-readable error codes used across the API. - - Rules: - - never rename an existing value (treat as public API) - - only add new values - """ - - # Generic / defaults - AUTHENTICATION_ERROR = "AUTHENTICATION_ERROR" - AUTHENTICATION_FAILED = "AUTHENTICATION_FAILED" - FORBIDDEN = "FORBIDDEN" - HTTP_ERROR = "HTTP_ERROR" - INTERNAL_SERVER_ERROR = "INTERNAL_SERVER_ERROR" - - # Auth/JWT - INVALID_TOKEN = "INVALID_TOKEN" - SESSION_ALREADY_ACTIVE = "SESSION_ALREADY_ACTIVE" - - # Authorization - NOT_STAFF = "NOT_STAFF" - NOT_IN_ACTIVE_PGROUP = "NOT_IN_ACTIVE_PGROUP" - NOT_BATON_HOLDER = "NOT_BATON_HOLDER" - - -class AareErrorCode(StrEnum): - """New error-code enum -- one value per concrete exception class in the - AareException hierarchy. - - Naming: SCREAMING_SNAKE_CASE that mirrors the class name (e.g. - ``TellCommunicationError`` → ``TELL_COMMUNICATION_ERROR``). - - Rules: - - never rename an existing value (treat as public API) - - only add new values - - one value per concrete exception class; family base classes do NOT get - values (they're abstract for isinstance matching only) - """ - - # Tell family - TELL_COMMUNICATION_ERROR = "TELL_COMMUNICATION_ERROR" - TELL_CONNECTION_EXCEPTION = "TELL_CONNECTION_EXCEPTION" - CRITICAL_TELL_EXCEPTION = "CRITICAL_TELL_EXCEPTION" - WARNING_TELL_EXCEPTION = "WARNING_TELL_EXCEPTION" - TELL_COMMAND_WHILE_BUSY_EXCEPTION = "TELL_COMMAND_WHILE_BUSY_EXCEPTION" - MOUNTING_FAILED = "MOUNTING_FAILED" - UNMOUNTING_FAILED = "UNMOUNTING_FAILED" - - # Other device-comm families - SMARGON_COMMUNICATION_ERROR = "SMARGON_COMMUNICATION_ERROR" - AEROTECH_COMMUNICATION_ERROR = "AEROTECH_COMMUNICATION_ERROR" - JF_JOCH_COMMUNICATION_ERROR = "JF_JOCH_COMMUNICATION_ERROR" - BEC_COMMUNICATION_ERROR = "BEC_COMMUNICATION_ERROR" - AARE_DB_COMMUNICATION_ERROR = "AARE_DB_COMMUNICATION_ERROR" - - # Beamline-state family - STATE_TRANSITION_FAILED = "STATE_TRANSITION_FAILED" - MAINTENANCE_STATE_EXCEPTION = "MAINTENANCE_STATE_EXCEPTION" - BEAMLINE_BUSY_EXCEPTION = "BEAMLINE_BUSY_EXCEPTION" - BEAMLINE_BUSY_TIMEOUT_EXCEPTION = "BEAMLINE_BUSY_TIMEOUT_EXCEPTION" - - # Data-collection family - DATA_COLLECTION_EXCEPTION = "DATA_COLLECTION_EXCEPTION" - RASTER_SCAN_EXCEPTION = "RASTER_SCAN_EXCEPTION" - - # Other automation - LOOP_CENTERING_FAILED = "LOOP_CENTERING_FAILED" - AXC_FAILED = "AXC_FAILED" - AUTO_RASTER_SAMPLE_SKIPPED = "AUTO_RASTER_SAMPLE_SKIPPED" - TRANSFORMATION_INVALID_EXCEPTION = "TRANSFORMATION_INVALID_EXCEPTION" - MAGNET_POSITION_SENSOR_ERORR = "MAGNET_POSITION_SENSOR_ERORR" # NOTE: class name "Erorr" has a typo; preserved - SMART_MAGNET_FAULT_EXCEPTION = "SMART_MAGNET_FAULT_EXCEPTION" - DOOR_SAFETY_ERROR = "DOOR_SAFETY_ERROR" - - # User errors - MANUAL_MOUNT_EXCEPTION = "MANUAL_MOUNT_EXCEPTION" - SAMPLE_EXCEPTION = "SAMPLE_EXCEPTION" - - # Auth errors -- usually surfaced via finer-grained AuthErrorCode, but - # these provide a class-level fallback when no specific auth code is set. - AUTHENTICATION_EXCEPTION = "AUTHENTICATION_EXCEPTION" - USER_RIGHTS_EXCEPTION = "USER_RIGHTS_EXCEPTION" - - # Fallback for unmapped exceptions caught by the bare-Exception handler - INTERNAL_ERROR = "INTERNAL_ERROR" - - -_CAMEL_BOUNDARY_1 = re.compile(r"(.)([A-Z][a-z]+)") -_CAMEL_BOUNDARY_2 = re.compile(r"([a-z0-9])([A-Z])") - - -def code_for_exception_class(cls_name: str) -> str: - """Convert a CamelCase exception class name to SCREAMING_SNAKE_CASE. - - Examples: - TellCommunicationError → TELL_COMMUNICATION_ERROR - AareDBCommunicationError → AARE_DB_COMMUNICATION_ERROR - AXCFailed → AXC_FAILED - """ - s1 = _CAMEL_BOUNDARY_1.sub(r"\1_\2", cls_name) - s2 = _CAMEL_BOUNDARY_2.sub(r"\1_\2", s1) - return s2.upper() - - -_ERROR_CODE_HELP: dict[str, str] = { - # Auth/JWT - AuthErrorCode.AUTHENTICATION_ERROR: ( - "Generic authentication problem. Usually means the request lacked valid credentials " - "(expired/invalid token, missing Authorization header, etc.)." - ), - AuthErrorCode.AUTHENTICATION_FAILED: ( - "Authentication failed during login/token creation. Typically incorrect credentials " - "or an inability to validate the user." - ), - AuthErrorCode.INVALID_TOKEN: ( - "The provided token could not be decoded/validated (bad signature, expired, malformed). " - "Re-authenticate to obtain a new token." - ), - AuthErrorCode.SESSION_ALREADY_ACTIVE: ( - "A different session currently owns control. Use “force current session” (if allowed) " - "or wait for the active session to expire/end." - ), - AuthErrorCode.NOT_BATON_HOLDER: ( - "A different session currently owns the baton. Request the baton or wait for active session to expire/end"), - # Authorization - AuthErrorCode.FORBIDDEN: ( - "Generic permissions failure. The user is authenticated but not allowed to perform this action." - ), - AuthErrorCode.NOT_STAFF: ( - "This action requires staff privileges. Log in with a staff account or ask staff to perform it." - ), - AuthErrorCode.NOT_IN_ACTIVE_PGROUP: ( - "You are not a member of the currently active p-group. Change p-group or use an account " - "that belongs to the active group." - ), - # Generic - AuthErrorCode.HTTP_ERROR: ( - "Generic HTTP error wrapper. The server returned an HTTPException that wasn’t mapped to a more specific code." - ), - AuthErrorCode.INTERNAL_SERVER_ERROR: ( - "Unhandled server error. Check server logs for a stack trace and context." - ), -} - -def error_code_help(code: str) -> str | None: - """ - Return a human help message for a code string, if known. - Accepts either enum value strings or raw strings. - """ - if not code: - return None - return _ERROR_CODE_HELP.get(str(code)) - - -def export_error_code_help() -> dict[str, str]: - """ - Export help text as {"CODE": "help text", ...} - """ - return {str(k): str(v) for k, v in _ERROR_CODE_HELP.items()} - -def export_error_codes_grouped() -> dict[str, dict[str, str]]: - """ - Export codes grouped by enum class name: - - { - "AuthErrorCode": {"INVALID_TOKEN": "INVALID_TOKEN", ...}, - "AareErrorCode": {"TELL_COMMUNICATION_ERROR": "TELL_COMMUNICATION_ERROR", ...} - } - """ - enums: tuple[type[StrEnum], ...] = (AuthErrorCode, AareErrorCode) - return {e.__name__: {c.name: str(c.value) for c in e} for e in enums} - -def export_error_codes() -> dict[str, str]: - """ - Backwards-compatible, flat export used by older clients/tests/docs: - - {"INVALID_TOKEN": "INVALID_TOKEN", ...} - - NOTE: This intentionally exports only AuthErrorCode to avoid breaking - existing consumers that assume a flat map and/or specific keys. - """ - return {c.name: str(c.value) for c in AuthErrorCode} diff --git a/src/aare/common/exception_handler.py b/src/aare/common/exception_handler.py deleted file mode 100644 index b8581145..00000000 --- a/src/aare/common/exception_handler.py +++ /dev/null @@ -1,517 +0,0 @@ -from __future__ import annotations - -import time -from typing import ClassVar - -from aare.common.logger_config import setup_logger -from aare.common.error_codes import AuthErrorCode - -logger = setup_logger("aareDAQ") - - -class AareException(Exception): - """Root of the aare exception hierarchy. - - Class-level ``critical`` is the default; pass ``critical=...`` to the - constructor of an "optionally critical" subclass to override on a single - raise site. - """ - - critical: ClassVar[bool] = False - - def __init__(self, *args, critical: bool | None = None, **kwargs): - super().__init__(*args, **kwargs) - if critical is not None: - self.critical = critical - - -class AutomationError(AareException): - """Errors that happen during a DAQ operation; route through automation flows.""" - - -class AareUserError(AareException): - """Bad input / mode misuse; routes through the input-correction flow.""" - - -class AareAuthError(AareException): - """Authentication / authorization failures; routes through re-auth flow.""" - - -# --------------------------------------------------------------------------- -# Device-family base classes -# -# These exist purely so watchers can match a whole device family via -# ``isinstance(e, TellException)`` etc. Per the plan they are intentionally -# empty -- do not add behavior. -# --------------------------------------------------------------------------- - -class TellException(AutomationError): - pass - - -class SmargonException(AutomationError): - pass - - -class AerotechException(AutomationError): - pass - - -class JFJochException(AutomationError): - pass - - -class BECException(AutomationError): - pass - - -class BeamlineStateException(AutomationError): - pass - - -# --------------------------------------------------------------------------- -# Concrete automation exceptions -# --------------------------------------------------------------------------- - -class TransformationInvalidException(AutomationError): - critical: ClassVar[bool] = True - - def __init__(self, message: str = "Transformation is not implemented", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class StateTransitionFailed(BeamlineStateException): - critical: ClassVar[bool] = True - - def __init__(self, message: str = "Beamline state transition failed", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class MaintenanceStateException(BeamlineStateException): - critical: ClassVar[bool] = True - - def __init__(self, message: str = "Beamline is in Maintenance state", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class DataCollectionException(AutomationError): - """Group parent for data-collection-time exceptions; also raisable directly.""" - - def __init__(self, message: str = "Data collection failed", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class RasterScanException(DataCollectionException): - def __init__(self, message: str = "Data collection failed", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class AutoRasterSampleSkipped(AutomationError): - """Raised when automation should skip the current sample because auto-raster is too large.""" - - -class LoopCenteringFailed(AutomationError): - def __init__(self, message: str = "Loop Centering did not detect a sample", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class UnmountingFailed(TellException): - """Sample failed to unmount. Optionally critical via instance flag.""" - - def __init__(self, message: str = "A sample was not unmounted", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class MountingFailed(TellException): - """Sample failed to mount. Optionally critical via instance flag. - """ - - def __init__(self, message: str = "A sample was not mounted", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class ManualMountException(AareUserError): - """Custom exception for manual mounting""" - def __init__(self, message: str = "Manual mounting failed", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class SmartMagnetFaultException(AutomationError): - """Custom exception for smart magnet fault""" - - critical: ClassVar[bool] = True - - def __init__(self, message: str = "Smart magnet fault", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class DoorSafetyError(AutomationError): - """Hutch personnel-safety system does not permit robot motion. - - Raised before a mount/unmount when ``…-EH1-PSYS:PROHIBITED-STATE`` is not in - the prohibited state (door safety could not be activated, so the robot will - not move) or when ``…-EH1-PSYS:ALARM-STATE`` reports an active alarm. Always - critical so automation halts and the GUI shows a pop-up. - """ - - critical: ClassVar[bool] = True - - def __init__(self, message: str = "Door safety could not be activated", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class TellCommandWhileBusyException(TellException): - """Custom exception for trying to move Tell when it is busy""" - - def __init__(self, message: str = "Tell is busy", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class TellConnectionException(TellException): - """Custom exception for connection problems""" - - critical: ClassVar[bool] = True - - def __init__(self, message: str = "Lost connection to Tell", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class WarningTellException(TellException): - def __init__(self, message: str = "Warning error in TELL", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class CriticalTellException(TellException): - critical: ClassVar[bool] = True - - def __init__(self, message: str = "Critical error in TELL", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class AXCFailed(AutomationError): - def __init__(self, message: str = "Auto X-ray centering failed", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class BeamlineBusyException(BeamlineStateException): - critical: ClassVar[bool] = True - - def __init__(self, message: str = "Beamline is in busy state", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class BeamlineBusyTimeoutException(BeamlineStateException): - critical: ClassVar[bool] = True - - def __init__(self, message: str = "Beamline busy state expired during operation", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class SampleException(AareUserError): - """Requested sample could not be located. - - NOTE: not listed in the redesign hierarchy (§1.2); parented under - AareUserError pending confirmation -- "sample not found" reads as bad - input rather than a runtime failure.""" - - def __init__(self, message: str = "Sample not found", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -# --------------------------------------------------------------------------- -# Auth exceptions -# --------------------------------------------------------------------------- - -class AuthenticationException(AareAuthError): - def __init__(self, - message: str = "Authentication failed.", - *, - status_code: int = 401, - headers: dict[str, str] | None = None, - code: AuthErrorCode = AuthErrorCode.AUTHENTICATION_FAILED, - critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - self.status_code = status_code - self.headers = headers - self.code = code - - def __str__(self) -> str: - return self.message - - -class UserRightsException(AareAuthError): - _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.", - *, - status_code: int = 403, - headers: dict[str, str] | None = None, - code: AuthErrorCode = AuthErrorCode.FORBIDDEN, - critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - self.status_code = status_code - self.headers = headers - self.code = code - - def __str__(self) -> str: - return self.message - - -# --------------------------------------------------------------------------- -# Device communication exceptions -# --------------------------------------------------------------------------- - -class SmargonCommunicationError(SmargonException): - """ - Raised when Smargon HTTP communication fails (connection refused, timeout, bad HTTP status, etc). - Keep the original exception in `__cause__` by using `raise ... from e`. - """ - - critical: ClassVar[bool] = True - - def __init__( - self, - message: str = "Smargon communication error", - *, - endpoint: str | None = None, - base_url: str | None = None, - operation: str | None = None, - status_code: int | None = None, - critical: bool | None = None, - ): - super().__init__(message, critical=critical) - self.message = message - self.endpoint = endpoint - self.base_url = base_url - self.operation = operation - self.status_code = status_code - - def __str__(self) -> str: - return self.message - - -class TellCommunicationError(TellException): - """ - Raised when TELL HTTP/PShell communication fails (timeouts, connection refused, etc). - Intended to be caught centrally by FastAPI exception handlers. - """ - - critical: ClassVar[bool] = True - - def __init__( - self, - message: str = "TELL communication error", - *, - endpoint: str | None = None, - base_url: str | None = None, - operation: str | None = None, - critical: bool | None = None, - ): - super().__init__(message, critical=critical) - self.message = message - self.endpoint = endpoint - self.base_url = base_url - self.operation = operation - - def __str__(self) -> str: - return self.message - - -class JFJochCommunicationError(JFJochException): - """ - Raised when JFJoch HTTP/API communication fails. - Intended for scan-time fallbacks and GUI-visible alerts. - """ - - critical: ClassVar[bool] = True - - 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, - critical: bool | None = None, - ): - super().__init__(message, critical=critical) - self.message = message - self.operation = operation - self.endpoint = endpoint - self.base_url = base_url - self.status_code = status_code - - def __str__(self) -> str: - return self.message - - -class AareDBCommunicationError(AutomationError): - """Raised when AareDB HTTPS communication fails. - - Default not critical -- DB hiccups don't always halt automation. Raise with - ``critical=True`` at call sites where a DB read failure should escalate.""" - - 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, - critical: bool | None = None, - ): - super().__init__(message, critical=critical) - self.message = message - self.operation = operation - self.endpoint = endpoint - self.base_url = base_url - self.status_code = status_code - - def __str__(self) -> str: - return self.message - - -class AerotechCommunicationError(AerotechException): - """ - 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`. - """ - - critical: ClassVar[bool] = True - - 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, - critical: bool | None = None, - ): - super().__init__(message, critical=critical) - self.message = message - self.endpoint = endpoint - self.base_url = base_url - self.operation = operation - self.status_code = status_code - - def __str__(self) -> str: - return self.message - - -class MagnetPositionSensorErorr(AutomationError): - critical: ClassVar[bool] = True - - def __init__(self, message: str = "Magnet position sensor error", *, critical: bool | None = None): - super().__init__(message, critical=critical) - self.message = message - - def __str__(self) -> str: - return self.message - - -class BECCommunicationError(BECException): - critical: ClassVar[bool] = True - - def __init__( - self, - message: str = "BEC communication error", - *, - exception: Exception | None = None, - operation: str | None = None, - endpoint: str | None = None, - base_url: str | None = None, - critical: bool | None = None, - ): - super().__init__(message, critical=critical) - self.message = message - self.operation = operation - self.exception = exception - self.endpoint = endpoint - self.base_url = base_url - - def __str__(self) -> str: - return self.message diff --git a/src/aare/common/find_xtal.py b/src/aare/common/find_xtal.py deleted file mode 100644 index 0e3acf0e..00000000 --- a/src/aare/common/find_xtal.py +++ /dev/null @@ -1,626 +0,0 @@ -from typing import List, Optional, Callable - -import numpy as np -from scipy import ndimage - -from aare.common.models import CrystalSize -from aare.common.raster_grid import RasterGridRequest, CenterOfMassModel -from aare.common.logger_config import setup_logger - -logger = setup_logger('aareDAQ') - -def identify_crystal_raster(result, r: RasterGridRequest) -> CenterOfMassModel | None: - images = result.images - if images and any(getattr(img, "spots_low_res", 0) for img in images): - # - # indexed_images = [img for img in images if img.index and img.spots_low_res > 4 and img.bkg > 4.5] - # - # if indexed_images: - # # indexed_images = [img for img in indexed_images if img.spots_indexed > 10] - # images = indexed_images - # - # filtered_images = [img for img in images - # if img.spots_ice is not None and img.spots_low_res > 4 and ( - # img.spots_ice / img.spots_low_res) < 5.0 - # and (img.spots_ice / img.spots_low_res) != 1] - # - # if indexed_images: - # logger.debug(f"Find image by maximum number of spots indexed") - # max_image = max(images, key=lambda img: img.spots_indexed) - # max_spots = max_image.spots_indexed - # max_images = [img for img in images if img.spots_indexed == max_spots] - # max_image = max_images[len(max_images) // 2] - # - # logger.debug(f"Image with maximum spots_low_res: {max_image}") - # logger.debug(f"Maximum spots_indexed value: {max_image.spots_indexed}") - # logger.debug(f"Maximum spots_low_res value: {max_image.spots_low_res}") - # else: - logger.debug(f"Find image by maximum number of low resolution spots") - max_image = max(images, key=lambda img: img.spots_low_res) - logger.debug(f"Image with maximum spots_low_res: {max_image}") - logger.debug(f"Maximum spots_low_res value: {max_image.spots_low_res}") - logger.debug(f"Maximum image found at grid coordiantes {max_image.nx}, {max_image.ny}") - logger.debug(f"Maximum image found at umL {max_image.nx * r.grid_size_mm.x}, {max_image.ny * r.grid_size_mm.y}") - com = CenterOfMassModel(n_x=max_image.nx, n_y=max_image.ny, max_image=max_image.number) - com_mm = com.get_com_mm(r) - grid_mm_x = com_mm.x - grid_mm_y = com_mm.y - - logger.debug(f"Grid coordinates in mm: x={grid_mm_x}, y={grid_mm_y}") - return com - else: - return None - -def rebuild_array_from_scan_results(scan_results: List, - value_field: str, - array_shape: Optional[tuple] = None, - nx_field: str = 'nx', - ny_field: str = 'ny', - default_value: float = 0.0, - threshold: Optional[float] = None, - condition_func: Optional[Callable] = None, - apply_filter_before: bool = True - ) -> np.ndarray: - # Extract coordinates and values - positions = [] - values = [] - - for result in scan_results: - nx = getattr(result, nx_field) - ny = getattr(result, ny_field) - value = getattr(result, value_field) - - # Skip if coordinates are None - if nx is None or ny is None: - continue - - positions.append((int(nx), int(ny))) # Note: (row, col) = (ny, nx) - if not value: - value = 0.0 - values.append(float(value)) - - if not positions: - logger.error("No valid positions found in scan results") - raise ValueError("No valid positions found in scan results") - - # Determine array shape - if array_shape is None: - max_row = max(pos[0] for pos in positions) - max_col = max(pos[1] for pos in positions) - array_shape = (max_row + 1, max_col + 1) - - # Initialize array with default values - result_array = np.full(array_shape, default_value, dtype=float) - - # Apply pre-filtering if requested - if apply_filter_before: - filtered_data = [] - for pos, val in zip(positions, values): - keep_value = True - - # Apply threshold filter - if threshold is not None and val < threshold: - keep_value = False - - # Apply custom condition - if condition_func is not None and not condition_func(val): - keep_value = False - - if keep_value: - filtered_data.append((pos, val)) - else: - filtered_data.append((pos, 0.0)) - - # Fill array with filtered values - for pos, val in filtered_data: - if 0 <= pos[0] < array_shape[0] and 0 <= pos[1] < array_shape[1]: - result_array[pos[0], pos[1]] = val - else: - # Fill array first, then apply filters - for pos, val in zip(positions, values): - if 0 <= pos[0] < array_shape[0] and 0 <= pos[1] < array_shape[1]: - result_array[pos[0], pos[1]] = val - - # Apply post-filtering - if threshold is not None: - result_array[result_array < threshold] = 0.0 - - if condition_func is not None: - mask = np.vectorize(condition_func)(result_array) - result_array[~mask] = 0.0 - - return result_array - -def create_quality_filtered_array(scan_results: List, - value_field: str, - min_spots: Optional[int] = None, - min_efficiency: Optional[float] = 1.0, - min_background: Optional[float] = None, - exclude_ice: Optional[bool] = True, - min_low_res_spots: Optional[float] = 10.0, - **kwargs - ) -> np.ndarray: - """ - Create array with comprehensive quality filtering - """ - - def quality_condition(result, max_spots_low_res, min_background, min_spots, min_efficiency, min_low_res_spots, max_filter: float = 0.3): - if exclude_ice and (result.spots_ice / max(result.spots_low_res, 1.0)) == 1.0: - # print(f"all ice for {result.number}") - return False - if min_low_res_spots and result.spots_low_res < min_low_res_spots: - return False - if result.spots_low_res < (max_spots_low_res * max_filter): - return False - if exclude_ice and result.spots_ice > result.spots * 0.8: # More than 50% ice - # print(f"more than 80% ice for {result.number}") - return False - if result.index: - # print(f"index is True for {result.number}") - return True - if min_spots and result.spots < min_spots: - # print(f"{result.spots} is less than {min_spots} for {result.number}") - return False - if result.spots_low_res < min_background: - return False - - if result.efficiency < min_efficiency: - return False - - return True - - # Filter results first - filtered_results = [] - - if min_spots is None: - min_spots = min((result.spots for result in scan_results if result.spots is not None), default=1) - if min_low_res_spots is None: - min_low_res_spots = min((result.spots_low_res for result in scan_results if result.spots_low_res is not None), default=1) - if min_background is None: - min_background = min((result.bkg for result in scan_results if result.bkg is not None), default=1) - if min_efficiency is None: - min_efficiency = 1.0 - max_spots_low_res = max((result.spots_low_res for result in scan_results if result.spots is not None), default=1) - - for result in scan_results: - if result.nx is not None and result.ny is not None: - if quality_condition(result, max_spots_low_res=max_spots_low_res, min_background=min_background, min_spots=min_spots, - min_efficiency=min_efficiency, min_low_res_spots=min_low_res_spots): - filtered_results.append(result) - else: - # Create a copy with zero value for filtered positions - import copy - zero_result = copy.copy(result) - setattr(zero_result, value_field, 0) - filtered_results.append(zero_result) - - return rebuild_array_from_scan_results(filtered_results, value_field, **kwargs) - - -def get_xtal_size(crystal_size, result_array, r:RasterGridRequest): - # Optional: get bounding box of the largest object - try: - labeled_array, num_objects = ndimage.label(result_array) - areas = ndimage.sum(np.ones_like(result_array, dtype=np.int32), labeled_array, - index=range(1, num_objects + 1)) - largest_idx = int(np.argmax(areas)) + 1 # +1 because labels start at 1 - largest_area = int(areas[largest_idx - 1]) - logger.info(f"Largest object label: {largest_idx}, area (px): {largest_area}") - - object_mask = labeled_array == largest_idx - rows = np.any(object_mask, axis=1) - cols = np.any(object_mask, axis=0) - row_min, row_max = np.where(rows)[0][[0, -1]] - col_min, col_max = np.where(cols)[0][[0, -1]] - logger.info(f"Largest bbox: width={col_max - col_min}, height={row_max - row_min}") - logger.info( - f"Largest bbox: width={(col_max - col_min) * r.grid_size_mm.x}, y={(row_max - row_min) * r.grid_size_mm.y}") - if r.n_x == 1: - logger.debug("calculating z") - crystal_size = CrystalSize(x=crystal_size.x, y=crystal_size.y, - z=(col_max - col_min) * r.grid_size_mm.y * 1000) - logger.info(f"Crystal Size: x={crystal_size.x}, y={crystal_size.y}, z={crystal_size.z}") - - else: - logger.debug("calculating x and y") - crystal_size = CrystalSize(x=(row_max - row_min) * r.grid_size_mm.x * 1000, - y=(col_max - col_min) * r.grid_size_mm.y * 1000, - z=crystal_size.z) - logger.info(f"Crystal Size: x={crystal_size.x}, y={crystal_size.y}, z={crystal_size.z}") - except ValueError as e: - logger.error(f"error calculating xtal size: {e}") - crystal_size = CrystalSize(x=0,y=0,z=0) - - return crystal_size - - -def get_best_b_factor(result_list: List): - if not result_list: - return None - best_b_factor = min((img for img in result_list if img.b is not None), - key=lambda img: img.b, - default=None) - if best_b_factor is None: - return None - logger.info(f"Best b: {best_b_factor.b}") - return best_b_factor.b - -def get_best_res(result_list: List): - if not result_list: - return None - best_res = min((img for img in result_list if img.res is not None), - key=lambda img: img.res, - default=None) - if best_res is None: - return None - logger.info(f"Best res: {best_res.res}") - return best_res.res - -def com_nan_check(com): - if np.isnan(com.n_x) or np.isnan(com.n_y): - return False - else: - return True - -def get_result_list_from_com(images, com: CenterOfMassModel): - if not com_nan_check(com): - return None - cx, cy = com.n_x, com.n_y - start_x, end_x = round(cx - 1), round(cx + 1) - start_y, end_y = round(cy - 1), round(cy + 1) - logger.info(f"range x {start_x} {end_x}, y {start_y} {end_y}") - result_list = [img for img in images - if start_x <= img.nx <= end_x and start_y <= img.ny <= end_y] - return result_list - -def get_com_image_number(com, images): - for image in images: - if image.nx == round(com[0]) and image.ny == round(com[1]): - logger.info(f"com found for image: {image.number}") - return image.number - return None - -def _to_com_model(coords, images, label: str) -> CenterOfMassModel | None: - """Validate (n_x, n_y) grid coords and wrap them in a CenterOfMassModel. - - Shared by raster_centre_of_mass and raster_highest_score, which differ only - in how they pick the target cell. - """ - if coords and not np.isnan(coords[0]) and not np.isnan(coords[1]): - logger.info(f"{label}: {coords}") - max_image = get_com_image_number(coords, images) - return CenterOfMassModel(n_x=coords[0], n_y=coords[1], max_image=max_image) - elif coords and np.isnan(coords[0]) and np.isnan(coords[1]): - logger.warning(f"{label} is nan: {coords[0]}, {coords[1]}") - return None - else: - logger.warning(f"No valid {label} found") - return None - - -def raster_centre_of_mass(result_array, images) -> CenterOfMassModel | None: - # grid_mm_x and grid_mm_y are relative to the top left corner of raster grid - com = ndimage.center_of_mass(result_array) - return _to_com_model(com, images, "Center of mass") - - -def raster_highest_score(images) -> CenterOfMassModel | None: - """Target the grid cell with the highest crystal score (compute_crystal_score_array).""" - score_array = compute_crystal_score_array(images) - return _to_com_model(_max_cell(score_array), images, "Highest score") - -def has_sufficient_low_res_spots( - result_array: np.ndarray, - min_spots_low_res: float, - ) -> bool: - if result_array is None or result_array.size == 0: - logger.warning("Empty result_array passed to has_sufficient_low_res_spots") - return False - - max_val = float(np.nanmax(result_array)) - logger.info(f"Max spots_low_res in raster result_array: {max_val}") - - if max_val < min_spots_low_res: - logger.info( - "Raster rejected: max spots_low_res " - f"{max_val} < required {min_spots_low_res}" - ) - return False - - return True - - -def compute_crystal_score_array( - scan_results: List, - w_bkg: float = 0.25, - w_low_res: float = 0.55, - w_indexed: float = 0.20, -) -> np.ndarray: - """Combine bkg (25%), spots_low_res (55%), and spots_indexed (20%) into a 0–100 score. - - Each field is min-max normalised to [0, 100] within the grid before weighting, - so the final score is the probability (0–100) that a pixel belongs to a crystal. - """ - arr_bkg = rebuild_array_from_scan_results(scan_results, "bkg") - arr_low = rebuild_array_from_scan_results(scan_results, "spots_low_res") - arr_idx = rebuild_array_from_scan_results(scan_results, "spots_indexed") - - def _norm(a: np.ndarray) -> np.ndarray: - mn, mx = float(a.min()), float(a.max()) - if mx == mn: - return np.zeros_like(a, dtype=float) - return (a - mn) / (mx - mn) * 100.0 - - return w_bkg * _norm(arr_bkg) + w_low_res * _norm(arr_low) + w_indexed * _norm(arr_idx) - - -def _max_cell(arr: np.ndarray) -> tuple[int, int]: - """Return (nx, ny) of the cell with the highest value.""" - idx = np.unravel_index(np.argmax(arr), arr.shape) - return int(idx[0]), int(idx[1]) - - -def _draw_panel( - ax, - arr: np.ndarray, - mask: np.ndarray, - label: str, - threshold: float, - grid_size_mm: Optional[tuple[float, float]] = None, - cbar_label: str = "spots_low_res", -) -> None: - """Shared helper: heatmap + crystal contour + max-cell square on one Axes. - - If grid_size_mm=(step_x_mm, step_y_mm) is provided, the crystal size in µm - is computed via get_xtal_size and shown in the panel title. - """ - import matplotlib.patches as mpatches - import matplotlib.pyplot as plt - from aare.common.coordinate import Coordinate - - n_nx, n_ny = arr.shape - max_nx, max_ny = _max_cell(arr) - n_cells = int(mask.sum()) - - im = ax.imshow( - arr.T, - origin="lower", - cmap="viridis", - aspect="equal", - extent=[-0.5, n_nx - 0.5, -0.5, n_ny - 0.5], - ) - plt.colorbar(im, ax=ax, label=cbar_label, fraction=0.046, pad=0.04) - ax.contour( - np.arange(n_nx), np.arange(n_ny), mask.T.astype(float), - levels=[0.5], - colors="cyan", - linewidths=1.8, - ) - rect = mpatches.Rectangle( - (max_nx - 0.5, max_ny - 0.5), 1, 1, - linewidth=2, edgecolor="red", facecolor="none", - ) - ax.add_patch(rect) - - size_line = "" - if grid_size_mm is not None and n_cells > 0: - r = RasterGridRequest( - exp_time_s=0.0, - n_x=n_nx, - n_y=n_ny, - grid_size_mm=Coordinate(x=grid_size_mm[0], y=grid_size_mm[1]), - smargon_top_left=None, - ) - xtal_size = get_xtal_size(CrystalSize(x=0, y=0, z=0), mask.astype(float), r) - size_line = f"\n{xtal_size.x:.0f}×{xtal_size.y:.0f} µm" - - ax.set_title( - f"{label}\nthresh≈{threshold:.1f} | {n_cells} cells{size_line}", - fontsize=8, - ) - ax.set_xlabel("nx", fontsize=7) - ax.set_ylabel("ny", fontsize=7) - ax.tick_params(labelsize=6) - - -# ── Clustering / thresholding methods ───────────────────────────────────────── - -def crystal_mask_corner_background(arr: np.ndarray) -> tuple[np.ndarray, float]: - """Threshold = max(far-corner value, floor=10). - - Simple, parameter-free. Returns too many cells when corners are zero - because any nonzero value passes. - """ - n_nx, n_ny = arr.shape - corners = [arr[0, 0], arr[n_nx - 1, 0], arr[0, n_ny - 1], arr[n_nx - 1, n_ny - 1]] - threshold = max(float(np.max(corners)), 10.0) - return arr > threshold, threshold - - -def crystal_mask_otsu_nonzero(arr: np.ndarray) -> tuple[np.ndarray, float]: - """Otsu threshold computed only on the nonzero values. - - Finds the natural gap in the signal distribution. - Ignores the large mass of background zeros so the threshold is - placed within the diffraction-signal population. - """ - from skimage.filters import threshold_otsu - - nonzero = arr[arr > 0] - if nonzero.size == 0: - return np.zeros_like(arr, dtype=bool), 0.0 - threshold = float(threshold_otsu(nonzero)) - return arr > threshold, threshold - - -def crystal_mask_mean_sigma(arr: np.ndarray, n_sigma: float = 0.5) -> tuple[np.ndarray, float]: - """Threshold = mean + n_sigma * std of nonzero values. - - n_sigma=0.5 keeps cells within ~1/2 std above average signal. - Raise n_sigma to tighten the region around the strongest-diffracting core. - """ - nonzero = arr[arr > 0] - if nonzero.size == 0: - return np.zeros_like(arr, dtype=bool), 0.0 - threshold = float(nonzero.mean() + n_sigma * nonzero.std()) - return arr > threshold, threshold - - -def crystal_mask_signal_percentile(arr: np.ndarray, percentile: float = 60.0) -> tuple[np.ndarray, float]: - """Threshold = given percentile of nonzero values. - - percentile=60 keeps the top 40 % of signal cells; raise to sharpen. - Unlike mean+sigma this is robust to heavy-tailed distributions. - """ - nonzero = arr[arr > 0] - if nonzero.size == 0: - return np.zeros_like(arr, dtype=bool), 0.0 - threshold = float(np.percentile(nonzero, percentile)) - return arr > threshold, threshold - - -def crystal_mask_dbscan(arr: np.ndarray, min_signal: float = 10.0, - eps: float = 1.5, min_samples: int = 3) -> tuple[np.ndarray, float]: - """DBSCAN spatial clustering on cells with signal > min_signal. - - Groups adjacent diffracting cells into clusters; the largest cluster - (by total signal weight) is labelled as the crystal. - eps=1.5 connects cells that are one grid step apart (including diagonal). - """ - from sklearn.cluster import DBSCAN - - ys, xs = np.where(arr > min_signal) - if len(ys) == 0: - return np.zeros_like(arr, dtype=bool), min_signal - - coords = np.column_stack([ys, xs]).astype(float) - labels = DBSCAN(eps=eps, min_samples=min_samples).fit_predict(coords) - - best_label, best_weight = -1, -1.0 - for lbl in set(labels): - if lbl == -1: - continue - weight = float(arr[ys[labels == lbl], xs[labels == lbl]].sum()) - if weight > best_weight: - best_label, best_weight = lbl, weight - - mask = np.zeros_like(arr, dtype=bool) - if best_label != -1: - sel = labels == best_label - mask[ys[sel], xs[sel]] = True - return mask, min_signal - - -def crystal_mask_kmeans(arr: np.ndarray, n_clusters: int = 3) -> tuple[np.ndarray, float]: - """K-means on signal values (1-D feature). - - Splits cells into n_clusters groups; the cluster with the highest - centroid is the crystal. n_clusters=3 separates background / fringe / crystal. - """ - from sklearn.cluster import KMeans - - flat = arr.flatten().reshape(-1, 1) - km = KMeans(n_clusters=n_clusters, random_state=0, n_init="auto").fit(flat) - crystal_label = int(np.argmax(km.cluster_centers_.flatten())) - threshold = float(sorted(km.cluster_centers_.flatten())[-2]) - mask = (km.labels_.reshape(arr.shape) == crystal_label) - return mask, threshold - - -CRYSTAL_METHOD_MAP: dict = { - "corner": ("Corner background", crystal_mask_corner_background), - "otsu": ("Otsu (nonzero)", crystal_mask_otsu_nonzero), - "mean_sigma": ("Mean + 0.5σ", lambda arr: crystal_mask_mean_sigma(arr, n_sigma=0.5)), - "percentile": ("Top-40% signal", lambda arr: crystal_mask_signal_percentile(arr, percentile=60.0)), - "dbscan": ("DBSCAN (spatial)", crystal_mask_dbscan), - "kmeans": ("K-means (k=3)", crystal_mask_kmeans), -} - - -def compare_crystal_methods( - results: List, - arr: Optional[np.ndarray] = None, - grid_size_mm: Optional[tuple[float, float]] = None, -) -> None: - """Plot each crystal-detection method side-by-side for visual comparison. - - Cyan contour = detected crystal region. - Red square = cell with maximum spots_low_res. - grid_size_mm = (step_x_mm, step_y_mm) — when provided, crystal size in µm - is calculated via get_xtal_size and shown in each panel title. - """ - import matplotlib.pyplot as plt - - if arr is None: - arr = rebuild_array_from_scan_results(results, "spots_low_res") - - methods = [ - ("Corner background", *crystal_mask_corner_background(arr)), - ("Otsu (nonzero)", *crystal_mask_otsu_nonzero(arr)), - ("Mean + 0.5σ", *crystal_mask_mean_sigma(arr, n_sigma=0.5)), - ("Top-40% signal", *crystal_mask_signal_percentile(arr, percentile=60.0)), - ("DBSCAN (spatial)", *crystal_mask_dbscan(arr)), - ("K-means (k=3)", *crystal_mask_kmeans(arr)), - ] - - fig, axes = plt.subplots(2, 3, figsize=(14, 9)) - for ax, (label, mask, threshold) in zip(axes.flat, methods): - _draw_panel(ax, arr, mask, label, threshold, grid_size_mm=grid_size_mm) - - fig.suptitle("Crystal region detection — method comparison\n" - "cyan = crystal outline | red square = max spots_low_res cell", - fontsize=10) - plt.tight_layout() - plt.show() - - -def compare_crystal_methods_scored( - results: List, - score_arr: Optional[np.ndarray] = None, - method: Optional[str] = None, - grid_size_mm: Optional[tuple[float, float]] = None, - w_bkg: float = 0.20, - w_low_res: float = 0.60, - w_indexed: float = 0.20, -) -> None: - """Plot each crystal-detection method applied to the composite 0–100 crystal score. - - The score combines bkg (w_bkg), spots_low_res (w_low_res), and spots_indexed - (w_indexed), each min-max normalised. All six methods are shown side-by-side. - - Args: - results: scan result list (used to build score if score_arr is None) - score_arr: pre-computed score array; computed from results when None - method: if given (one of 'corner','otsu','mean_sigma','percentile', - 'dbscan','kmeans'), that panel is highlighted with a yellow border - grid_size_mm: (step_x_mm, step_y_mm) for crystal-size annotation - w_bkg, w_low_res, w_indexed: composite score weights (must sum to 1.0) - """ - import matplotlib.pyplot as plt - - if score_arr is None: - score_arr = compute_crystal_score_array(results, w_bkg=w_bkg, w_low_res=w_low_res, w_indexed=w_indexed) - - method_entries = [ - (name, label, fn) - for name, (label, fn) in CRYSTAL_METHOD_MAP.items() - ] - - fig, axes = plt.subplots(2, 3, figsize=(15, 9)) - for ax, (name, label, fn) in zip(axes.flat, method_entries): - mask, threshold = fn(score_arr) - _draw_panel(ax, score_arr, mask, label, threshold, - grid_size_mm=grid_size_mm, cbar_label="crystal score (0–100)") - if method and name == method: - for spine in ax.spines.values(): - spine.set_edgecolor("yellow") - spine.set_linewidth(3) - - weight_note = f"bkg×{w_bkg:.0%} + spots_low_res×{w_low_res:.0%} + spots_indexed×{w_indexed:.0%}" - fig.suptitle( - f"Crystal region detection on composite score [{weight_note}]\n" - "cyan = crystal outline | red square = max-score cell | score range 0–100", - fontsize=10, - ) - plt.tight_layout() - plt.show() \ No newline at end of file diff --git a/src/aare/common/logger_config.py b/src/aare/common/logger_config.py deleted file mode 100644 index eced4277..00000000 --- a/src/aare/common/logger_config.py +++ /dev/null @@ -1,180 +0,0 @@ -import logging -import logging.config -import os -from pathlib import Path - -import yaml - - -class IgnoreSuccessfulStatusAccessFilter(logging.Filter): - # Should be disabled if debugging status calls. - # In particular if this is causing server/GUI lag - _SUPPRESSED_SUCCESS_MESSAGES = ( - '"GET /status HTTP/1.1" 200', - '"GET /sample/reference_tools HTTP/1.1" 200', - '"GET /sample/spreadsheet HTTP/1.1" 200', - '"GET /local_contact/device_state HTTP/1.1" 200', - '"POST /scan/smart_params HTTP/1.1" 200', - ) - - _SUPPRESSED_EXPECTED_403_MESSAGES = ( - '"GET /sse/face_detection HTTP/1.1" 403', - '"GET /sse/automation_progress HTTP/1.1" 403', - ) - - def filter(self, record: logging.LogRecord) -> bool: - message = record.getMessage() - return not any( - entry in message - for entry in ( - *self._SUPPRESSED_SUCCESS_MESSAGES, - *self._SUPPRESSED_EXPECTED_403_MESSAGES, - ) - ) - - -def get_uvicorn_logging_config() -> dict: - return { - "version": 1, - "disable_existing_loggers": False, - "filters": { - "ignore_successful_status_access": { - "()": "aare.common.logger_config.IgnoreSuccessfulStatusAccessFilter", - }, - }, - "formatters": { - "default": { - "()": "uvicorn.logging.DefaultFormatter", - "fmt": "%(levelprefix)s %(message)s", - "use_colors": True, - }, - }, - "handlers": { - "default": { - "formatter": "default", - "class": "logging.StreamHandler", - "stream": "ext://sys.stdout", - }, - "access": { - "formatter": "default", - "class": "logging.StreamHandler", - "stream": "ext://sys.stdout", - "filters": ["ignore_successful_status_access"], - }, - }, - "loggers": { - "uvicorn": { - "handlers": ["default"], - "level": "DEBUG", - "propagate": False, - }, - "uvicorn.access": { - "handlers": ["access"], - "level": "INFO", - "propagate": False, - }, - }, - } - - -def get_config_path(env: str = "dev") -> Path: - module_dir = Path(__file__).parent - config_filename = f"logging_{env}.yaml" - config_path = module_dir / "config" / config_filename - - if not config_path.exists(): - raise FileNotFoundError(f"Logging config not found: {config_path}") - - return config_path - - -# TODO fix logging. -def setup_logger( - name="aareDAQ", - base_dir: str | None = f"~/tmp/mxlogs/", - config_path: str | None = None, -): - # switch to production mode using: $ APP_ENV=prod python main.py - env = os.getenv("APP_ENV", "dev") # default: dev - - if config_path is None: - config_file = get_config_path(env) - else: - config_file = Path(config_path) - - with open(config_file, "r") as f: - config = yaml.safe_load(f.read()) - - # Force base directory to be under the user's home dir - effective_base = base_dir + f"{name}" or f"/tmp/logs/mxlogs/{name}" - effective_base = os.path.abspath( - os.path.expanduser(os.path.expandvars(effective_base)) - ) - - # Ensure directories for file handlers exist - handlers = config.get("handlers", {}) - - for h in handlers.values(): - filename = h.get("filename") - if not filename: - continue - abs_filename = os.path.join(effective_base, os.path.basename(filename)) - - h["filename"] = abs_filename - log_dir = os.path.dirname(abs_filename) - if log_dir and not os.path.exists(log_dir): - os.makedirs(log_dir, exist_ok=True) - - logging.config.dictConfig(config) - logging.getLogger("redis_lock").setLevel(logging.WARNING) - logging.getLogger("urllib3").setLevel(logging.WARNING) - logging.getLogger("matplotlib.font_manager").setLevel(logging.WARNING) - logging.getLogger("aaredaq").setLevel(logging.DEBUG) # DAQ - logging.getLogger("aaregui").setLevel(logging.INFO) - return logging.getLogger(name) - - -def find_existing_formatter( - default_fmt="%(asctime)s [%(levelname)s] %(name)s: %(message)s", - default_datefmt=None, -): - root = logging.getLogger() - for h in root.handlers: - if h.formatter: - return h.formatter - - for logger_name, logger in logging.Logger.manager.loggerDict.items(): - if isinstance(logger, logging.Logger): - for h in logger.handlers: - if h.formatter: - return h.formatter - - return logging.Formatter(default_fmt, datefmt=default_datefmt) - - -def attach_to_logger(logger_name: str = "", handler: logging.Handler = None): - logging.getLogger(logger_name).addHandler(handler) - - -if __name__ == "__main__": - os.environ["APP_ENV"] = "dev" # or "prod" - # os.environ["LOG_DIR"] = "/tmp/mxlogs" # optional; matches your setup code - - log = setup_logger( - "aareDAQ", base_dir="/tmp/test_logs" - ) # use the configured logger name - - log.debug("debug smoke test") - log.error("error smoke test") - log.info("info smoke test") - log.warning("warning smoke test") - log.critical("critical smoke test") - - for h in log.handlers: - if hasattr(h, "flush"): - h.flush() - - for h in log.handlers: - log.debug("debug test") - log.error("error test") - print(type(h), getattr(h, "baseFilename", None), h.level) diff --git a/src/aare/common/logger_events.py b/src/aare/common/logger_events.py deleted file mode 100644 index 6a1ee5a4..00000000 --- a/src/aare/common/logger_events.py +++ /dev/null @@ -1,158 +0,0 @@ -import functools -import time -from typing import Any, Callable -import logging - -from aare.common.raster_grid import RasterGridRequest -from aare.common.rotation_scan import RotationScanRequest - - -def log_timing( - logger: logging.Logger, - message_prefix: str = "", - level: int = logging.DEBUG, - extra: dict[str, Any] | None = None, -) -> Callable: - """Decorator to time a function call and log the duration.""" - def decorator(func: Callable) -> Callable: - @functools.wraps(func) - def wrapper(*args: Any, **kwargs: Any) -> Any: - start = time.perf_counter() - prefix = f"{message_prefix}: " if message_prefix else "" - logger.log(level, f"{prefix}Starting {func.__name__}") - try: - result = func(*args, **kwargs) - duration = time.perf_counter() - start - logger.log( - level, - f"{prefix}Finished {func.__name__} in {duration:.4f} seconds", - extra=merge_log_context(extra, duration_s=duration) - ) - return result - except Exception as e: - duration = time.perf_counter() - start - logger.log( - level, - f"{prefix}{func.__name__} FAILED after {duration:.4f} seconds with error: {e}", - extra=merge_log_context(extra, duration_s=duration) - ) - raise - return wrapper - return decorator - - -def merge_log_context(*parts: dict[str, Any] | None, **extra: Any) -> dict[str, Any]: - payload: dict[str, Any] = {} - for part in parts: - if part: - payload.update(part) - payload.update(extra) - return payload - - -def sample_log_context(sample) -> dict[str, Any]: - return { - "sample_id": getattr(sample, "db_id", None), - "sample_name": getattr(sample, "sample_name", None), - } - - -def raster_request_log_context(request: RasterGridRequest | None) -> dict[str, Any]: - if request is None: - return {} - - return { - "file_prefix": request.file_prefix, - "omega_deg": request.omega_deg, - "dtz": request.dtz, - "transmission": request.transmission, - "n_x": request.n_x, - "n_y": request.n_y, - "grid_size_x_mm": getattr(request.grid_size_mm, "x", None), - "grid_size_y_mm": getattr(request.grid_size_mm, "y", None), - } - - -def rotation_request_log_context( - request: RotationScanRequest | None, - *, - total_time_s: float | None = None, -) -> dict[str, Any]: - if request is None: - return {} - - payload = { - "file_prefix": request.file_prefix, - "start_omega_deg": request.start_omega_deg, - "dtz": request.dtz, - "exp_time_s": request.exp_time_s, - "incr_omega_deg": request.incr_omega_deg, - "steps": request.steps, - "transmission": request.transmission, - "screening": getattr(request, "screening", None), - } - - if total_time_s is not None: - payload["total_time_s"] = total_time_s - - return payload - - -def geom_log_context(geom, *, prefix: str = "") -> dict[str, Any]: - smargon = getattr(geom, "smargon", None) - sh_mm = getattr(smargon, "sh_mm", None) - - return { - f"{prefix}omega_deg": getattr(geom, "omega_deg", None), - f"{prefix}beam_x_pxl": getattr(getattr(geom, "beam_location_pxl", None), "x", None), - f"{prefix}beam_y_pxl": getattr(getattr(geom, "beam_location_pxl", None), "y", None), - f"{prefix}pixel_in_mm": getattr(geom, "pixel_in_mm", None), - f"{prefix}sh_x_mm": getattr(sh_mm, "x", None), - f"{prefix}sh_y_mm": getattr(sh_mm, "y", None), - f"{prefix}sh_z_mm": getattr(sh_mm, "z", None), - } - - -def ml_bundle_meta_log_context( - *, - target_point: tuple[float, float] | None = None, - focus: float | None = None, -) -> dict[str, Any]: - return { - "target_point": target_point, - "focus": focus, - } - - -def log_ml_bundle_meta( - logger: logging.Logger, - context: str, - *, - target_point: tuple[float, float] | None = None, - focus: float | None = None, -) -> None: - if target_point is None and focus is None: - return - - logger.debug( - "ML bundle metadata", - extra=merge_log_context( - {"context": context}, - ml_bundle_meta_log_context(target_point=target_point, focus=focus), - ), - ) - - -def log_duration( - logger: logging.Logger, - message: str, - duration_s: float, - *, - level: int = logging.INFO, - extra: dict[str, Any] | None = None, -) -> None: - logger.log( - level, - message, - extra=merge_log_context(extra, duration_s=duration_s), - ) \ No newline at end of file diff --git a/src/aare/common/models.py b/src/aare/common/models.py deleted file mode 100644 index 9d693471..00000000 --- a/src/aare/common/models.py +++ /dev/null @@ -1,689 +0,0 @@ -from pathlib import Path -import re -from enum import Enum -from typing import Annotated, Literal, Tuple, List, Optional -from dataclasses import dataclass -from pydantic import BaseModel, Field, field_validator, AfterValidator, ConfigDict, AliasChoices - -from aare.common.coordinate import Coordinate, positive_coords -from aare.common.diffraction_geometry import DiffractionGeometry -from aare.common.sample_geometry import SampleGeometryModel -from jfjoch_client.models.scan_result import ScanResult - -from aare.common.beamline import MXBeamline -from aare.common.tell_models import TellStateModel - - -class StagePositionEnum(Enum): - MEASURE = 0 - PARK = 1 - DOWN = 2 - UNKNOWN = 3 - - -class TokenData(BaseModel): - sub: str # Username - pgroups: List[str] - session: int - staff: bool = False - - -class DewarAddress(BaseModel): - segment: Literal["A", "B", "C", "D", "E", "F", "X", "R"] - pos: Annotated[int, Field(ge=1, le=5)] - - -class SampleDewarAddress(BaseModel): - puck: DewarAddress - pin: Annotated[int, Field(ge=1, le=16)] - - -# From the database for puck loading -class PuckInfo(BaseModel): - db_id: int - puck_name: str - dewar_name: str - user: str = "" - location: Optional[DewarAddress] = None - - -# From TELL to database after loading -class PuckLoadedInfo(BaseModel): - puck_name: str - location: DewarAddress - - -class DataCollectionParameters(BaseModel): - model_config = ConfigDict(from_attributes=True, populate_by_name=True) - - directory: Optional[str] = None - oscillation: Optional[float] = None # Only accept positive float - exposure: Optional[float] = None # Only accept positive floats between 0 and 1 - totalangle: Optional[int] = Field( # was totalrange - default=None, - validation_alias=AliasChoices('totalangle', 'totalrange') - ) # Only accept positive integers between 0 and 360 - transmission: Optional[int] = None # Only accept positive integers between 0 and 100 - targetresolution: Optional[float] = None # Only accept positive float - beamsize: Optional[str] = None - aperture: Optional[int] = None # Optional string field - datacollectiontype: Optional[str] = None # Only accept "standard", other types might be added later - processingpipeline: Optional[str] = "" # Only accept "gopy", "autoproc", "xia2dials" - spacegroupnumber: Optional[int] = None # Only accept positive integers between 1 and 230 - unitcell: Optional[str] = Field( # was cellparameters - default=None, - validation_alias=AliasChoices('unitcell', 'cellparameters') - ) # Must be a set of six positive floats or integers - rescutkey: Optional[str] = None # Only accept "is" or "cchalf" - rescutvalue: Optional[float] = None # Must be a positive float if rescutkey is provided - processingresolution: Optional[float] = Field( # was userresolution - default=None, - validation_alias=AliasChoices('processingresolution', 'userresolution') - ) - pdbid: Optional[str] = "" # Accepts either the format of the protein data bank code or {provided} - autoprocfull: Optional[bool] = None - procfull: Optional[bool] = None - adpenabled: Optional[bool] = None - noano: Optional[bool] = None - ffcscampaign: Optional[bool] = None - trustedhigh: Optional[float] = None # Should be a float between 0 and 2.0 - autoprocextraparams: Optional[str] = None # Optional string field - chiphiangles: Optional[float] = None # Optional float field between 0 and 30 - dose: Optional[float] = None # Optional float field - cloud: bool = True - pdbmodel: Optional[str] = None - - @field_validator("directory", mode="after") - @classmethod - def directory_characters(cls, v): - - # Default directory value if empty - if not v: # Handles None or empty cases - default_value = "{date}/{prefix}" - return default_value - - # Strip trailing slashes and store original value for comparison - v = str(v).strip("/") # Ensure it's a string and no trailing slashes - original_value = v - - # Replace spaces with underscores - v = v.replace(" ", "_") - - # Validate directory pattern with macros and allowed characters - valid_macros = [ - # Current macros - "{puck}", - "{position}", - "{prefix}", - "{date}", - "{run}", - "{beamline}", - # Legacy macros — accepted for back-compat, no longer documented - "{sgPuck}", - "{sgPosition}", - "{sgPrefix}", - "{sgPriority}", - "{protein}", - "{method}", - ] - valid_macro_pattern = re.compile( - "|".join(re.escape(macro) for macro in valid_macros) - ) - - # Check if the value contains valid macros - allowed_chars_pattern = "[a-z0-9_.+-/]" - v_without_macros = valid_macro_pattern.sub("macro", v) - - allowed_path_pattern = re.compile( - f"^(({allowed_chars_pattern}+|macro)*/*)*$", re.IGNORECASE - ) - if not allowed_path_pattern.match(v_without_macros): - raise ValueError( - f"'{v}' is not valid. Value must be a valid path or macro." - ) - return v - - @field_validator("unitcell", mode="before") - @classmethod - def unitcell_format(cls, v): - if v: - tokens = v.replace(",", " ").split() - try: - values = [float(i) for i in tokens] - except ValueError: - raise ValueError( - f"'{v}' is not valid. " - "Value must be six positive floats or integers (space or comma separated)." - ) - if len(values) != 6 or any(val <= 0 for val in values): - raise ValueError( - f"'{v}' is not valid. " - "Value must be six positive floats or integers (space or comma separated)." - ) - return v - - @field_validator("cloud", mode="before") - @classmethod - def coerce_cloud_default(cls, v): - if v in ("", None): - return True - if isinstance(v, bool): - return v - v_str = str(v).strip().lower() - if v_str in {"true", "yes", "1"}: - return True - if v_str in {"false", "no", "0"}: - return False - raise ValueError("cloud must be blank for default, or True/False") - - -# From database to TELL after loading -class SampleShortInfo(BaseModel): - db_id: int - puck_name: str - dewar_name: str - sample_name: str - run_number: int - aaredb_params: Optional[DataCollectionParameters] = None - user: str = "" - pin: Annotated[int, Field(ge=1, le=16)] - location: DewarAddress | None = None - priority: Optional[float] = 1.0 - comment: Optional[str] = None - mount_count: int = 0 - rotation_count: int = 0 - raster_count: int = 0 - screening_count: int = 0 - - def tell_address(self) -> SampleDewarAddress: - return SampleDewarAddress(puck=self.location, pin=self.pin) - - def loc_str(self) -> str: - if self.location is None: - return "-" - else: - return f"{self.location.segment}{self.location.pos}-{self.pin}" - - def loc_str_sort(self) -> str: - if self.location is None: - return "" - else: - return f"{self.location.segment}{self.location.pos}-{self.pin:02d}" - - @classmethod - def from_dict(cls, data: dict): - return cls(**data) - - -class SampleShortInfoList(BaseModel): - s: List[SampleShortInfo] - - -class BeamMarkCoeffModel(BaseModel): - """ - Model to calculate beam center for a given zoom level based on a quadratic approximation - For both x and y it contains tuple of coefficients a, b, c (a*zoom^2 + b*zoom + c) - Default is beam center in 1000,1000 at any zoom level - """ - - coeff_x: Tuple[float, float, float] = (0, 0, 1000.0) - coeff_y: Tuple[float, float, float] = (0, 0, 1000.0) - - def apply(self, zoom: float): - return Coordinate( - x=self.coeff_x[0] * zoom**2 + self.coeff_x[1] * zoom + self.coeff_x[2], - y=self.coeff_y[0] * zoom**2 + self.coeff_y[1] * zoom + self.coeff_y[2], - ) - - -class FluorescenceSpectrumParameterModel(BaseModel): - erase: bool = True - acq_time_s: float - transmission: Annotated[float, Field(ge=0,le=1)] | None = None - - -class FluorescenceSpectrumOutputModel(BaseModel): - bkg: list[float] | None = None - spectrum: list[float] - energy_eV: list[float] - average_dead_time: Annotated[float, Field(ge=0,le=1)] - - -class MLBoxType(Enum): - LOOP_ALL = 0 - PIN = 1 - CRYSTAL = 2 - LOOP_FACE = 3 - ICE = 4 - NEEDLE = 5 - - -class BoundingBoxModel(BaseModel): - top_x: float - top_y: float - bottom_x: float - bottom_y: float - - -class MLBoxModel(BaseModel): - cls: MLBoxType - box: BoundingBoxModel - conf: float - - @staticmethod - def from_tuple(cls:MLBoxType, box_tuple: tuple[float, float, float, float], conf: float) -> "MLBoxModel": - x1, y1, x2, y2 = box_tuple - return MLBoxModel( - cls=cls, - box=BoundingBoxModel(top_x=float(x1), top_y=float(y1), bottom_x=float(x2), bottom_y=float(y2)), - conf=float(conf) - ) - - -class MLOutputModel(BaseModel): - boxes: dict[str, MLBoxModel] = {} - - @staticmethod - def get_class_str(cls: MLBoxType) -> str: - if cls == MLBoxType.LOOP_ALL: - return "Loop_all" - if cls == MLBoxType.PIN: - return "Pin" - if cls == MLBoxType.CRYSTAL: - return "Crystal" - if cls == MLBoxType.LOOP_FACE: - return "Loop_face" - if cls == MLBoxType.ICE: - return "Ice" - if cls == MLBoxType.NEEDLE: - return "Needle" - return "Unknown" - - def _next_unique_key(self, base: str) -> str: - if base not in self.boxes: - return base - i = 2 - while f"{base}_{i}" in self.boxes: - i += 1 - return f"{base}_{i}" - - def add_box(self, cls: MLBoxType, box_tuple: tuple[float, float, float, float], - conf: float | None = None) -> str: - key_base = self.get_class_str(cls) - key = self._next_unique_key(key_base) - self.boxes[key] = MLBoxModel.from_tuple(cls, box_tuple, conf) - return key - - def get_box_model(self, key: str) -> MLBoxModel | None: - return self.boxes.get(key) - - def get_box_tuple(self, key: str) -> tuple[float, float, float, float] | None: - m = self.get_box_model(key) - if not m or not m.box: - return None - return (m.box.top_x, m.box.top_y, m.box.bottom_x, m.box.bottom_y) - - def get_box_tuple_with_conf(self, key: str) -> tuple[float, float, float, float, float] | None: - m = self.get_box_model(key) - if not m or not m.box: - return None - conf = float(m.conf) if m.conf is not None else 0.0 - return (m.box.top_x, m.box.top_y, m.box.bottom_x, m.box.bottom_y, conf) - - def get_keys_for_class(self, cls: MLBoxType) -> list[str]: - base = self.get_class_str(cls) - return [k for k, v in self.boxes.items() if v.cls == cls and (k == base or k.startswith(base + "_"))] - - def get_models_for_class(self, cls: MLBoxType) -> list[MLBoxModel]: - keys = self.get_keys_for_class(cls) - return [self.boxes[k] for k in keys] - - def get_tuples_for_class(self, cls: MLBoxType) -> list[tuple[float, float, float, float]]: - out: list[tuple[float, float, float, float]] = [] - for m in self.get_models_for_class(cls): - if m.box: - out.append((m.box.top_x, m.box.top_y, m.box.bottom_x, m.box.bottom_y)) - return out - - def get_tuples_with_conf_for_class(self, cls: MLBoxType) -> list[tuple[float, float, float, float, float]]: - out: list[tuple[float, float, float, float, float]] = [] - for m in self.get_models_for_class(cls): - if m.box: - out.append((m.box.top_x, m.box.top_y, m.box.bottom_x, m.box.bottom_y, float(m.conf))) - return out - - def get_best_for_class(self, cls: MLBoxType) -> MLBoxModel | None: - models = self.get_models_for_class(cls) - if not models: - return None - return max(models, key=lambda m: (m.conf if m.conf is not None else 0.0)) - - -class BeamlineStateEnum(Enum): - Maintenance = 1 - SampleExchange = 2 - SampleAlignment = 3 - DataCollection = 4 - DewarTransfer = 5 - XrayFluorescence = 6 - BeamLocation = 7 - Moving = 8 - RobotSampleExchange = 9 - XtalSnapshot = 10 - BeamstopAlignment = 11 - FluxMeasurement = 12 - - def display_name(self) -> str: - return { - BeamlineStateEnum.Maintenance: "Maintenance", - BeamlineStateEnum.SampleExchange: "Sample exchange", - BeamlineStateEnum.SampleAlignment: "Sample alignment", - BeamlineStateEnum.DataCollection: "Data collection", - BeamlineStateEnum.DewarTransfer: "Dewar transfer", - BeamlineStateEnum.XrayFluorescence: "X-ray fluorescence", - BeamlineStateEnum.BeamLocation: "Beam location", - BeamlineStateEnum.Moving: "Moving", - BeamlineStateEnum.RobotSampleExchange: "Robot sample exchange", - BeamlineStateEnum.XtalSnapshot: "Xtal snapshot", - BeamlineStateEnum.BeamstopAlignment: "Beamstop alignment", - BeamlineStateEnum.FluxMeasurement: "Flux measurement", - }.get(self, "-") - - -class DAQOperation(str, Enum): - AUTOMATION = "automation" - MOUNT = "mount" - UNMOUNT = "unmount" - LOOP_CENTERING = "loop_centering" - FACE_CENTERING = "face_centering" - RASTER = "raster" - ROTATION = "rotation" - MEASURE = "measure" - - -class SessionsStateEnum(Enum): - Vacant = 0 - OwnedByYou = 1 - OwnedByElse = 2 - PendingYouToElse = 3 - PendingElseToYou = 4 - - -class SampleCameraSettings(BaseModel): - gain: float - exposure: float - - -class ZoomModeEnum(Enum): - User = 1 - LoopCenter = 2 - BeamLocation = 3 - - -class ZoomModel(BaseModel): - z:dict[float, SampleCameraSettings] - - def get_camera_settings(self, zoom_value: float) -> SampleCameraSettings | tuple[float, float]: - if not self.z: - raise ValueError("No zoom data available") - - if zoom_value in self.z: - elem = self.z[zoom_value] - return SampleCameraSettings(gain=elem.gain, exposure=elem.exposure) - - sorted_zooms = sorted(self.z.keys()) - - if zoom_value <= sorted_zooms[0]: - elem = self.z[sorted_zooms[0]] - return SampleCameraSettings(gain=elem.gain, exposure=elem.exposure) - - if zoom_value >= sorted_zooms[-1]: - elem = self.z[sorted_zooms[-1]] - return SampleCameraSettings(gain=elem.gain, exposure=elem.exposure) - print('interpolating zoom') - for i in range(len(sorted_zooms) - 1): - if sorted_zooms[i] <= zoom_value <= sorted_zooms[i + 1]: - lower_zoom = sorted_zooms[i] - upper_zoom = sorted_zooms[i + 1] - - lower_elem = self.z[lower_zoom] - upper_elem = self.z[upper_zoom] - - t = (zoom_value - lower_zoom) / (upper_zoom - lower_zoom) - - interpolated_gain = lower_elem.gain + t * (upper_elem.gain - lower_elem.gain) - interpolated_exp = lower_elem.exposure + t * (upper_elem.exposure - lower_elem.exposure) - - return SampleCameraSettings(gain=interpolated_gain, exposure=interpolated_exp) - - closest = min(self.z.keys(), key=lambda x: abs(x - zoom_value)) - elem = self.z[closest] - return SampleCameraSettings(gain=elem.gain, exposure=elem.exposure) - -def def_zoom(beamline) -> ZoomModel: - print(f'user zoom for {beamline}') - if beamline == MXBeamline.X06DA: - return ZoomModel( - z={ - 1: SampleCameraSettings(gain=0, exposure=0.05), - 280: SampleCameraSettings(gain=0, exposure=0.05), - 500: SampleCameraSettings(gain=0, exposure=0.1), - 700: SampleCameraSettings(gain=0, exposure=0.15), - 800: SampleCameraSettings(gain=0, exposure=0.2), - 1000: SampleCameraSettings(gain=0, exposure=0.25) - } - ) - elif beamline == MXBeamline.X10SA: - return ZoomModel(z={1: SampleCameraSettings(gain=0, exposure=0.002)}) - elif beamline == MXBeamline.X06SA: - return ZoomModel(z={1: SampleCameraSettings(gain=0, exposure=0.002)}) - elif beamline == MXBeamline.SIMULATED: - return ZoomModel(z={1: SampleCameraSettings(gain=0, exposure=0.002)}) - else: - raise ValueError(f"Invalid beamline: {beamline}") - -def def_bl_zoom(beamline) -> ZoomModel: - if beamline == MXBeamline.X06DA: - # Sensible starting presets so the beam is visible at every zoom; these - # are tuned live and persisted to Redis from the GUI (config.zoom_settings). - return ZoomModel( - z={ - 1: SampleCameraSettings(gain=0, exposure=0.05), - 280: SampleCameraSettings(gain=0, exposure=0.05), - 500: SampleCameraSettings(gain=0, exposure=0.1), - 700: SampleCameraSettings(gain=0, exposure=0.15), - 800: SampleCameraSettings(gain=0, exposure=0.2), - 1000: SampleCameraSettings(gain=0, exposure=0.25) - } - ) - elif beamline == MXBeamline.X10SA: - return ZoomModel(z={1: SampleCameraSettings(gain=0, exposure=0.002)}) - elif beamline == MXBeamline.X06SA: - return ZoomModel(z={1: SampleCameraSettings(gain=0, exposure=0.002)}) - elif beamline == MXBeamline.SIMULATED: - return ZoomModel(z={1: SampleCameraSettings(gain=0, exposure=0.002)}) - else: - raise ValueError(f"Invalid beamline: {beamline}") - -def def_loop_centering_zoom(beamline) -> ZoomModel: - if beamline == MXBeamline.X06DA: - return ZoomModel(z={1: SampleCameraSettings(gain=0, exposure=0.05)}) - elif beamline == MXBeamline.X10SA: - return ZoomModel(z={1: SampleCameraSettings(gain=0, exposure=0.002), - 280: SampleCameraSettings(gain=0, exposure=0.002)}) - elif beamline == MXBeamline.X06SA: - return ZoomModel(z={1: SampleCameraSettings(gain=0, exposure=0.05)}) - elif beamline == MXBeamline.SIMULATED: - return ZoomModel(z={1: SampleCameraSettings(gain=0, exposure=0.05)}) - else: - raise ValueError(f"Invalid beamline: {beamline}") - -def zoom_manager(mode: ZoomModeEnum = ZoomModeEnum.User, beamline: MXBeamline = None) -> ZoomModel | None: - if beamline is None or beamline == MXBeamline.SIMULATED: - print('SIMULATED') - return def_zoom(beamline) - if mode == ZoomModeEnum.BeamLocation: - return def_bl_zoom(beamline) - elif mode == ZoomModeEnum.LoopCenter: - return def_loop_centering_zoom(beamline) - elif mode == ZoomModeEnum.User: - return def_zoom(beamline) - else: - raise ValueError(f"Invalid zoom mode: {mode}") - - -class AutofocusSettings(BaseModel): - center_x_pxl: float | None # Use beam center - center_y_pxl: float | None # Use beam center - radius_pxl: float - z_range_um: float - z_steps: int - - -class BeamlineStatus(BaseModel): - name: str - ring_current_mA: float - front_light: Annotated[float, Field(ge=0.0, le=100.0)] - back_light: Annotated[float, Field(ge=0.0, le=100.0)] - cryojet_K: float - shutter_open: bool - exp_shutter_open: bool | None - flux_ph_s: float - sample_camera: SampleCameraSettings - transmission: Annotated[float, Field(ge=0.0, le=1.0)] | None - zoom: float - commissioning_mode: bool - dtz_min: float - dtz_max: float - # Hutch personnel-safety system state. ``pss_prohibited`` is True when the - # hutch is interlocked so the robot may move (PROHIBITED-STATE); the GUI - # blocks a mount when it is False. ``pss_alarm`` is True when ALARM-STATE - # != 0 (warning). Defaults keep older payloads/constructors valid and avoid - # the GUI false-blocking when an old server omits the field. - pss_prohibited: bool = True - pss_alarm: bool = False - - -class SessionStatus(BaseModel): - session: SessionsStateEnum = SessionsStateEnum.Vacant - current_pgroup: str | None = None - staff: bool = False - - -class OpenGuiSessionInfo(BaseModel): - session: int - username: str - staff: bool = False - last_seen_ts: float - last_interaction_ts: float | None = None - close_requested: bool = False - close_requested_by: str | None = None - close_requested_at: float | None = None - close_grace_seconds: int | None = None - holds_baton: bool = False - - -class CrystalSize(BaseModel): - x: float = 0.0 - y: float = 0.0 - z: float = 0.0 - - -class DAQStatusModel(BaseModel): - geom: SampleGeometryModel - diffraction: DiffractionGeometry - bl: BeamlineStatus - state: BeamlineStateEnum - busy: bool - sample: SampleShortInfo | None = None - session: SessionStatus - open_guis: list[OpenGuiSessionInfo] = [] - box: BoundingBoxModel | None = None - last_best_res: float | None = None - last_best_b_factor: float | None = None - crystal_size: CrystalSize = CrystalSize(x=0,y=0,z=0) - - tell_connected: bool = True - tell_error: str | None = None - tell_state: TellStateModel | None = None - - 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 - dtz_collection: float | None = 130.0 - dtz_park: float | None = 150.0 - dtz_wash_sample_distance: float | None = 150.0 - dtz_bsz_safety_margin: float | None = 50.0 - bsz: float | None = 25.0 - - camera_max_magnification: float | None = 1.0 - camera_min_magnification: float | None = 500.0 - camera_translation_factor_a: float | None = 0.00253 - camera_translation_factor_b: float | None = 512.0 - - -class CryojetSettingsModel(BaseModel): - cryojet_park_position: float | None = 12.0 - cryojet_measurement_position: float | None = 5.0 - cryojet_in_use: bool | None = True - - -class SimpleStrategyInputModel(BaseModel): - last_best_res: float | None = None - last_best_b_factor: float | None = None - crystal_size: CrystalSize = CrystalSize(x=0, y=0, z=0) - angular_range: int = 360 - start_angle: float = 0 - incr_omega_deg:float = 0.2 - d_vis: float = 1.5 - current_temp_k: float = 100 - filename: Optional[str] = None - - -class SimpleScanParameters(BaseModel): - filename: Optional[str] = None - dtz: float = 150 - exp_time_s: float = 0.02 - start_omega_deg:float = 0 - incr_omega_deg: float = 0.2 - steps: int = 1800 - transmission: Annotated[float, Field(ge=0.0, le=1.0)] = 1.0 - last_best_res: float | None = None - last_best_b_factor: float | None = None - crystal_size: CrystalSize = CrystalSize(x=0, y=0, z=0) - flux_ph_s: float | None = None - calculated_dose_Mgy: float | None = None - calculated_dose_rate_MGy_s: float | None = None - xtal_size_dose_rate_MGy_s: float | None = None - target_dose_MGy: float | None = None - beam_size_x_um: float | None = None - beam_size_y_um: float | None = None - dose_rate_MGy_s: float | None = None - d_vis: float | None = None - d_tar: float | None = None - - -class ScanResultPayloadModel(BaseModel): - result: ScanResult - sample_id: int - attach_image: bool = True - beam_mark_pxl: tuple[float, float] - beam_size_mm: Annotated[Coordinate, AfterValidator(positive_coords)] - - -class RecoveryActionRequest(BaseModel): - confirmation_code: str - - -@dataclass -class LoopCenteringResult: - success: bool - comment: str | None = None - error: Exception | None = None \ No newline at end of file diff --git a/src/aare/common/raster_grid.py b/src/aare/common/raster_grid.py deleted file mode 100644 index 67c6c559..00000000 --- a/src/aare/common/raster_grid.py +++ /dev/null @@ -1,140 +0,0 @@ -from typing import Annotated, List - -import numpy as np - -from jfjoch_client.models.scan_result import ScanResult -from pydantic import Field, AfterValidator, BaseModel - -from aare.common.coordinate import Coordinate, positive_coords, SmargonCoordinate -from aare.common.sample_geometry import SampleGeometryModel - - -class RasterGridRequest(BaseModel): - dtz: float | None = None - transmission: float | None = None - file_prefix: str | None = None - exp_time_s: float - - # Number of grid elements - 0 is allowed for empty grid - n_x: Annotated[int, Field(ge=0)] - n_y: Annotated[int, Field(ge=0)] - - # Size of grid elements - grid_size_mm: Annotated[Coordinate, AfterValidator(positive_coords)] - - # Coordinates of top left corner in Smargon (SH...) coordinates. - # If None, measure from the current position. - smargon_top_left: SmargonCoordinate | None - - - # omega angle - omega_deg: float = 0.0 - - visible: bool = True - - def get_image_number(self) -> int: - return self.n_x * self.n_y - - def grid_size_pxl(self, geom: SampleGeometryModel) -> Coordinate: - return self.grid_size_mm / geom.pixel_in_mm - - def start_pxl(self, geom: SampleGeometryModel) -> Coordinate: - if self.smargon_top_left is None: - return geom.smargon_to_beamline(geom.smargon.sh_mm) - return geom.smargon_to_beamline(self.smargon_top_left.sh_mm) - -class CenterOfMassModel(BaseModel): - n_x: float - n_y: float - max_image: int | None = None - - @classmethod - def model_validate_maybe(cls, n_x: float, n_y: float) -> "CenterOfMassModel | None": - # Return None if either coordinate is NaN - if np.isnan(n_x) or np.isnan(n_y): - return None - return cls(n_x=float(n_x), n_y=float(n_y)) - - def get_com_mm(self, r:RasterGridRequest) -> Coordinate: - #defines COM in mm at beam position, relative to the top left corner of grid - return Coordinate(x=(self.n_x+0.5)*r.grid_size_mm.x, y=(self.n_y+0.5)*r.grid_size_mm.y) - - def get_com_pxl(self, r:RasterGridRequest, geom:SampleGeometryModel) -> Coordinate: - #defines COM in pixels at beam position, relative to the top left corner of grid - com_mm = self.get_com_mm(r) - return Coordinate(x=com_mm.x/geom.pixel_in_mm, y=com_mm.y/geom.pixel_in_mm) - - def com_nan_check(self): - # True if both values are finite (not NaN) - return not (np.isnan(self.n_x) or np.isnan(self.n_y)) - -class CompletedRasterGridElem(BaseModel): - request: RasterGridRequest - result: ScanResult - centre_of_mass: CenterOfMassModel | None - -class CompletedRasterGrid(BaseModel): - r: List[CompletedRasterGridElem] - -class RasterPayloadModel(BaseModel): - request: RasterGridRequest - result: ScanResult - sample_id: int - attach_image: bool = True - centre_of_mass: CenterOfMassModel | None = None - raster_score: List[float | None] | None = None - center_pxl: Coordinate | None - start_pxl: Coordinate - cell_size_pxl: Annotated[Coordinate, AfterValidator(positive_coords)] - beam_mark_pxl: tuple[float, float] - beam_size_mm: Annotated[Coordinate, AfterValidator(positive_coords)] - -def grid_to_image_id(grid_x: int, grid_y: int, number_of_cols: int) -> int: - """ - Get the associated image id for a given grid position (grid_x, grid_y) - using serpentine row ordering. - - grid_x: 1-based column index - grid_y: 1-based row index - number_of_cols: total number of columns in the grid - """ - if number_of_cols < 1: - raise ValueError("number_of_cols must be >= 1") - if grid_x < 1: - raise ValueError("grid_x must be >= 1") - if grid_y < 1: - raise ValueError("grid_y must be >= 1") - if grid_x > number_of_cols: - raise ValueError("grid_x cannot be greater than number_of_cols") - - col = grid_x - 1 - row = grid_y - 1 - - if row % 2 == 0: # left -> right - return row * number_of_cols + col - else: # right -> left - return row * number_of_cols + (number_of_cols - 1 - col) - - -def image_id_to_grid(image_id: int, number_of_cols: int) -> tuple[int, int]: - """ - Get the associated grid position (grid_x, grid_y) for a given image id - using serpentine row ordering. - """ - if number_of_cols < 1: - raise ValueError("number_of_cols must be >= 1") - if image_id < 0: - raise ValueError("image_id must be >= 0") - - row = image_id // number_of_cols - pos_in_row = image_id % number_of_cols - - if row % 2 == 0: # left -> right - col = pos_in_row - else: # right -> left - col = number_of_cols - 1 - pos_in_row - - grid_x = col + 1 - grid_y = row + 1 - - return grid_x, grid_y \ No newline at end of file diff --git a/src/aare/common/recurrence_watcher.py b/src/aare/common/recurrence_watcher.py deleted file mode 100644 index 05240849..00000000 --- a/src/aare/common/recurrence_watcher.py +++ /dev/null @@ -1,136 +0,0 @@ -from __future__ import annotations - -from dataclasses import dataclass -from typing import Callable, TypeAlias - -from aare.common import exception_handler - -from aare.common.exception_handler import ( - AareException, - LoopCenteringFailed, - SmargonException, - TellException, - TransformationInvalidException, -) - -ExceptionType: TypeAlias = type[AareException] - - -@dataclass(frozen=True) -class WatcherTrip: - name: str - threshold: int - streak: int - exception_class: str - - def message(self) -> str: - return ( - f"{self.streak} consecutive {self.name} errors across samples " - f"({self.exception_class}) — automation halted" - ) - - -class RecurrenceWatcher: - def __init__(self, *, name: str, observes: ExceptionType, threshold: int): - self.name = name - self._observes = observes - self.threshold = max(1, int(threshold)) - self._streak = 0 - - @property - def streak(self) -> int: - return self._streak - - def reset(self) -> None: - self._streak = 0 - - def observe(self, exception_class: type | None) -> bool: - if exception_class is None: - self._streak = 0 - return False - if not isinstance(exception_class, type) or not issubclass(exception_class, self._observes): - return False - self._streak += 1 - return self._streak >= self.threshold - - def maybe_trip(self, exception_class: type | None) -> WatcherTrip | None: - if not self.observe(exception_class): - return None - class_name = exception_class.__name__ if isinstance(exception_class, type) else "Unknown" - return WatcherTrip( - name=self.name, - threshold=self.threshold, - streak=self._streak, - exception_class=class_name, - ) - - -DEFAULT_WATCHERS: tuple[tuple[str, ExceptionType, int], ...] = ( - ("tell", TellException, 5), - ("alc", LoopCenteringFailed, 3), - ("transformation", TransformationInvalidException, 3), - ("smargon", SmargonException, 3), -) - - -def create_default_watchers( - overrides: dict[str, int] | None = None, -) -> list[RecurrenceWatcher]: - thresholds = dict(overrides or {}) - return [ - RecurrenceWatcher( - name=name, - observes=observes, - threshold=thresholds.get(name, threshold), - ) - for name, observes, threshold in DEFAULT_WATCHERS - ] - - -def redis_key_to_env_var(key: str) -> str: - """Adapter from canonical Redis key shape to an env-var-safe name. - - ``aare:watchers:{bl}:{watcher}:threshold`` → - ``AARE_WATCHERS_{BL}_{WATCHER}_THRESHOLD``. - - Used by the GUI today since it has no Redis client; the loader's - ``get_value`` argument keeps the Redis key shape authoritative so the - DAQ side can later swap in ``cfg.get`` without churn (plan §7).""" - return key.replace(":", "_").upper() - - -def load_watcher_threshold_overrides( - *, - beamline: str, - watcher_names: list[str], - get_value: Callable[[str], str | bytes | None], -) -> dict[str, int]: - overrides: dict[str, int] = {} - for watcher_name in watcher_names: - key = f"aare:watchers:{beamline}:{watcher_name}:threshold" - raw_value = get_value(key) - if raw_value is None: - continue - if isinstance(raw_value, bytes): - raw_value = raw_value.decode("utf-8", errors="ignore") - try: - parsed = int(str(raw_value).strip()) - except (TypeError, ValueError): - continue - if parsed > 0: - overrides[watcher_name] = parsed - return overrides - - -_EXCEPTION_CLASS_BY_NAME: dict[str, type] = { - name: klass - for name, klass in vars(exception_handler).items() - if isinstance(klass, type) - and issubclass(klass, AareException) -} - - -def resolve_exception_class(exception_class_name: str | None) -> type | None: - if not exception_class_name: - return None - return _EXCEPTION_CLASS_BY_NAME.get(exception_class_name) \ No newline at end of file diff --git a/src/aare/common/rotation_scan.py b/src/aare/common/rotation_scan.py deleted file mode 100644 index 471e7a9a..00000000 --- a/src/aare/common/rotation_scan.py +++ /dev/null @@ -1,25 +0,0 @@ -from jfjoch_client.models.scan_result import ScanResult -from pydantic import BaseModel - -from aare.common.coordinate import SmargonCoordinate - - -class RotationScanRequest(BaseModel): - dtz: float | None = None - transmission: float | None = None - file_prefix: str | None = None - exp_time_s: float - - start: SmargonCoordinate | None = None # None = start where you are - end: SmargonCoordinate | None = None - - start_omega_deg: float = 0 - incr_omega_deg: float - steps: int - - screening: bool = False - wedge_omega_deg: float | None = None - -class CompletedRotationScan(BaseModel): - request: RotationScanRequest - result: ScanResult \ No newline at end of file diff --git a/src/aare/common/sample_geometry.py b/src/aare/common/sample_geometry.py deleted file mode 100644 index c9a4a473..00000000 --- a/src/aare/common/sample_geometry.py +++ /dev/null @@ -1,97 +0,0 @@ -from typing import Annotated - -import numpy as np -from pydantic import BaseModel, Field, AfterValidator - -from aare.common.coordinate import Coordinate, SmargonCoordinate, positive_coords - - -class SampleGeometryModel(BaseModel): - # Beam location in the camera coordinates - beam_location_pxl: Coordinate - - # 1 pixel in mm - pixel_in_mm: Annotated[float, Field(gt=0.0)] - - aerotech: Coordinate - - aerotech_meas: Coordinate - # Current smargon coordinates/target - smargon: SmargonCoordinate - - omega_deg: float - - # Beam size in mm - beam_size_mm: Annotated[Coordinate, AfterValidator(positive_coords)] - - @property - def beam_size_pxl(self) -> Coordinate: - return self.beam_size_mm / self.pixel_in_mm - - - def smargon_z(self, pic_y: float) -> Coordinate: - y_mm = (pic_y - self.beam_location_pxl.y) * self.pixel_in_mm - return self.smargon.sh_mm + self.smargon_nudge(Coordinate(z=y_mm)) - - # This goes from image coordinates to relative beam coordinates in image plane - def picture_to_sample(self, pic_pxl: Coordinate) -> Coordinate: - # Relative to beam location - beam_mm = (pic_pxl - self.beam_location_pxl) * self.pixel_in_mm + self.aerotech - return beam_mm - - # This goes from relative beam coordinates in image plane to image - def sample_to_picture(self, beam_mm: Coordinate) -> Coordinate: - beam_pxl = (beam_mm - self.aerotech) / self.pixel_in_mm - return beam_pxl + self.beam_location_pxl - - def translate_smargon(self, coord: Coordinate) -> SmargonCoordinate: - return SmargonCoordinate(sh_mm=self.smargon.sh_mm + self.smargon_nudge(coord), - phi_deg=self.smargon.phi_deg, - chi_deg=self.smargon.chi_deg) - - def smargon_nudge(self, coord: Coordinate) -> Coordinate: - phi = np.radians(np.around(self.smargon.phi_deg, decimals=1)) - chi = np.radians(np.around(self.smargon.chi_deg, decimals=1)) - omega = np.radians(np.around(self.omega_deg, decimals=1)) - - offset = Coordinate() - co = np.cos(omega) - so = np.sin(omega) - cp = np.cos(phi) - sp = np.sin(phi) - cc = np.cos(chi) - sc = np.sin(chi) - - offset.x = ( - -coord.z * co * sp - - coord.y * so * sp - + coord.x * cp * sc - - coord.y * co * cc * cp - + coord.z * so * cc * cp - ) - offset.y = ( - -coord.z * co * cp - - coord.y * so * cp - - coord.x * sc * sp - + coord.y * co * cc * sp - - coord.z * so * cc * sp - ) - offset.z = -coord.x * cc - coord.y * co * sc + coord.z * so * sc - return offset - - def beamline_to_smargon(self, coord: Coordinate) -> Coordinate: - return self.smargon.sh_mm + self.smargon_nudge(coord) - - def smargon_to_beamline(self, coord: Coordinate) -> Coordinate: - rel = coord - self.smargon.sh_mm - - vec_x = self.smargon_nudge(Coordinate(x=1)).normalize() - vec_y = self.smargon_nudge(Coordinate(y=1)).normalize() - - return Coordinate(x=rel * vec_x, y=rel * vec_y) - - def picture_to_smargon(self, coord: Coordinate) -> Coordinate: - return self.beamline_to_smargon(self.picture_to_sample(coord)) - - def smargon_to_picture(self, coord: Coordinate) -> Coordinate: - return self.sample_to_picture(self.smargon_to_beamline(coord)) diff --git a/src/aare/common/simulate_raster.py b/src/aare/common/simulate_raster.py deleted file mode 100644 index f69b613c..00000000 --- a/src/aare/common/simulate_raster.py +++ /dev/null @@ -1,217 +0,0 @@ -""" -Synthetic raster scan data for no-beam / offline testing. - -Entry point ------------ -generate_no_beam_scan_result(request, seed=None) -> ScanResult - -The returned ScanResult is structurally identical to a real one, so all -find_xtal.py functions consume it without modification. - -Pre-defined seeds ------------------ -Each seed fixes cluster geometry so results are reproducible. The six -scenarios below cover the main cases needed for unit-testing crystal-finding -algorithms. - -Seed Scenario ----- -------- -1 Single crystal, centred – baseline positive detection -2 Single crystal, off-centre – tests COM accuracy near an edge -3 Two well-separated crystals – multi-crystal detection -4 Two overlapping crystals – segmentation challenge -5 Two clusters at ~90 ° – twinned / differently oriented crystal -6 Pure background only – true negative (no crystal) -""" - -from __future__ import annotations - -from dataclasses import dataclass, field -from typing import List - -import numpy as np -from jfjoch_client import ScanResult -from jfjoch_client.models.scan_result_images_inner import ScanResultImagesInner - -from aare.common.raster_grid import RasterGridRequest - -# --------------------------------------------------------------------------- -# Cluster geometry descriptor -# --------------------------------------------------------------------------- - -@dataclass -class _ClusterParams: - """Fractional coordinates and shape of one elliptical crystal cluster. - - cx, cy – cluster centre as a fraction of (n_x, n_y) [0..1] - ax, ay – semi-axes as a fraction of (n_x, n_y) - theta – rotation of the ellipse in radians - peak – peak spots_low_res at cluster centre (integer) - """ - cx: float - cy: float - ax: float - ay: float - theta: float - peak: int - - -@dataclass -class _SeedConfig: - clusters: List[_ClusterParams] = field(default_factory=list) - - -# --------------------------------------------------------------------------- -# Pre-defined seed catalogue -# --------------------------------------------------------------------------- - -SEED_CATALOGUE: dict[int, _SeedConfig] = { - # --- 1: single crystal, centred, compact, strong signal ---------------- - 1: _SeedConfig(clusters=[ - _ClusterParams(cx=0.50, cy=0.50, ax=0.15, ay=0.15, theta=0.0, peak=120), - ]), - - # --- 2: single crystal, off-centre, slightly elongated ----------------- - 2: _SeedConfig(clusters=[ - _ClusterParams(cx=0.25, cy=0.70, ax=0.12, ay=0.18, theta=0.35, peak=90), - ]), - - # --- 3: two well-separated crystals ------------------------------------ - 3: _SeedConfig(clusters=[ - _ClusterParams(cx=0.20, cy=0.20, ax=0.12, ay=0.12, theta=0.0, peak=100), - _ClusterParams(cx=0.75, cy=0.72, ax=0.15, ay=0.10, theta=0.2, peak=80), - ]), - - # --- 4: two overlapping crystals (centres ~1 sigma apart) -------------- - 4: _SeedConfig(clusters=[ - _ClusterParams(cx=0.40, cy=0.45, ax=0.18, ay=0.15, theta=0.0, peak=110), - _ClusterParams(cx=0.58, cy=0.55, ax=0.16, ay=0.18, theta=0.5, peak=95), - ]), - - # --- 5: two clusters with ~90° different orientations (twinned) -------- - 5: _SeedConfig(clusters=[ - _ClusterParams(cx=0.30, cy=0.40, ax=0.25, ay=0.08, theta=0.0, peak=105), - _ClusterParams(cx=0.68, cy=0.62, ax=0.08, ay=0.25, theta=0.0, peak=100), - ]), - - # --- 6: pure background, no crystal – true negative -------------------- - 6: _SeedConfig(clusters=[]), -} - - -# --------------------------------------------------------------------------- -# Core generation -# --------------------------------------------------------------------------- - -def _gaussian_cluster( - ix: np.ndarray, - iy: np.ndarray, - p: _ClusterParams, - n_x: int, - n_y: int, -) -> np.ndarray: - """Return a 2-D Gaussian signal array for one cluster. - - ix, iy are integer coordinate grids with shape (n_x, n_y). - """ - cx = p.cx * (n_x - 1) - cy = p.cy * (n_y - 1) - - dx = (ix - cx) / max(p.ax * n_x, 0.5) - dy = (iy - cy) / max(p.ay * n_y, 0.5) - - cos_t = np.cos(p.theta) - sin_t = np.sin(p.theta) - dx_r = dx * cos_t + dy * sin_t - dy_r = -dx * sin_t + dy * cos_t - - return p.peak * np.exp(-2.0 * (dx_r ** 2 + dy_r ** 2)) - - -def generate_no_beam_scan_result( - request: RasterGridRequest, - seed: int | None = None, -) -> ScanResult: - """Build a synthetic ScanResult for no-beam / offline operation. - - Parameters - ---------- - request: - The grid request whose n_x, n_y, and file_prefix are used. - seed: - Integer 1-6 selects a pre-defined scenario from SEED_CATALOGUE. - Any other value (or None) draws cluster parameters randomly using - the seed as an RNG seed (None → fully random). - - Returns - ------- - ScanResult - Populated with one ScanResultImagesInner per grid cell, with - realistic background noise and optional crystal cluster(s). - """ - n_x = max(request.n_x, 1) - n_y = max(request.n_y, 1) - - rng = np.random.default_rng(seed) - - # Coordinate grids - ix, iy = np.meshgrid(np.arange(n_x), np.arange(n_y), indexing="ij") - - # Background: Poisson noise, mean ≈ 3 counts - bkg = rng.poisson(lam=3.0, size=(n_x, n_y)).astype(float) - - # Crystal signal layer (spots_low_res) - signal = np.zeros((n_x, n_y), dtype=float) - - if seed in SEED_CATALOGUE: - cfg = SEED_CATALOGUE[seed] - clusters = cfg.clusters - else: - # Random fallback: 1 or 2 clusters - n_clusters = rng.integers(1, 3) - clusters = [ - _ClusterParams( - cx=rng.uniform(0.15, 0.85), - cy=rng.uniform(0.15, 0.85), - ax=rng.uniform(0.12, 0.30), - ay=rng.uniform(0.12, 0.30), - theta=rng.uniform(0, np.pi), - peak=int(rng.integers(50, 150)), - ) - for _ in range(n_clusters) - ] - - for p in clusters: - signal += _gaussian_cluster(ix, iy, p, n_x, n_y) - - # Add Poisson noise on top of the crystal signal - spots_low_res = rng.poisson(lam=np.maximum(signal, 0)).astype(int) - - # Assemble image list – one entry per grid cell in serpentine order: - # row 0 left→right, row 1 right→left, row 2 left→right, … - images: list[ScanResultImagesInner] = [] - img_number = 0 - for xi in range(n_x): - yi_range = range(n_y) - for yi in yi_range: - images.append(ScanResultImagesInner( - number=img_number, - nx=xi, - ny=yi, - efficiency=1.0, - bkg=float(bkg[xi, yi]), - spots=int(spots_low_res[xi, yi]), - spots_low_res=int(spots_low_res[xi, yi]), - spots_indexed=0, - spots_ice=0, - index=0, - b=0.0, - res=None, - pixel_sum=None, - max=None, - sat=None, - err=None, - )) - img_number += 1 - - return ScanResult(file_prefix=request.file_prefix, images=images) diff --git a/src/aare/common/tell_models.py b/src/aare/common/tell_models.py deleted file mode 100644 index e52b98e8..00000000 --- a/src/aare/common/tell_models.py +++ /dev/null @@ -1,62 +0,0 @@ -from enum import Enum - -from pydantic import BaseModel - - -class TellActivityEnum(str, Enum): - IDLE = "idle" - MOUNTING = "mounting" - UNMOUNTING = "unmounting" - DRYING = "drying" - COOLING = "cooling" - ERROR = "error" - - def display_name(self) -> str: - return { - TellActivityEnum.IDLE: "Idle", - TellActivityEnum.MOUNTING: "Mounting", - TellActivityEnum.UNMOUNTING: "Unmounting", - TellActivityEnum.DRYING: "Drying", - TellActivityEnum.COOLING: "Cooling", - TellActivityEnum.ERROR: "Error", - }.get(self, str(self.value).capitalize()) - - -class TellPhaseEnum(str, Enum): - IDLE = "idle" - PREPARING = "preparing" - AUTO_UNMOUNT = "auto_unmount" - RETURNING_OLD_SAMPLE = "returning_old_sample" - OLD_SAMPLE_RETURNED = "old_sample_returned" - PICKING_NEW_SAMPLE = "picking_new_sample" - PLACING_NEW_SAMPLE = "placing_new_sample" - FINALIZING = "finalizing" - COMPLETE = "complete" - FAILED = "failed" - - def display_name(self) -> str: - return { - TellPhaseEnum.IDLE: "Idle", - TellPhaseEnum.PREPARING: "Preparing", - TellPhaseEnum.AUTO_UNMOUNT: "Auto-unmount", - TellPhaseEnum.RETURNING_OLD_SAMPLE: "Returning old sample", - TellPhaseEnum.OLD_SAMPLE_RETURNED: "Old sample returned", - TellPhaseEnum.PICKING_NEW_SAMPLE: "Picking new sample", - TellPhaseEnum.PLACING_NEW_SAMPLE: "Placing sample on gonio", - TellPhaseEnum.FINALIZING: "Finalizing", - TellPhaseEnum.COMPLETE: "Complete", - TellPhaseEnum.FAILED: "Failed", - }.get(self, str(self.value).replace("_", " ").capitalize()) - - -class TellStateModel(BaseModel): - activity: TellActivityEnum = TellActivityEnum.IDLE - message: str | None = None - last_event_class: str | None = None - last_event_value: str | None = None - last_update_ts: str | None = None - mount_success: bool | None = None - mount_error: str | None = None - sample_position: str | None = None - operation: str | None = None - phase: TellPhaseEnum | None = None \ No newline at end of file diff --git a/src/aare/daq/aaredb.py b/src/aare/daq/aaredb.py index 5cd69dc9..2bc745be 100644 --- a/src/aare/daq/aaredb.py +++ b/src/aare/daq/aaredb.py @@ -10,36 +10,39 @@ import aareDB import cv2 import numpy as np import requests +from aarecommon.config.logger import setup_logger +from aarecommon.config.logger_events import log_timing +from aarecommon.math.coordinate import Coordinate +from aarecommon.math.find_xtal import compute_crystal_score_array +from aarecommon.math.sample_geometry import SampleGeometryModel +from aarecommon.models.beamline import MXBeamline +from aarecommon.models.models import ( + DAQStatusModel, + PuckLoadedInfo, + SampleShortInfo, + ScanResultPayloadModel, +) +from aarecommon.models.raster_grid import ( + CenterOfMassModel, + RasterGridRequest, + RasterPayloadModel, +) +from aarecommon.models.rotation_scan import RotationScanRequest from aareDB import ( - SetTellPosition, + BeamlineParametersInput, + CharacterizationParameters, + Datasets, + ExperimentParametersCreate, + GridScanParameters, + RotationParameters, SampleEventCreate, SampleEventType, + SetTellPosition, SetTellPositionRequest, - CharacterizationParameters, - RotationParameters, - GridScanParameters, - Datasets, - Detector as DetectorParameters, - BeamlineParametersInput, - ExperimentParametersCreate) -from pydantic import StrictInt - -from aare.common.coordinate import Coordinate -from aare.common.find_xtal import compute_crystal_score_array -from aare.common.logger_config import setup_logger -from aare.common.logger_events import log_timing -from aare.common.models import ( - SampleShortInfo, - PuckLoadedInfo, - DAQStatusModel, ScanResultPayloadModel, ) - -from aare.common.beamline import MXBeamline -from aare.common.raster_grid import RasterGridRequest, RasterPayloadModel, CenterOfMassModel -from aare.common.rotation_scan import RotationScanRequest -from aare.common.sample_geometry import SampleGeometryModel - +from aareDB import Detector as DetectorParameters from jfjoch_client.models import ScanResult +from pydantic import StrictInt logger = setup_logger("aareDAQ") @@ -55,18 +58,26 @@ class AareWrapper: # --- 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.ssl_ca_cert = ( + "/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem" + ) # 2. Present Machine Identity (The certs that worked in curl) beamline_name = bl.value.lower() - configuration.cert_file = f"/etc/ssl/certs/secrets/mx-{beamline_name}-queue-01_from_dmz-01.crt" - configuration.key_file = f"/etc/ssl/certs/secrets/mx-{beamline_name}-queue-01_from_dmz-01.key" + configuration.cert_file = ( + f"/etc/ssl/certs/secrets/mx-{beamline_name}-queue-01_from_dmz-01.crt" + ) + configuration.key_file = ( + f"/etc/ssl/certs/secrets/mx-{beamline_name}-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.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) @@ -89,7 +100,7 @@ class AareWrapper: puck_in_segment=i.location.pos, ) o.append(t) - payload = SetTellPositionRequest(pucks = o, tell=self.__bl.value.upper()) + payload = SetTellPositionRequest(pucks=o, tell=self.__bl.value.upper()) ret = self.__tell_api.set_tell_positions( set_tell_position_request=payload, ) @@ -102,7 +113,7 @@ class AareWrapper: manual_sample = ManualSampleCreate( pgroup=s.user, sample_name=s.sample_name, - data_collection_parameters=s.aaredb_params + data_collection_parameters=s.aaredb_params, ) try: @@ -119,29 +130,42 @@ class AareWrapper: ) -> None: if sample_id is None or sample_id < 0: if sample_id is None: - logger.debug(f"Sample db_id is None, skipping sample event {event_type!s}") + logger.debug( + f"Sample db_id is None, skipping sample event {event_type!s}" + ) elif sample_id < 0: logger.debug( f"Sample db_id is invalid ({sample_id}), skipping sample event {event_type!s}" ) return try: - self.__sample_api.create_sample_event(sample_id=sample_id, sample_event_create=SampleEventCreate(event_type=event_type, comment=comment)) + self.__sample_api.create_sample_event( + sample_id=sample_id, + sample_event_create=SampleEventCreate( + event_type=event_type, comment=comment + ), + ) except Exception as e: logger.error(f"Error sending sample event {event_type!s} to db: {e}") @log_timing(logger, "AareDB call") - def upload_image(self, sample_id: int, filename: str, bgr_image: np.ndarray, message: Optional[str] = None): - _, buffer = cv2.imencode('.jpg', bgr_image) + def upload_image( + self, + sample_id: int, + filename: str, + bgr_image: np.ndarray, + message: Optional[str] = None, + ): + _, buffer = cv2.imencode(".jpg", bgr_image) jpeg_bytes = io.BytesIO(buffer) url = f"{self.__host}/protected_router/sample_runner/{sample_id}/upload-images" headers = { "accept": "application/json", - "X-Shared-Password": os.getenv("AAREDB_SHARED_PASSWORD") + "X-Shared-Password": os.getenv("AAREDB_SHARED_PASSWORD"), } request_kwargs = { - "files": {'uploaded_file': (filename + ".jpg", jpeg_bytes, "image/jpeg")}, + "files": {"uploaded_file": (filename + ".jpg", jpeg_bytes, "image/jpeg")}, "verify": self.__ssl_ca_cert, "cert": (self.__cert_file, self.__key_file), "headers": headers, @@ -152,15 +176,17 @@ class AareWrapper: logger.debug(f"Response status code: {response.status_code}") @log_timing(logger, "AareDB call") - def upload_jpg(self, sample_id: int, filename: str, jpg_image, message: Optional[str] = None): + def upload_jpg( + self, sample_id: int, filename: str, jpg_image, message: Optional[str] = None + ): logger.debug(f"jppg_image of type: {type(jpg_image)}") url = f"{self.__host}/protected_router/sample_runner/{sample_id}/upload-images" headers = { "accept": "application/json", - "X-Shared-Password": os.getenv("AAREDB_SHARED_PASSWORD") + "X-Shared-Password": os.getenv("AAREDB_SHARED_PASSWORD"), } request_kwargs = { - "files": {'uploaded_file': (filename + ".jpg", jpg_image, "image/jpeg")}, + "files": {"uploaded_file": (filename + ".jpg", jpg_image, "image/jpeg")}, "verify": self.__ssl_ca_cert, "cert": (self.__cert_file, self.__key_file), "headers": headers, @@ -171,28 +197,30 @@ class AareWrapper: logger.debug(f"Response status code: {response.status_code}") @log_timing(logger, "AareDB call") - def create_rotation_run(self, s: Optional[SampleShortInfo], r:RotationScanRequest, d:DAQStatusModel): + def create_rotation_run( + self, s: Optional[SampleShortInfo], r: RotationScanRequest, d: DAQStatusModel + ): if s is None: return try: if r.screening: characterization = CharacterizationParameters( - omegaStart_deg=round(r.start_omega_deg,3), + omegaStart_deg=round(r.start_omega_deg, 3), omegaStep=r.incr_omega_deg, - phi=round(d.geom.smargon.phi_deg,3), - chi=round(d.geom.smargon.chi_deg,3), + phi=round(d.geom.smargon.phi_deg, 3), + chi=round(d.geom.smargon.chi_deg, 3), numberOfImages=r.steps, exposureTime_s=r.exp_time_s, - oscillation_deg=r.wedge_omega_deg + oscillation_deg=r.wedge_omega_deg, ) rotation = None else: rotation = RotationParameters( - omegaStart_deg=round(r.start_omega_deg,3), + omegaStart_deg=round(r.start_omega_deg, 3), omegaStep=r.incr_omega_deg, - phi=round(d.geom.smargon.phi_deg,3), - chi=round(d.geom.smargon.chi_deg,3), + phi=round(d.geom.smargon.phi_deg, 3), + chi=round(d.geom.smargon.chi_deg, 3), numberOfImages=r.steps, exposureTime_s=r.exp_time_s, ) @@ -200,7 +228,7 @@ class AareWrapper: dataset = Datasets( filepath=r.file_prefix, status="written", - written_at=datetime.datetime.now() + written_at=datetime.datetime.now(), ) detector_data = DetectorParameters( manufacturer="DECTRIS", @@ -213,7 +241,7 @@ class AareWrapper: beamCenterY_px=d.diffraction.beam_center_pxl[1], pixelSizeX_um=d.diffraction.pixel_size_mm * 1000, pixelSizeY_um=d.diffraction.pixel_size_mm * 1000, - dataset=dataset + dataset=dataset, ) beamline_params = BeamlineParametersInput( synchrotron="Swiss Light Source", @@ -231,55 +259,55 @@ class AareWrapper: beamSizeHeight=d.geom.beam_size_mm.y * 1000, cryojetTemperature_K=d.bl.cryojet_K, rotation=rotation, - characterization=characterization + characterization=characterization, ) experiment_params_payload = ExperimentParametersCreate( - type="standard", - beamline_parameters=beamline_params, - sample_id=s.db_id + type="standard", beamline_parameters=beamline_params, sample_id=s.db_id ) response = self.__sample_api.create_experiment_parameters_for_sample( sample_id=s.db_id, - experiment_parameters_create=experiment_params_payload + experiment_parameters_create=experiment_params_payload, ) - #logger.debug("Experiment parameters created:", response) + # logger.debug("Experiment parameters created:", response) except Exception as e: logger.error(e) @log_timing(logger, "AareDB call") - def create_gridscan_run(self, s: Optional[SampleShortInfo], r:RasterGridRequest, d:DAQStatusModel): + def create_gridscan_run( + self, s: Optional[SampleShortInfo], r: RasterGridRequest, d: DAQStatusModel + ): if s is None: return try: gridscan = GridScanParameters( - #xStart=90.0, + # xStart=90.0, xStep=r.grid_size_mm.x, - #yStart=0.0, - yStep= r.grid_size_mm.y, + # yStart=0.0, + yStep=r.grid_size_mm.y, x_col=r.n_x, y_row=r.n_y, - omegaStart_deg=round(r.omega_deg,2), - numberOfImages=round(r.n_x * r.n_y,0), - exposureTime_s=round(r.exp_time_s,4), + omegaStart_deg=round(r.omega_deg, 2), + numberOfImages=round(r.n_x * r.n_y, 0), + exposureTime_s=round(r.exp_time_s, 4), ) dataset = Datasets( filepath=r.file_prefix, status="written", - written_at=datetime.datetime.now() + written_at=datetime.datetime.now(), ) detector_data = DetectorParameters( manufacturer="DECTRIS", model=d.diffraction.detector_description, type="photon-counting", serialNumber=d.diffraction.detector_serial_number, - detectorDistance_mm=round(r.dtz,3), - resolution_at_edge_Ang=round(d.diffraction.max_resolution_angstrom,3), - beamCenterX_px=round(d.diffraction.beam_center_pxl[0],2), - beamCenterY_px=round(d.diffraction.beam_center_pxl[1],2), - pixelSizeX_um=round(d.diffraction.pixel_size_mm * 1000,2), - pixelSizeY_um=round(d.diffraction.pixel_size_mm * 1000,2), - dataset=dataset + detectorDistance_mm=round(r.dtz, 3), + resolution_at_edge_Ang=round(d.diffraction.max_resolution_angstrom, 3), + beamCenterX_px=round(d.diffraction.beam_center_pxl[0], 2), + beamCenterY_px=round(d.diffraction.beam_center_pxl[1], 2), + pixelSizeX_um=round(d.diffraction.pixel_size_mm * 1000, 2), + pixelSizeY_um=round(d.diffraction.pixel_size_mm * 1000, 2), + dataset=dataset, ) beamline_params = BeamlineParametersInput( synchrotron="Swiss Light Source", @@ -287,7 +315,7 @@ class AareWrapper: detector=detector_data, energy_keV=d.diffraction.energy_keV, wavelength_Ang=d.diffraction.wavelength_angstrom, - ringCurrent_mA=round(d.bl.ring_current_mA,3), + ringCurrent_mA=round(d.bl.ring_current_mA, 3), ringMode="Beamline Development", monochromator="Si111", transmission=d.bl.transmission, @@ -296,30 +324,36 @@ class AareWrapper: beamSizeWidth=d.geom.beam_size_mm.x * 1000, beamSizeHeight=d.geom.beam_size_mm.y * 1000, cryojetTemperature_K=d.bl.cryojet_K, - gridScan=gridscan + gridScan=gridscan, ) experiment_params_payload = ExperimentParametersCreate( - type="standard", - beamline_parameters=beamline_params, - sample_id=s.db_id + type="standard", beamline_parameters=beamline_params, sample_id=s.db_id ) response = self.__sample_api.create_experiment_parameters_for_sample( sample_id=s.db_id, - experiment_parameters_create=experiment_params_payload + experiment_parameters_create=experiment_params_payload, ) - #logger.info("Experiment parameters created:", response) + # logger.info("Experiment parameters created:", response) except Exception as e: logger.debug(e) @log_timing(logger, "AareDB call") - def ingest_gridscan(self, sample: Optional[SampleShortInfo], raster_result: ScanResult, - raster_request: RasterGridRequest, geom: SampleGeometryModel, - com: Optional[CenterOfMassModel], beam_mark_pxl:tuple[float,float]): + def ingest_gridscan( + self, + sample: Optional[SampleShortInfo], + raster_result: ScanResult, + raster_request: RasterGridRequest, + geom: SampleGeometryModel, + com: Optional[CenterOfMassModel], + beam_mark_pxl: tuple[float, float], + ): if sample is None: return - payload_model = self.format_gridscan_payload(sample, raster_result, raster_request, geom, com, beam_mark_pxl) + payload_model = self.format_gridscan_payload( + sample, raster_result, raster_request, geom, com, beam_mark_pxl + ) if payload_model is None: return payload = payload_model.model_dump() @@ -327,26 +361,36 @@ class AareWrapper: url = f"{self.__host}/protected_router/gridscan_runner/ingest" headers = { "accept": "application/json", - "X-Shared-Password": os.getenv("AAREDB_SHARED_PASSWORD") + "X-Shared-Password": os.getenv("AAREDB_SHARED_PASSWORD"), } - response = requests.post(url, + response = requests.post( + url, auth=(os.getenv("AAREDB_USERNAME"), os.getenv("AAREDB_PASSWORD")), - headers=headers, data=json.dumps(payload), timeout=30, + headers=headers, + data=json.dumps(payload), + timeout=30, verify=self.__ssl_ca_cert, - cert=(self.__cert_file, self.__key_file)) + cert=(self.__cert_file, self.__key_file), + ) response.raise_for_status() logger.info(f"Response status code: {response.status_code}") - - def format_gridscan_payload(self, sample: Optional[SampleShortInfo], raster_result:ScanResult, - raster_request:RasterGridRequest, geom:SampleGeometryModel, - com: Optional[CenterOfMassModel], beam_mark_pxl:tuple[float,float]) -> RasterPayloadModel|None: + def format_gridscan_payload( + self, + sample: Optional[SampleShortInfo], + raster_result: ScanResult, + raster_request: RasterGridRequest, + geom: SampleGeometryModel, + com: Optional[CenterOfMassModel], + beam_mark_pxl: tuple[float, float], + ) -> RasterPayloadModel | None: try: - - cell_size_pxl = Coordinate(x=raster_request.grid_size_mm.x/geom.pixel_in_mm, - y=raster_request.grid_size_mm.y/geom.pixel_in_mm) + cell_size_pxl = Coordinate( + x=raster_request.grid_size_mm.x / geom.pixel_in_mm, + y=raster_request.grid_size_mm.y / geom.pixel_in_mm, + ) if raster_request.smargon_top_left is None: beam_loc = geom.smargon.sh_mm @@ -359,30 +403,36 @@ class AareWrapper: com_grid_pxl = com.get_com_pxl(raster_request, geom) x = start_pxl.x + com_grid_pxl.x y = start_pxl.y + com_grid_pxl.y - center_pxl = Coordinate(x=x,y=y) + center_pxl = Coordinate(x=x, y=y) else: center_pxl = None try: score_arr = compute_crystal_score_array(raster_result.images) - score = [float(score_arr[img.nx, img.ny]) if img.nx is not None and img.ny is not None else None - for img in raster_result.images] + score = [ + float(score_arr[img.nx, img.ny]) + if img.nx is not None and img.ny is not None + else None + for img in raster_result.images + ] except Exception as e: - logger.warning(f"raster score computation failed, sending null score: {e}") + logger.warning( + f"raster score computation failed, sending null score: {e}" + ) score = None payload = RasterPayloadModel( - request = raster_request, - result = raster_result, - sample_id = sample.db_id, - attach_image = True, - centre_of_mass = com, - raster_score = score, # TODO: Check whether AareDB needs to change for taking this input - start_pxl = start_pxl, - center_pxl = center_pxl, - cell_size_pxl = cell_size_pxl, - beam_mark_pxl = beam_mark_pxl, - beam_size_mm = geom.beam_size_mm + request=raster_request, + result=raster_result, + sample_id=sample.db_id, + attach_image=True, + centre_of_mass=com, + raster_score=score, # TODO: Check whether AareDB needs to change for taking this input + start_pxl=start_pxl, + center_pxl=center_pxl, + cell_size_pxl=cell_size_pxl, + beam_mark_pxl=beam_mark_pxl, + beam_size_mm=geom.beam_size_mm, ) return payload @@ -392,8 +442,13 @@ class AareWrapper: raise e @log_timing(logger, "AareDB call") - def ingest_scan(self, sample: Optional[SampleShortInfo], result: ScanResult, - geom: SampleGeometryModel, beam_mark_pxl:tuple[float,float]): + def ingest_scan( + self, + sample: Optional[SampleShortInfo], + result: ScanResult, + geom: SampleGeometryModel, + beam_mark_pxl: tuple[float, float], + ): if sample is None: return @@ -406,28 +461,35 @@ class AareWrapper: url = f"{self.__host}/protected_router/scan_runner/ingest" headers = { "accept": "application/json", - "X-Shared-Password": os.getenv("AAREDB_SHARED_PASSWORD") + "X-Shared-Password": os.getenv("AAREDB_SHARED_PASSWORD"), } - response = requests.post(url, + response = requests.post( + url, auth=(os.getenv("AAREDB_USERNAME"), os.getenv("AAREDB_PASSWORD")), - headers=headers, data=json.dumps(payload), timeout=30, + headers=headers, + data=json.dumps(payload), + timeout=30, verify=self.__ssl_ca_cert, - cert=(self.__cert_file, self.__key_file)) + cert=(self.__cert_file, self.__key_file), + ) response.raise_for_status() logger.info(f"Response status code: {response.status_code}") - def format_scan_payload(self, sample: Optional[SampleShortInfo], result:ScanResult, - geom:SampleGeometryModel, - beam_mark_pxl:tuple[float,float]) -> ScanResultPayloadModel|None: + def format_scan_payload( + self, + sample: Optional[SampleShortInfo], + result: ScanResult, + geom: SampleGeometryModel, + beam_mark_pxl: tuple[float, float], + ) -> ScanResultPayloadModel | None: try: - payload = ScanResultPayloadModel( - result = result, - sample_id = sample.db_id, - attach_image = True, - beam_mark_pxl = beam_mark_pxl, - beam_size_mm = geom.beam_size_mm + result=result, + sample_id=sample.db_id, + attach_image=True, + beam_mark_pxl=beam_mark_pxl, + beam_size_mm=geom.beam_size_mm, ) return payload diff --git a/src/aare/daq/auth.py b/src/aare/daq/auth.py index 118d91b5..908ff930 100644 --- a/src/aare/daq/auth.py +++ b/src/aare/daq/auth.py @@ -3,31 +3,40 @@ import ipaddress import logging import os import pwd -import uuid -from datetime import datetime, timedelta, UTC -from typing import List import time - -logger = logging.getLogger("aareDAQ") +import uuid +from datetime import UTC, datetime, timedelta +from typing import List import jwt +from aarecommon.errors.exception_handler import ( + AuthenticationException, + AuthErrorCode, + UserRightsException, +) +from aarecommon.models.auth import ( + BatonRequest, + BatonRequestStatus, + BatonStatus, + BatonTransferQueue, +) +from aarecommon.models.models import SessionsStateEnum from fastapi import Depends, HTTPException, Request, status -from fastapi.security import OAuth2PasswordRequestForm, OAuth2PasswordBearer +from fastapi.security import OAuth2PasswordBearer, 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 - +logger = logging.getLogger("aareDAQ") if os.environ.get("JWT_AAREDAQ_KEY") is None: - raise Exception("JWT_AAREDAQ_KEY environment variable not set, cannot guarantee safe authentication.") + raise Exception( + "JWT_AAREDAQ_KEY environment variable not set, cannot guarantee safe authentication." + ) SECRET_KEY = os.environ.get("JWT_AAREDAQ_KEY") ALGORITHM = "HS256" -ACCESS_TOKEN_EXPIRE_MINUTES = 24 * 60 * 7 # 1 week +ACCESS_TOKEN_EXPIRE_MINUTES = 24 * 60 * 7 # 1 week SESSION_EXPIRE_SECONDS = 60 * 10 BATON_REQUEST_TIMEOUT_SECONDS = 30 @@ -36,8 +45,9 @@ SUPER_USERS = ["e10019", "e11206", "e18747"] oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token") + class TokenData(BaseModel): - sub: str # Username + sub: str # Username pgroups: List[str] session: int staff: bool = False @@ -45,7 +55,9 @@ class TokenData(BaseModel): def create_access_token(token: TokenData): to_encode = token.model_dump() - to_encode.update({"exp": datetime.now(UTC) + timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES)}) + to_encode.update( + {"exp": datetime.now(UTC) + timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES)} + ) encoded_jwt = jwt.encode(to_encode, SECRET_KEY, algorithm=ALGORITHM) return encoded_jwt @@ -70,9 +82,13 @@ def authenticate_from_proxy_header(request: Request) -> str: only the Apache proxy running on the same host can supply it. """ client_host = request.client.host if request.client else None - logger.debug(f"[auth] /token client_host={client_host!r} request.client={request.client!r}") + logger.debug( + f"[auth] /token client_host={client_host!r} request.client={request.client!r}" + ) if not _is_loopback(client_host): - logger.warning(f"[auth] Rejecting X-Remote-User: client_host {client_host!r} is not loopback") + logger.warning( + f"[auth] Rejecting X-Remote-User: client_host {client_host!r} is not loopback" + ) raise AuthenticationException( message="X-Remote-User header is only trusted from the localhost proxy", status_code=401, @@ -96,12 +112,15 @@ def authenticate_user(cfg: BeamlineConfig, username: str) -> str: groups = os.getgrouplist(user_info.pw_name, user_info.pw_gid) supplementary_groups = [grp.getgrgid(g).gr_name.lower() for g in groups] super_user = username in SUPER_USERS - pgroups = [group for group in supplementary_groups if group.startswith('p')] - staff = "unx-mxgroup" in supplementary_groups or "unx-sls_mx" in supplementary_groups or super_user - token = TokenData(sub=username, - pgroups=pgroups, - staff=staff, - session=cfg.generate_session()) + pgroups = [group for group in supplementary_groups if group.startswith("p")] + staff = ( + "unx-mxgroup" in supplementary_groups + or "unx-sls_mx" in supplementary_groups + or super_user + ) + token = TokenData( + sub=username, pgroups=pgroups, staff=staff, session=cfg.generate_session() + ) return create_access_token(token) @@ -146,6 +165,7 @@ def check_jwt_rw(cfg: BeamlineConfig, data: TokenData) -> None: # In case something is wrong but you are holder (maybe redis expiry?) cfg.try_set_active_session(data.session, SESSION_EXPIRE_SECONDS) + def check_jwt_staff_only(data: TokenData) -> None: if not data.staff: raise UserRightsException( @@ -154,31 +174,36 @@ def check_jwt_staff_only(data: TokenData) -> None: code=AuthErrorCode.NOT_STAFF, ) + def check_jwt_staff(cfg: BeamlineConfig, data: TokenData) -> None: check_jwt_staff_only(data) holder = cfg.baton_holder if holder and holder.session != data.session: - # Staff can take over if they don't have it, but they need to use force_current_session - # or request_baton (which does staff override). - # If they are calling a RW endpoint, they SHOULD already be the holder. - pass + # Staff can take over if they don't have it, but they need to use force_current_session + # or request_baton (which does staff override). + # If they are calling a RW endpoint, they SHOULD already be the holder. + pass try: cfg.try_extend_active_session(data.session, SESSION_EXPIRE_SECONDS) except Exception: cfg.try_set_active_session(data.session, SESSION_EXPIRE_SECONDS) + def force_current_sesion(cfg: BeamlineConfig, data: TokenData) -> None: 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 + expiry_sec=SESSION_EXPIRE_SECONDS, ) -def _finalize_expired_baton_request(cfg: BeamlineConfig, pending: BatonRequest) -> BatonStatus: + +def _finalize_expired_baton_request( + cfg: BeamlineConfig, pending: BatonRequest +) -> BatonStatus: """ Resolve an expired baton request in one place. @@ -195,7 +220,7 @@ def _finalize_expired_baton_request(cfg: BeamlineConfig, pending: BatonRequest) to_username=pending.requester_username, to_is_staff=requester_is_staff, to_pgroup=cfg.pgroup, - expiry_sec=SESSION_EXPIRE_SECONDS + expiry_sec=SESSION_EXPIRE_SECONDS, ) else: cfg.queued_baton_transfer = BatonTransferQueue( @@ -204,16 +229,20 @@ def _finalize_expired_baton_request(cfg: BeamlineConfig, pending: BatonRequest) target_is_staff=requester_is_staff, target_pgroup=cfg.pgroup, queued_at=time.time(), - reason="timeout_beamline_busy" + 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 - )) + 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: """ @@ -230,6 +259,7 @@ 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: """ Build baton status scoped to the requesting session. @@ -267,6 +297,7 @@ def get_baton_status(cfg: BeamlineConfig, data: TokenData) -> BatonStatus: allow_non_staff_request=cfg.allow_non_staff_request_from_staff, ) + def request_baton(cfg: BeamlineConfig, data: TokenData) -> dict: resolve_baton_timeout_if_needed(cfg) @@ -326,10 +357,19 @@ def request_baton(cfg: BeamlineConfig, data: TokenData) -> dict: to_pgroup=cfg.pgroup, expiry_sec=SESSION_EXPIRE_SECONDS, ) - return {"granted": True, "override": True, "message": "Staff override - baton acquired"} + return { + "granted": True, + "override": True, + "message": "Staff override - baton acquired", + } # Policy check: non-staff requesting from staff - if holder and holder.is_staff and not data.staff and not cfg.allow_non_staff_request_from_staff: + if ( + holder + and holder.is_staff + and not data.staff + and not cfg.allow_non_staff_request_from_staff + ): return { "error": True, "message": "Requesting baton from staff is disabled by backend policy.", @@ -362,7 +402,9 @@ def request_baton(cfg: BeamlineConfig, data: TokenData) -> dict: msg = f"Request sent to {holder.username if holder else 'current holder'}" if is_busy: - msg += " (Note: beamline is currently busy, transfer will be queued if accepted)" + msg += ( + " (Note: beamline is currently busy, transfer will be queued if accepted)" + ) return { "pending": True, @@ -372,7 +414,10 @@ def request_baton(cfg: BeamlineConfig, data: TokenData) -> dict: "beamline_busy": is_busy, } -def respond_to_baton_request(cfg: BeamlineConfig, data: TokenData, accept: bool) -> dict: + +def respond_to_baton_request( + cfg: BeamlineConfig, data: TokenData, accept: bool +) -> dict: """ Current baton holder responds to a pending request. """ @@ -394,9 +439,13 @@ def respond_to_baton_request(cfg: BeamlineConfig, data: TokenData, accept: bool) to_username=pending.requester_username, to_is_staff=pending.requester_is_staff, to_pgroup=cfg.pgroup, - expiry_sec=SESSION_EXPIRE_SECONDS + expiry_sec=SESSION_EXPIRE_SECONDS, ) - return {"accepted": True, "transferred": True, "message": "Baton transferred"} + return { + "accepted": True, + "transferred": True, + "message": "Baton transferred", + } else: # Beamline busy, queue the transfer cfg.queued_baton_transfer = BatonTransferQueue( @@ -405,7 +454,7 @@ def respond_to_baton_request(cfg: BeamlineConfig, data: TokenData, accept: bool) target_is_staff=pending.requester_is_staff, target_pgroup=cfg.pgroup, queued_at=time.time(), - reason="accepted_beamline_busy" + reason="accepted_beamline_busy", ) cfg.clear_pending_baton_request() return { @@ -419,6 +468,7 @@ def respond_to_baton_request(cfg: BeamlineConfig, data: TokenData, accept: bool) cfg.set_pending_baton_request(pending, timeout_sec=10) return {"refused": True, "message": "Baton request refused"} + def release_baton(cfg: BeamlineConfig, data: TokenData) -> dict: """ Voluntarily release the baton (set session to free). @@ -438,6 +488,7 @@ def release_baton(cfg: BeamlineConfig, data: TokenData) -> dict: 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. @@ -452,4 +503,4 @@ def cancel_baton_request(cfg: BeamlineConfig, data: TokenData) -> dict: 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 + return {"cancelled": True, "message": "Request cancelled"} diff --git a/src/aare/daq/config.py b/src/aare/daq/config.py index ee0d038e..36ee8161 100644 --- a/src/aare/daq/config.py +++ b/src/aare/daq/config.py @@ -1,58 +1,56 @@ import base64 import io import json -from datetime import datetime -from typing import Tuple, List -from dataclasses import asdict, is_dataclass import time +from dataclasses import asdict, is_dataclass +from datetime import datetime +from typing import List, Tuple import numpy as np import redis import redis_lock -from aare.common.coordinate import Coordinate, AerotechCoordinate -from aare.common.automation_models import AutomationProgress -from aare.common.models import ( - BeamlineSettingsModel, - BeamMarkCoeffModel, - ZoomModeEnum, - zoom_manager, - SampleShortInfo, - SessionStatus, - BeamlineStateEnum, - SessionsStateEnum, - SampleShortInfoList, - CryojetSettingsModel, - ZoomModel, - SampleCameraSettings, - FluorescenceSpectrumOutputModel, - CrystalSize, - SimpleStrategyInputModel, - SimpleScanParameters, - OpenGuiSessionInfo -) - -from aare.common.auth_models import ( - BatonRequest, +from aarecommon.config.beamline import cfg_get +from aarecommon.config.logger import setup_logger +from aarecommon.errors.exception_handler import BeamlineBusyException +from aarecommon.math.coordinate import AerotechCoordinate, Coordinate +from aarecommon.models.auth import ( BatonHolderInfo, + BatonRequest, BatonRequestStatus, BatonTransferQueue, ) -from aare.common.beamline import MXBeamline, cfg_get -from aare.common.logger_config import setup_logger +from aarecommon.models.automation import AutomationProgress +from aarecommon.models.beamline import MXBeamline +from aarecommon.models.models import ( + BeamlineSettingsModel, + BeamlineStateEnum, + BeamMarkCoeffModel, + CryojetSettingsModel, + CrystalSize, + FluorescenceSpectrumOutputModel, + OpenGuiSessionInfo, + SampleCameraSettings, + SampleShortInfo, + SampleShortInfoList, + SessionsStateEnum, + SessionStatus, + SimpleScanParameters, + SimpleStrategyInputModel, + ZoomModeEnum, + ZoomModel, + zoom_manager, +) -from aare.common.exception_handler import BeamlineBusyException from aare.daq.config_model import LocalContactConfigModel -#TODO WHAT SHOULD THIS BE? This should be in the YAMl file it is beamline specific -ABR_POS_MOUNT = AerotechCoordinate( - at_mm=Coordinate(x=0, y=0, z=0), - omega_deg=0 -) +# TODO WHAT SHOULD THIS BE? This should be in the YAMl file it is beamline specific +ABR_POS_MOUNT = AerotechCoordinate(at_mm=Coordinate(x=0, y=0, z=0), omega_deg=0) ABR_OMEGA_MOUNT = 0.0 DEFAULT_LENS_MAGNIFICATION = 10.0 logger = setup_logger("aareDAQ") + # Serialize the NumPy array to a Base64 string def numpy_to_base64(array: np.ndarray) -> str: buffer = io.BytesIO() # Create an in-memory buffer @@ -123,7 +121,7 @@ class BeamlineConfig: f"{raw_dtz_safe_position!r} ({e})" ) - #GUI session management + # GUI session management def _gui_sessions_index_key(self) -> str: return f"{self.__bl}:gui_sessions" @@ -132,9 +130,9 @@ class BeamlineConfig: return f"{self.__bl}:gui_session:{session}" def _write_gui_session( - self, - payload: OpenGuiSessionInfo, - expiry_sec: int | None = None, + self, + payload: OpenGuiSessionInfo, + expiry_sec: int | None = None, ) -> None: expiry = int(expiry_sec or self.GUI_SESSION_EXPIRE_SECONDS) @@ -165,12 +163,12 @@ class BeamlineConfig: return None def touch_gui_session( - self, - *, - session: int, - username: str, - staff: bool = False, - expiry_sec: int, + self, + *, + session: int, + username: str, + staff: bool = False, + expiry_sec: int, ) -> OpenGuiSessionInfo: existing = self._read_gui_session(session) holder = self.baton_holder @@ -181,11 +179,19 @@ class BeamlineConfig: username=username, staff=staff, last_seen_ts=time.time(), - last_interaction_ts=existing.last_interaction_ts if existing is not None else None, + last_interaction_ts=existing.last_interaction_ts + if existing is not None + else None, close_requested=existing.close_requested if existing is not None else False, - close_requested_by=existing.close_requested_by if existing is not None else None, - close_requested_at=existing.close_requested_at if existing is not None else None, - close_grace_seconds=existing.close_grace_seconds if existing is not None else None, + close_requested_by=existing.close_requested_by + if existing is not None + else None, + close_requested_at=existing.close_requested_at + if existing is not None + else None, + close_grace_seconds=existing.close_grace_seconds + if existing is not None + else None, holds_baton=holds_baton, ) @@ -193,7 +199,9 @@ class BeamlineConfig: self.purge_expired_gui_sessions() return payload - def update_gui_interaction(self, session: int, last_interaction_ts: float) -> OpenGuiSessionInfo | None: + def update_gui_interaction( + self, session: int, last_interaction_ts: float + ) -> OpenGuiSessionInfo | None: payload = self._read_gui_session(session) if payload is None: self.__client.srem(self._gui_sessions_index_key(), session) @@ -205,11 +213,11 @@ class BeamlineConfig: return payload def request_gui_close( - self, - *, - session: int, - requested_by: str, - grace_seconds: int = 60, + self, + *, + session: int, + requested_by: str, + grace_seconds: int = 60, ) -> OpenGuiSessionInfo | None: payload = self._read_gui_session(session) if payload is None: @@ -273,7 +281,7 @@ class BeamlineConfig: for session_id in sorted((int(s) for s in session_ids)): payload = self._read_gui_session(session_id) if payload is not None: - payload.holds_baton = (payload.session == holder_session) + payload.holds_baton = payload.session == holder_session sessions.append(payload) return sessions @@ -285,7 +293,9 @@ class BeamlineConfig: return None holder = self.baton_holder - payload.holds_baton = bool(holder is not None and holder.session == payload.session) + payload.holds_baton = bool( + holder is not None and holder.session == payload.session + ) return payload # Session and authentication management @@ -321,8 +331,9 @@ class BeamlineConfig: return baton.session def session_status(self, session: int) -> SessionStatus: - return SessionStatus(session=self.session_state(session), - current_pgroup=self.pgroup) + return SessionStatus( + session=self.session_state(session), current_pgroup=self.pgroup + ) def session_state(self, session: int) -> SessionsStateEnum: holder = self.baton_holder @@ -333,18 +344,26 @@ class BeamlineConfig: if holder.session == session: # You are the holder. Check if someone else requested from you. - if pending and pending.status == BatonRequestStatus.PENDING and pending.holder_session == session: + if ( + pending + and pending.status == BatonRequestStatus.PENDING + and pending.holder_session == session + ): return SessionsStateEnum.PendingElseToYou return SessionsStateEnum.OwnedByYou else: # Someone else is the holder. Check if you requested from them. - if pending and pending.status == BatonRequestStatus.PENDING and pending.requester_session == session: + if ( + pending + and pending.status == BatonRequestStatus.PENDING + and pending.requester_session == session + ): return SessionsStateEnum.PendingYouToElse return SessionsStateEnum.OwnedByElse def try_set_active_session(self, session: int, expiry_sec: int) -> None: with redis_lock.Lock( - self.__client, f"{self.__bl}:active_session_lock", expire=10 + self.__client, f"{self.__bl}:active_session_lock", expire=10 ): active = self.active_session if active is None: @@ -355,10 +374,10 @@ class BeamlineConfig: ) self.__client.expire(f"{self.__bl}:active_session", expiry_sec) - #TODO finish setting this up! + # TODO finish setting this up! def try_extend_active_session(self, session: int, expiry_sec: int) -> None: with redis_lock.Lock( - self.__client, f"{self.__bl}:active_session_lock", expire=10 + self.__client, f"{self.__bl}:active_session_lock", expire=10 ): active = self.active_session if active is None: @@ -432,7 +451,9 @@ class BeamlineConfig: except Exception: return None - def set_pending_baton_request(self, request: BatonRequest | None, timeout_sec: int = 30) -> 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") @@ -460,7 +481,9 @@ class BeamlineConfig: 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()) + 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).""" @@ -473,19 +496,19 @@ class BeamlineConfig: return True def execute_baton_transfer( - self, - to_session: int, - to_username: str, - to_is_staff: bool, - to_pgroup: str | None, - expiry_sec: int + 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, 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) @@ -493,7 +516,7 @@ class BeamlineConfig: username=to_username, session=to_session, is_staff=to_is_staff, - pgroup=to_pgroup + pgroup=to_pgroup, ) # Clear any pending request or queued transfer self.clear_pending_baton_request() @@ -516,7 +539,7 @@ class BeamlineConfig: to_username=queued.target_username, to_is_staff=queued.target_is_staff, to_pgroup=queued.target_pgroup, - expiry_sec=expiry_sec + expiry_sec=expiry_sec, ) return True @@ -564,7 +587,7 @@ class BeamlineConfig: raise Exception("Beamline is not in a proper state") def start_moving( - self, target: BeamlineStateEnum, timeout: int | None = None + self, target: BeamlineStateEnum, timeout: int | None = None ) -> BeamlineStateEnum: self.try_set_busy(timeout=timeout) curr_state = self.state @@ -617,7 +640,9 @@ class BeamlineConfig: ) # Apply lens magnification correction relative to the default 10x lens. # A lower magnification lens (e.g. 5x) makes each pixel cover more physical space. - lens_magnification = cfg_get("daq.hardware.lens_magnification", DEFAULT_LENS_MAGNIFICATION) + lens_magnification = cfg_get( + "daq.hardware.lens_magnification", DEFAULT_LENS_MAGNIFICATION + ) try: lens_magnification = float(lens_magnification) except (TypeError, ValueError): @@ -634,11 +659,15 @@ class BeamlineConfig: z = ln(lens_factor / (b * target)) / a """ if target_pixel_in_mm <= 0: - raise ValueError(f"target_pixel_in_mm must be > 0, got {target_pixel_in_mm}") + raise ValueError( + f"target_pixel_in_mm must be > 0, got {target_pixel_in_mm}" + ) cfg = self.settings a = cfg.camera_translation_factor_a b = cfg.camera_translation_factor_b - lens_magnification = cfg_get("daq.hardware.lens_magnification", DEFAULT_LENS_MAGNIFICATION) + lens_magnification = cfg_get( + "daq.hardware.lens_magnification", DEFAULT_LENS_MAGNIFICATION + ) try: lens_magnification = float(lens_magnification) except (TypeError, ValueError): @@ -704,7 +733,9 @@ class BeamlineConfig: def settings(self, data: BeamlineSettingsModel): with redis_lock.Lock(self.__client, f"{self.__bl}:settings_lock", expire=10): current = self.__get_settings() - updated_data = current.model_copy(update=data.model_dump(exclude_unset=True)) + updated_data = current.model_copy( + update=data.model_dump(exclude_unset=True) + ) self.__client.set(f"{self.__bl}:settings", updated_data.model_dump_json()) @property @@ -721,10 +752,14 @@ class BeamlineConfig: self.__client.set(f"{self.__bl}:cryojet_settings", data.model_dump_json()) def get_alc_bkg(self, zoom: float, exp: float, gain: float) -> np.ndarray | None: - return base64_to_numpy(self.__client.get(f"{self.__bl}:bkg{zoom:.1f}_{exp:.2f}_{gain:.1f}")) + return base64_to_numpy( + self.__client.get(f"{self.__bl}:bkg{zoom:.1f}_{exp:.2f}_{gain:.1f}") + ) def put_alc_bkg(self, zoom: float, exp: float, gain: float, data: np.ndarray): - self.__client.set(f"{self.__bl}:bkg{zoom:.1f}_{exp:.2f}_{gain:.1f}", numpy_to_base64(data)) + self.__client.set( + f"{self.__bl}:bkg{zoom:.1f}_{exp:.2f}_{gain:.1f}", numpy_to_base64(data) + ) @property def spreadsheet(self) -> SampleShortInfoList: @@ -798,12 +833,12 @@ class BeamlineConfig: def beam_mark_coeff(self, data: BeamMarkCoeffModel): self.__client.set(f"{self.__bl}:beam_center_camera", data.model_dump_json()) -#TODO tidy up zoom functions + # TODO tidy up zoom functions @property def zoom_mode(self) -> ZoomModeEnum: raw_value = self.__client.get(f"{self.__bl}:zoom_mode") if raw_value is None: - print('no zoom mode given, defaulting to user mode') + print("no zoom mode given, defaulting to user mode") return ZoomModeEnum.User try: int_value = int(raw_value) # Ensure it's an integer @@ -814,7 +849,7 @@ class BeamlineConfig: ) @zoom_mode.setter - def zoom_mode(self, mode: ZoomModeEnum): + def zoom_mode(self, mode: ZoomModeEnum): self.__client.set(f"{self.__bl}:zoom_mode", mode.value) @staticmethod @@ -842,11 +877,13 @@ class BeamlineConfig: return ZoomModel(**data_dict) @zoom_settings.setter - def zoom_settings(self,data: ZoomModel): + def zoom_settings(self, data: ZoomModel): mode = self.zoom_mode if not mode or not isinstance(mode, ZoomModeEnum): raise Exception("incorrect zoom settings mode used") - self.__client.set(f"{self.__bl}:{self.zoom_setting_string(mode)}", data.model_dump_json()) + self.__client.set( + f"{self.__bl}:{self.zoom_setting_string(mode)}", data.model_dump_json() + ) def save_zoom_camera_setting( self, @@ -860,7 +897,11 @@ class BeamlineConfig: mode = mode or self.zoom_mode key = f"{self.__bl}:{self.zoom_setting_string(mode)}" tmp = self.__client.get(key) - model = ZoomModel(**json.loads(tmp)) if tmp is not None else zoom_manager(mode, self.__mxb) + model = ( + ZoomModel(**json.loads(tmp)) + if tmp is not None + else zoom_manager(mode, self.__mxb) + ) model.z[zoom_value] = settings self.__client.set(key, model.model_dump_json()) @@ -949,7 +990,7 @@ class BeamlineConfig: def crystal_size(self) -> CrystalSize: tmp = self.__client.get(f"{self.__bl}:crystal_size") if tmp is None: - return CrystalSize(x=0,y=0,z=0) + return CrystalSize(x=0, y=0, z=0) data_dict = json.loads(tmp) return CrystalSize(**data_dict) @@ -1006,6 +1047,7 @@ class BeamlineConfig: def record_mount_success(self): self.reset_mount_failure_streak() + @property def simple_input_parameters(self) -> SimpleStrategyInputModel | None: tmp = self.__client.get(f"{self.__bl}:simple_input_params") @@ -1015,11 +1057,13 @@ class BeamlineConfig: return SimpleStrategyInputModel(**data_dict) @simple_input_parameters.setter - def simple_input_parameters(self, input_params:SimpleStrategyInputModel | None): + def simple_input_parameters(self, input_params: SimpleStrategyInputModel | None): if input_params is None: self.__client.delete(f"{self.__bl}:simple_input_params") else: - self.__client.set(f"{self.__bl}:simple_input_params", input_params.model_dump_json()) + self.__client.set( + f"{self.__bl}:simple_input_params", input_params.model_dump_json() + ) @property def auto_params(self) -> SimpleScanParameters | None: @@ -1035,7 +1079,7 @@ class BeamlineConfig: return None @auto_params.setter - def auto_params(self, params:SimpleScanParameters | None): + def auto_params(self, params: SimpleScanParameters | None): if params is None: self.__client.delete(f"{self.__bl}:auto_params") else: @@ -1063,11 +1107,15 @@ class BeamlineConfig: "progress": progress, } - def set_automation_progress_state(self, progress: AutomationProgress | dict) -> dict: + def set_automation_progress_state( + self, progress: AutomationProgress | dict + ) -> dict: def _json_default(value): if isinstance(value, datetime): return value.isoformat() - raise TypeError(f"Object of type {type(value).__name__} is not JSON serializable") + raise TypeError( + f"Object of type {type(value).__name__} is not JSON serializable" + ) if isinstance(progress, dict): payload = dict(progress) @@ -1106,7 +1154,7 @@ class BeamlineConfig: return 0 @failed_mount_count.setter - def failed_mount_count(self, count:int): + def failed_mount_count(self, count: int): if count == 0: self.__client.delete(f"{self.__bl}:failed_mount_count") else: @@ -1214,9 +1262,15 @@ class BeamlineConfig: def set_detector_metadata(self, payload: dict) -> dict: safe_payload = dict(payload or {}) - safe_payload["dtz_low"] = self._coerce_optional_float(safe_payload.get("dtz_low")) - safe_payload["dtz_high"] = self._coerce_optional_float(safe_payload.get("dtz_high")) - safe_payload["pixel_size_mm"] = self._coerce_optional_float(safe_payload.get("pixel_size_mm")) + safe_payload["dtz_low"] = self._coerce_optional_float( + safe_payload.get("dtz_low") + ) + safe_payload["dtz_high"] = self._coerce_optional_float( + safe_payload.get("dtz_high") + ) + safe_payload["pixel_size_mm"] = self._coerce_optional_float( + safe_payload.get("pixel_size_mm") + ) safe_payload["updated_at"] = datetime.now().isoformat(timespec="seconds") self.__client.set(self._detector_metadata_key(), json.dumps(safe_payload)) return safe_payload @@ -1282,7 +1336,9 @@ class BeamlineConfig: logger.warning(f"Failed to read Local Contact config from Redis: {e}") return default - def set_local_contact_config(self, config: LocalContactConfigModel | dict) -> LocalContactConfigModel: + def set_local_contact_config( + self, config: LocalContactConfigModel | dict + ) -> LocalContactConfigModel: validated = LocalContactConfigModel.model_validate(config) try: redis_key = f"{self.__bl}:local_contact_config" @@ -1302,10 +1358,12 @@ class BeamlineConfig: def local_contact_config(self, value: LocalContactConfigModel | dict) -> None: self.set_local_contact_config(value) + if __name__ == "__main__": - from aare.common.beamline import mx_beamline + from aarecommon.modelconfigs.beamline import mx_beamline + cfg = BeamlineConfig(bl=mx_beamline()) # cfg.allow_non_staff_request_from_staff = True # cfg.state_busy = False - #fg.abr_meas_pos = AerotechCoordinate(at_mm=Coordinate(x=0.0,y=0.0,z=0.0)) - cfg.dtz_safe_position = 300.0 \ No newline at end of file + # fg.abr_meas_pos = AerotechCoordinate(at_mm=Coordinate(x=0.0,y=0.0,z=0.0)) + cfg.dtz_safe_position = 300.0 diff --git a/src/aare/daq/daq.py b/src/aare/daq/daq.py index 3fa09514..88ba79fe 100644 --- a/src/aare/daq/daq.py +++ b/src/aare/daq/daq.py @@ -5,117 +5,129 @@ import time from datetime import datetime, timezone from math import ceil from pathlib import Path -from typing import List, Tuple, Optional, Callable +from typing import Callable, List, Optional, Tuple import cv2 import numpy as np -from aareDB import SampleEventType -from jfjoch_client import ScanResult, ScanResultImagesInner - -from aare.daq import workflows -from aare.daq.aaredb import AareWrapper - -from aare.daq.config import BeamlineConfig, ABR_POS_MOUNT -from aare.daq.config import BeamlineStateEnum -from aare.daq.devices import BeamlineDevices -from aare.daq.mlbox import MlBox - -from aare.common.tell_models import TellStateModel, TellPhaseEnum -from aare.common.beamline import MXBeamline, cfg_get -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.logger_events import ( +from aarecommon.config.beamline import cfg_get +from aarecommon.config.logger import setup_logger +from aarecommon.config.logger_events import ( log_timing, merge_log_context, raster_request_log_context, rotation_request_log_context, sample_log_context, ) -from aare.common.models import ( - SampleShortInfo, - PuckLoadedInfo, - SampleShortInfoList, AutofocusSettings, - DAQStatusModel, BeamlineStatus, SessionStatus, SampleCameraSettings, ZoomModeEnum, - SimpleScanParameters, FluorescenceSpectrumParameterModel, - FluorescenceSpectrumOutputModel, DAQOperation) -from aare.common.automation_models import ( +from aarecommon.errors.exception_handler import ( + AerotechCommunicationError, + AutoRasterSampleSkipped, + BeamlineBusyException, + BeamlineBusyTimeoutException, + BECCommunicationError, + CriticalTellException, + DataCollectionException, + JFJochCommunicationError, + LoopCenteringFailed, + MagnetPositionSensorErorr, + MaintenanceStateException, + MountingFailed, + RasterScanException, + SmargonCommunicationError, + StateTransitionFailed, + TellCommunicationError, + TransformationInvalidException, + UnmountingFailed, + WarningTellException, +) +from aarecommon.math.coordinate import AerotechCoordinate, Coordinate, SmargonCoordinate +from aarecommon.math.diffraction_geometry import DiffractionGeometry +from aarecommon.math.sample_geometry import SampleGeometryModel +from aarecommon.models.automation import ( AutomationProgress, StepState, StepStatus, WorkflowStateKind, ) -from aare.common.raster_grid import RasterGridRequest, CompletedRasterGrid -from aare.common.rotation_scan import RotationScanRequest, CompletedRotationScan -from aare.common.sample_geometry import SampleGeometryModel -from aare.common.exception_handler import ( - TransformationInvalidException, - StateTransitionFailed, - MaintenanceStateException, - LoopCenteringFailed, - MountingFailed, - WarningTellException, - CriticalTellException, - SmargonCommunicationError, - TellCommunicationError, - BECCommunicationError, - JFJochCommunicationError, - AerotechCommunicationError, - MagnetPositionSensorErorr, - UnmountingFailed, - DataCollectionException, - RasterScanException, - BeamlineBusyTimeoutException, - BeamlineBusyException, - AutoRasterSampleSkipped +from aarecommon.models.beamline import MXBeamline +from aarecommon.models.models import ( + AutofocusSettings, + BeamlineStatus, + DAQOperation, + DAQStatusModel, + FluorescenceSpectrumOutputModel, + FluorescenceSpectrumParameterModel, + PuckLoadedInfo, + SampleCameraSettings, + SampleShortInfo, + SampleShortInfoList, + SessionStatus, + SimpleScanParameters, + ZoomModeEnum, ) +from aarecommon.models.raster_grid import CompletedRasterGrid, RasterGridRequest +from aarecommon.models.rotation_scan import CompletedRotationScan, RotationScanRequest +from aarecommon.models.tell import TellPhaseEnum, TellStateModel +from aareDB import SampleEventType +from jfjoch_client import ScanResult, ScanResultImagesInner + +from aare.daq import workflows +from aare.daq.aaredb import AareWrapper +from aare.daq.config import ABR_POS_MOUNT, BeamlineConfig, BeamlineStateEnum from aare.daq.config_model import LocalContactConfigModel -from aare.daq.operations.face_detection import FaceDetectionContext, FaceDetectionService, FaceDetectionResult +from aare.daq.devices import BeamlineDevices +from aare.daq.mlbox import MlBox +from aare.daq.operations.common.ml_bounding_box import get_ml_bounding_box +from aare.daq.operations.common.runtime import DAQRuntimeState +from aare.daq.operations.common.services import ( + DataCollectionPreparer, + FaceDetectionProgressEmitter, + OperationServices, + PredictionProvider, + SampleEventPublisher, + ScanIngestionService, + StateController, + TraceWriter, +) +from aare.daq.operations.common.simulate_scan_result import ( + build_fake_rotation_result, +) +from aare.daq.operations.face_detection import ( + FaceDetectionContext, + FaceDetectionResult, + FaceDetectionService, +) from aare.daq.operations.face_detection.models import ( FaceDetectionDependencies, FaceDetectionSettings, ) -from aare.daq.operations.loop_centering import LoopCenteringService, LoopCenteringContext +from aare.daq.operations.loop_centering import ( + LoopCenteringContext, + LoopCenteringService, +) from aare.daq.operations.loop_centering.models import ( + LoopCenteringDependencies, LoopCenteringSettings, - LoopCenteringDependencies, ) -from aare.daq.operations.mounting.service import MountingService from aare.daq.operations.mounting.models import ( - MountingResult, MountingContext, MountingDependencies, + MountingResult, MountingSettings, ) +from aare.daq.operations.mounting.service import MountingService from aare.daq.operations.raster.models import ( RasterContext, - RasterSettings, RasterDependencies, + RasterSettings, ) from aare.daq.operations.raster.service import RasterService -from aare.daq.operations.common.ml_bounding_box import get_ml_bounding_box from aare.daq.operations.rotation.models import ( - RotationSettings, - RotationDependencies, RotationContext, + RotationDependencies, + RotationSettings, ) from aare.daq.operations.rotation.service import RotationService from aare.daq.operations.screenshot.service import ScreenshotService -from aare.daq.operations.common.simulate_scan_result import ( - build_fake_rotation_result, -) -from aare.daq.operations.common.runtime import DAQRuntimeState -from aare.daq.operations.common.services import ( - FaceDetectionProgressEmitter, - OperationServices, - PredictionProvider, - StateController, - TraceWriter, - SampleEventPublisher, - ScanIngestionService, - DataCollectionPreparer, -) - from aare.devices.area_detector import AutoEnum from aare.devices.jfjoch import JFJochWrapper @@ -151,7 +163,9 @@ class _DAQNonCriticalRunner: description: str, sample: SampleShortInfo | None = None, ) -> object | None: - return self._daq._run_noncritical(action, description=description, sample=sample) + return self._daq._run_noncritical( + action, description=description, sample=sample + ) class _DAQScreenshotSampleProvider: @@ -201,7 +215,9 @@ class _DAQSampleEventSender: def __init__(self, daq: "AareDAQ"): self._daq = daq - def send_sample_event(self, sample_id: int, event_type, comment: str | None = None) -> None: + def send_sample_event( + self, sample_id: int, event_type, comment: str | None = None + ) -> None: self._daq._AareDAQ__aare.send_sample_event(sample_id, event_type, comment) @@ -217,7 +233,9 @@ class _DAQScanIngestor: beam_mark_pxl=beam_mark_pxl, ) - def ingest_gridscan(self, *, sample, raster_result, raster_request, geom, com, beam_mark_pxl) -> None: + def ingest_gridscan( + self, *, sample, raster_result, raster_request, geom, com, beam_mark_pxl + ) -> None: self._daq._AareDAQ__aare.ingest_gridscan( sample=sample, raster_result=raster_result, @@ -258,8 +276,8 @@ class _FaceDetectionProgressReporter: self._daq._emit_face_detection_progress(payload) -#TODO tidy up DAQ - migrate functions into different scripts, to reduce size? -#TODO investigate using a state machine within each operation to reduce callbacks? +# TODO tidy up DAQ - migrate functions into different scripts, to reduce size? +# TODO investigate using a state machine within each operation to reduce callbacks? class AareDAQ: """ Main Data Acquisition class for the Aare system. @@ -285,7 +303,9 @@ class AareDAQ: self._beamline = bl self.__aare = AareWrapper(bl) self.__saved_box = None - self._smargon_trace_path = Path("/sls/mx/applications/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._automation_progress_cb: Callable[[AutomationProgress], None] | None = None self._last_sample_sync_ts = 0.0 @@ -383,8 +403,14 @@ class AareDAQ: def restart_detector(self) -> dict[str, object]: self.__cfg.try_set_busy(timeout=360) try: - beamline = MXBeamline.SIMULATED if self.__cfg.simulated_detector else self._beamline - logger.info(f"Restarting JFJoch wrapper with simulated={self.__cfg.simulated_detector}") + beamline = ( + MXBeamline.SIMULATED + if self.__cfg.simulated_detector + else self._beamline + ) + logger.info( + f"Restarting JFJoch wrapper with simulated={self.__cfg.simulated_detector}" + ) self.__jfjoch = JFJochWrapper(beamline) return { "ok": True, @@ -449,7 +475,9 @@ class AareDAQ: self.__cfg.simulate_smargon = enabled return self.restart_smargon() - raise ValueError("Unknown simulation device. Expected one of: bec, detector, tell, aerotech, smargon") + raise ValueError( + "Unknown simulation device. Expected one of: bec, detector, tell, aerotech, smargon" + ) def _is_hardware_failure(self, error: Exception) -> bool: return isinstance( @@ -482,18 +510,17 @@ class AareDAQ: ) def _raise_if_critical_jfjoch_detector_error( - self, - error: Exception, - *, - command: str, + self, + error: Exception, + *, + command: str, ) -> None: if not isinstance(error, JFJochCommunicationError): return - if ( - getattr(error, "status_code", None) != 500 - and not self._is_jfjoch_detector_state_error(str(error)) - ): + if getattr( + error, "status_code", None + ) != 500 and not self._is_jfjoch_detector_state_error(str(error)): return message = ( @@ -513,10 +540,10 @@ class AareDAQ: ) from error def _raise_if_critical_bec_error( - self, - error: Exception, - *, - command: str, + self, + error: Exception, + *, + command: str, ) -> None: if not isinstance(error, BECCommunicationError): return @@ -550,11 +577,11 @@ class AareDAQ: ) def _run_noncritical( - self, - action: Callable[[], object], - *, - description: str, - sample: SampleShortInfo | None = None, + self, + action: Callable[[], object], + *, + description: str, + sample: SampleShortInfo | None = None, ) -> object | None: try: return action() @@ -570,11 +597,11 @@ class AareDAQ: return None def _run_critical( - self, - action: Callable[[], object], - *, - description: str, - sample: SampleShortInfo | None = None, + self, + action: Callable[[], object], + *, + description: str, + sample: SampleShortInfo | None = None, ) -> object: try: return action() @@ -623,7 +650,9 @@ class AareDAQ: return [] @staticmethod - def _tell_phase_confirms_previous_sample_unmounted(phase: TellPhaseEnum | None) -> bool: + def _tell_phase_confirms_previous_sample_unmounted( + phase: TellPhaseEnum | None, + ) -> bool: return phase in { TellPhaseEnum.OLD_SAMPLE_RETURNED, TellPhaseEnum.PICKING_NEW_SAMPLE, @@ -640,7 +669,9 @@ class AareDAQ: state_ts is not None and state_ts >= started_at and tell_state.operation == "mount" - and self._tell_phase_confirms_previous_sample_unmounted(tell_state.phase) + and self._tell_phase_confirms_previous_sample_unmounted( + tell_state.phase + ) ): return True @@ -662,12 +693,14 @@ class AareDAQ: return False - def set_face_detection_progress_callback(self, cb: Callable[[dict], None] | None) -> None: + def set_face_detection_progress_callback( + self, cb: Callable[[dict], None] | None + ) -> None: self._face_detection_progress_cb = cb def set_automation_progress_callback( - self, - cb: Callable[[AutomationProgress], None] | None, + self, + cb: Callable[[AutomationProgress], None] | None, ) -> None: self._automation_progress_cb = cb @@ -688,13 +721,13 @@ class AareDAQ: logger.warning(f"Failed to emit automation progress: {e}") def _record_best_effort_step_failure( - self, - *, - progress: AutomationProgress, - step: WorkflowStateKind, - error: Exception, - sample: SampleShortInfo | None, - code: str, + self, + *, + progress: AutomationProgress, + step: WorkflowStateKind, + error: Exception, + sample: SampleShortInfo | None, + code: str, ) -> None: reason = str(error) event_context = merge_log_context( @@ -726,17 +759,23 @@ class AareDAQ: current_step=None, steps=[ StepState(step=WorkflowStateKind.MOUNT, status=StepStatus.PENDING), - StepState(step=WorkflowStateKind.LOOP_CENTRE, status=StepStatus.PENDING), + StepState( + step=WorkflowStateKind.LOOP_CENTRE, status=StepStatus.PENDING + ), StepState(step=WorkflowStateKind.RASTER, status=StepStatus.PENDING), - StepState(step=WorkflowStateKind.DATA_COLLECTION, status=StepStatus.PENDING), + StepState( + step=WorkflowStateKind.DATA_COLLECTION, status=StepStatus.PENDING + ), StepState(step=WorkflowStateKind.FINAL, status=StepStatus.PENDING), ], finished=False, success=None, samples_in_queue=self._automation_samples_in_queue, avg_time_per_sample=( - self._automation_total_sample_time_s / self._automation_completed_samples - if self._automation_completed_samples > 0 else 0.0 + self._automation_total_sample_time_s + / self._automation_completed_samples + if self._automation_completed_samples > 0 + else 0.0 ), current_sample_name=self._automation_last_sample_name, ) @@ -753,9 +792,9 @@ class AareDAQ: return labels.get(step, str(step.value)) def _get_progress_step( - self, - progress: AutomationProgress, - step: WorkflowStateKind, + self, + progress: AutomationProgress, + step: WorkflowStateKind, ) -> StepState | None: for item in progress.steps: if item.step == step: @@ -763,11 +802,11 @@ class AareDAQ: return None def _set_progress_context( - self, - progress: AutomationProgress, - *, - current_sample_name: str | None = None, - samples_in_queue: int | None = None, + self, + progress: AutomationProgress, + *, + current_sample_name: str | None = None, + samples_in_queue: int | None = None, ) -> None: if current_sample_name is not None: safe_name = str(current_sample_name or "") @@ -783,33 +822,34 @@ class AareDAQ: if self._automation_completed_samples > 0: progress.avg_time_per_sample = ( - self._automation_total_sample_time_s / self._automation_completed_samples + self._automation_total_sample_time_s + / self._automation_completed_samples ) else: progress.avg_time_per_sample = 0.0 def _record_completed_sample_time( - self, - progress: AutomationProgress, - elapsed_s: float, + self, + progress: AutomationProgress, + elapsed_s: float, ) -> None: if elapsed_s <= 0: return self._automation_completed_samples += 1 self._automation_total_sample_time_s += float(elapsed_s) progress.avg_time_per_sample = ( - self._automation_total_sample_time_s / self._automation_completed_samples + self._automation_total_sample_time_s / self._automation_completed_samples ) def _set_progress_step( - self, - progress: AutomationProgress, - step: WorkflowStateKind, - status: StepStatus, - message: str = "", - *, - make_current: bool = False, - error_code: str | None = None, + self, + progress: AutomationProgress, + step: WorkflowStateKind, + status: StepStatus, + message: str = "", + *, + make_current: bool = False, + error_code: str | None = None, ) -> None: now = time.time() @@ -840,10 +880,10 @@ class AareDAQ: self._emit_automation_progress(progress) def _mark_progress_running( - self, - progress: AutomationProgress, - step: WorkflowStateKind, - message: str = "", + self, + progress: AutomationProgress, + step: WorkflowStateKind, + message: str = "", ) -> None: self._set_progress_step( progress, @@ -854,18 +894,18 @@ class AareDAQ: ) def _mark_progress_success( - self, - progress: AutomationProgress, - step: WorkflowStateKind, - message: str = "", + self, + progress: AutomationProgress, + step: WorkflowStateKind, + message: str = "", ) -> None: self._set_progress_step(progress, step, StepStatus.SUCCESS, message) def _mark_progress_failed( - self, - progress: AutomationProgress, - step: WorkflowStateKind, - message: str = "", + self, + progress: AutomationProgress, + step: WorkflowStateKind, + message: str = "", ) -> None: self._set_progress_step( progress, @@ -888,21 +928,23 @@ class AareDAQ: self._emit_automation_progress(progress) def _mark_progress_finished( - self, - progress: AutomationProgress, - success: bool, - message: str = "", + self, + progress: AutomationProgress, + success: bool, + message: str = "", ) -> None: final_status = StepStatus.SUCCESS if success else StepStatus.FAILED - self._set_progress_step(progress, WorkflowStateKind.FINAL, final_status, message) + self._set_progress_step( + progress, WorkflowStateKind.FINAL, final_status, message + ) progress.current_step = self._step_display_name(WorkflowStateKind.FINAL) progress.finished = True progress.success = success self._emit_automation_progress(progress) - #-------------------------------------------- + # -------------------------------------------- # Operation Services - #-------------------------------------------- + # -------------------------------------------- def _build_runtime_state(self) -> DAQRuntimeState: return DAQRuntimeState( @@ -1026,17 +1068,17 @@ class AareDAQ: logger=logger, ) - #-------------------------------------------- + # -------------------------------------------- # Operation Handlers - #-------------------------------------------- - #TODO make sure this is implemented currectly + # -------------------------------------------- + # TODO make sure this is implemented currectly def _handle_operation_error( - self, - operation: DAQOperation, - sample: Optional[SampleShortInfo], - error: Exception, - event_type: SampleEventType = SampleEventType.FAILED, - additional_comment: Optional[str] = None, + self, + operation: DAQOperation, + sample: Optional[SampleShortInfo], + error: Exception, + event_type: SampleEventType = SampleEventType.FAILED, + additional_comment: Optional[str] = None, ) -> None: """ Centralized databse maessage error handling for all operations. @@ -1085,11 +1127,11 @@ class AareDAQ: }, ) - #todo make sure these functions are correctly implemented! - #Operation handlers should handle database communication and beamline state changes, - #where possible/between operations. For example in raster, we may use XtalSnapshot to take a screenshot, then - #return to Datacollection. - #Busy stats are handled by the public function call i.e. sample() or by automation i.e. measure() + # todo make sure these functions are correctly implemented! + # Operation handlers should handle database communication and beamline state changes, + # where possible/between operations. For example in raster, we may use XtalSnapshot to take a screenshot, then + # return to Datacollection. + # Busy stats are handled by the public function call i.e. sample() or by automation i.e. measure() def _execute_mount_and_prepare(self, sample: SampleShortInfo | None) -> bool: """ Operation handler for executing mounting and take screenshot. @@ -1110,10 +1152,14 @@ class AareDAQ: clear_cached_on_empty=False, ) except Exception as sync_error: - logger.warning(f"Failed to reconcile previous sample from TELL before mount: {sync_error}") + logger.warning( + f"Failed to reconcile previous sample from TELL before mount: {sync_error}" + ) if previous_sample is not None and previous_sample.db_id is not None: - self.__aare.send_sample_event(previous_sample.db_id, SampleEventType.UNMOUNTING) + self.__aare.send_sample_event( + previous_sample.db_id, SampleEventType.UNMOUNTING + ) self.__set_state(BeamlineStateEnum.RobotSampleExchange) @@ -1121,7 +1167,9 @@ class AareDAQ: self.__aare.send_sample_event(sample.db_id, SampleEventType.MOUNTING) self.__devs.tell.blower_on() - mounting_result: MountingResult = self._create_mounting_service().execute(target=sample) + mounting_result: MountingResult = self._create_mounting_service().execute( + target=sample + ) if not mounting_result.success: raise mounting_result.error or MountingFailed( @@ -1132,17 +1180,23 @@ class AareDAQ: mounted_sample = mounting_result.mounted_sample or sample if ( - mounting_result.did_unmount_previous - and previous_sample is not None - and previous_sample.db_id is not None + mounting_result.did_unmount_previous + and previous_sample is not None + and previous_sample.db_id is not None ): - self.__aare.send_sample_event(previous_sample.db_id, SampleEventType.UNMOUNTED) + self.__aare.send_sample_event( + previous_sample.db_id, SampleEventType.UNMOUNTED + ) self.__set_state(BeamlineStateEnum.SampleAlignment) if mounted_sample is not None and mounted_sample.db_id is not None: - self.__aare.send_sample_event(mounted_sample.db_id, SampleEventType.MOUNTED) - self.save_screenshot_db(mounted_sample.db_id, f"{mounted_sample.db_id}_mounted") + self.__aare.send_sample_event( + mounted_sample.db_id, SampleEventType.MOUNTED + ) + self.save_screenshot_db( + mounted_sample.db_id, f"{mounted_sample.db_id}_mounted" + ) return True @@ -1150,7 +1204,9 @@ class AareDAQ: self._last_mount_error_message = str(e) or "Mount failed" logger.error(f"Mount failed due to invalid transformation: {e}") self._handle_operation_error( - operation=DAQOperation.MOUNT if sample is not None else DAQOperation.UNMOUNT, + operation=DAQOperation.MOUNT + if sample is not None + else DAQOperation.UNMOUNT, sample=sample, error=e, event_type=SampleEventType.MOUNTFAILED, @@ -1161,7 +1217,9 @@ class AareDAQ: self._last_mount_error_message = str(e) or "Mount failed" logger.error(f"Tell communication error occured: {e}") self._handle_operation_error( - operation=DAQOperation.MOUNT if sample is not None else DAQOperation.UNMOUNT, + operation=DAQOperation.MOUNT + if sample is not None + else DAQOperation.UNMOUNT, sample=sample, error=e, event_type=SampleEventType.MOUNTFAILED, @@ -1173,9 +1231,9 @@ class AareDAQ: logger.error(f"Mount failed: {e}") previous_sample_unmounted = ( - previous_sample is not None - and previous_sample.db_id is not None - and self._was_previous_sample_unmounted_since(mount_started_at) + previous_sample is not None + and previous_sample.db_id is not None + and self._was_previous_sample_unmounted_since(mount_started_at) ) if previous_sample_unmounted: @@ -1196,10 +1254,12 @@ class AareDAQ: logger.exception("Failed to set state to SampleAlignment") finally: self._handle_operation_error( - operation=DAQOperation.MOUNT if sample is not None else DAQOperation.UNMOUNT, + operation=DAQOperation.MOUNT + if sample is not None + else DAQOperation.UNMOUNT, sample=sample, error=e, - event_type=SampleEventType.MOUNTFAILED + event_type=SampleEventType.MOUNTFAILED, ) return False @@ -1233,7 +1293,9 @@ class AareDAQ: sample=sample, error=result.error or Exception("Loop centering failed"), event_type=SampleEventType.ALCFAILED, - additional_comment=result.comment if result.comment is not None else "", + additional_comment=result.comment + if result.comment is not None + else "", ) return False @@ -1244,7 +1306,9 @@ class AareDAQ: except Exception as e: logger.error(f"Loop centering failed: {e}") if result is not None and result.error is not None: - additional_comment = result.comment if result.comment is not None else "" + additional_comment = ( + result.comment if result.comment is not None else "" + ) self._handle_operation_error( operation=DAQOperation.LOOP_CENTERING, sample=sample, @@ -1255,12 +1319,12 @@ class AareDAQ: return False def _execute_face_detection( - self, - steps: int = 14, - step_size: int = 15, - face_min_ratio: float = 0.3, - report_error: bool = True, - sample: Optional[SampleShortInfo] = None, + self, + steps: int = 14, + step_size: int = 15, + face_min_ratio: float = 0.3, + report_error: bool = True, + sample: Optional[SampleShortInfo] = None, ) -> FaceDetectionResult: """ Execute face detection sequence through the face detection service. @@ -1275,7 +1339,9 @@ class AareDAQ: if sample is None: try: sample = self.sample - logger.debug(f"No sample provided, using current sample from DAQ {sample}") + logger.debug( + f"No sample provided, using current sample from DAQ {sample}" + ) except Exception: logger.error("Failed to get current sample") sample = None @@ -1340,8 +1406,9 @@ class AareDAQ: comment=additional_comment, ) - def _execute_raster_sequence(self, grid_request: RasterGridRequest, - auto_center: bool = False) -> CompletedRasterGrid | None: + def _execute_raster_sequence( + self, grid_request: RasterGridRequest, auto_center: bool = False + ) -> CompletedRasterGrid | None: """ Execute raster scan with optional auto-centering. @@ -1389,7 +1456,9 @@ class AareDAQ: self.__jfjoch.measure_raster(grid_request, status) logger.info("detector initialised") else: - logger.info("Simulated detector mode enabled; using fake raster result.") + logger.info( + "Simulated detector mode enabled; using fake raster result." + ) self.__set_state(BeamlineStateEnum.DataCollection) raster_result = raster_service.execute(grid_request) @@ -1404,13 +1473,17 @@ class AareDAQ: raster_request_log_context(grid_request), { "auto_center": auto_center, - "result_count": len(result.r) if hasattr(result, "r") and result.r is not None else None, + "result_count": len(result.r) + if hasattr(result, "r") and result.r is not None + else None, }, ), ) else: if self.sample is not None and self.sample.db_id is not None: - self.__aare.send_sample_event(self.sample.db_id, SampleEventType.RASTERINGFAILED) + self.__aare.send_sample_event( + self.sample.db_id, SampleEventType.RASTERINGFAILED + ) logger.error( "Raster sequence returned no result", extra=merge_log_context( @@ -1420,7 +1493,6 @@ class AareDAQ: ) return result - except JFJochCommunicationError as e: logger.exception( "Raster sequence failed due to JFJoch communication error", @@ -1434,7 +1506,7 @@ class AareDAQ: sample=self.sample, error=e, event_type=SampleEventType.RASTERINGFAILED, - additional_comment=f"JFJoch communication error: {e}" + additional_comment=f"JFJoch communication error: {e}", ) raise @@ -1476,7 +1548,7 @@ class AareDAQ: sample=self.sample, error=e, event_type=SampleEventType.RASTERINGFAILED, - additional_comment=f"Raster sequence failed: {e}" + additional_comment=f"Raster sequence failed: {e}", ) return None @@ -1530,7 +1602,9 @@ class AareDAQ: # ) # raise - def _execute_rotation_sequence(self, rotation_request: RotationScanRequest) -> CompletedRotationScan | None: + def _execute_rotation_sequence( + self, rotation_request: RotationScanRequest + ) -> CompletedRotationScan | None: """ Execute rotation scan. @@ -1543,13 +1617,15 @@ class AareDAQ: try: return self._create_rotation_service().run(rotation_request) except JFJochCommunicationError as e: - logger.error(f"Rotation sequence failed due to JFJoch Communication error: {e}") + logger.error( + f"Rotation sequence failed due to JFJoch Communication error: {e}" + ) self._handle_operation_error( operation=DAQOperation.ROTATION, sample=self.sample, error=e, event_type=SampleEventType.COLLECTIONFAILED, - additional_comment=f"JFJoch communication error: {e}" + additional_comment=f"JFJoch communication error: {e}", ) raise except Exception as e: @@ -1558,7 +1634,7 @@ class AareDAQ: operation=DAQOperation.ROTATION, sample=self.sample, error=e, - event_type=SampleEventType.COLLECTIONFAILED + event_type=SampleEventType.COLLECTIONFAILED, ) raise @@ -1596,19 +1672,21 @@ class AareDAQ: logger.warning(f"Failed to append smargon trace: {e}") def _sample_matches_mounted_address( - self, - sample: SampleShortInfo | None, - mounted_address, + self, + sample: SampleShortInfo | None, + mounted_address, ) -> bool: if sample is None or sample.location is None or mounted_address is None: return False return ( - sample.location.segment == mounted_address.puck.segment - and sample.location.pos == mounted_address.puck.pos - and sample.pin == mounted_address.pin + sample.location.segment == mounted_address.puck.segment + and sample.location.pos == mounted_address.puck.pos + and sample.pin == mounted_address.pin ) - def _find_sample_by_mounted_address(self, mounted_address) -> SampleShortInfo | None: + def _find_sample_by_mounted_address( + self, mounted_address + ) -> SampleShortInfo | None: for sample in self.__cfg.spreadsheet.s: if self._sample_matches_mounted_address(sample, mounted_address): return sample @@ -1619,7 +1697,9 @@ class AareDAQ: return None - def _placeholder_sample_from_mounted_address(self, mounted_address) -> SampleShortInfo: + def _placeholder_sample_from_mounted_address( + self, mounted_address + ) -> SampleShortInfo: return SampleShortInfo( db_id=-1, puck_name="", @@ -1643,9 +1723,9 @@ class AareDAQ: return "sample" def sync_current_sample_from_tell( - self, - force: bool = False, - clear_cached_on_empty: bool = True, + self, + force: bool = False, + clear_cached_on_empty: bool = True, ) -> SampleShortInfo | None: current_sample = self.__cfg.current_sample @@ -1653,7 +1733,10 @@ class AareDAQ: return current_sample now = time.monotonic() - if not force and (now - self._last_sample_sync_ts) < self._sample_sync_min_interval_s: + if ( + not force + and (now - self._last_sample_sync_ts) < self._sample_sync_min_interval_s + ): return current_sample self._last_sample_sync_ts = now @@ -1662,7 +1745,9 @@ class AareDAQ: if mounted_address is None: if current_sample is not None and current_sample.location is not None: if clear_cached_on_empty: - logger.warning("TELL reports no mounted sample; clearing cached current_sample") + logger.warning( + "TELL reports no mounted sample; clearing cached current_sample" + ) self.__cfg.current_sample = None else: logger.warning( @@ -1675,7 +1760,9 @@ class AareDAQ: resolved_sample = self._find_sample_by_mounted_address(mounted_address) if resolved_sample is None: - resolved_sample = self._placeholder_sample_from_mounted_address(mounted_address) + resolved_sample = self._placeholder_sample_from_mounted_address( + mounted_address + ) logger.warning( "Mounted sample from TELL was not found in known sample lists; using placeholder", extra={"mounted_address": str(mounted_address)}, @@ -1700,7 +1787,10 @@ class AareDAQ: @state.setter def state(self, target: BeamlineStateEnum): if target == BeamlineStateEnum.Moving: - logger.error(f"Cannot explicitly move to busy state", extra={"target":target, "state":self.__cfg.state}) + logger.error( + f"Cannot explicitly move to busy state", + extra={"target": target, "state": self.__cfg.state}, + ) raise Exception("Cannot explicitly move to busy state") start = time.perf_counter() @@ -1714,41 +1804,47 @@ class AareDAQ: self.last_time = end - start - def spreadsheet_params(self) -> tuple[Optional[SimpleScanParameters], str|None]: + def spreadsheet_params(self) -> tuple[Optional[SimpleScanParameters], str | None]: file_prefix = None if self.status.sample is None: return None, file_prefix logger.debug(f"generate params for sample: {self.status.sample}") - aaredb_params = self.status.sample.aaredb_params if hasattr(self.status.sample, "aaredb_params") else None + aaredb_params = ( + self.status.sample.aaredb_params + if hasattr(self.status.sample, "aaredb_params") + else None + ) if aaredb_params is None: return None, file_prefix - #if aaredb_params.directory: + # if aaredb_params.directory: # file_prefix = aaredb_params.directory if ( - getattr(aaredb_params, 'exposure', None) is None - and getattr(aaredb_params, 'transmission', None) is None - and getattr(aaredb_params, 'oscillation', None) is None - and getattr(aaredb_params, 'totalrange', None) is None - and getattr(aaredb_params, 'targetresolution', None) is None + getattr(aaredb_params, "exposure", None) is None + and getattr(aaredb_params, "transmission", None) is None + and getattr(aaredb_params, "oscillation", None) is None + and getattr(aaredb_params, "totalrange", None) is None + and getattr(aaredb_params, "targetresolution", None) is None ): return None, file_prefix params = SimpleScanParameters() - if (exp := getattr(aaredb_params, 'exposure', None)) is not None: + if (exp := getattr(aaredb_params, "exposure", None)) is not None: params.exp_time_s = exp - if (trans := getattr(aaredb_params, 'transmission', None)) is not None: + if (trans := getattr(aaredb_params, "transmission", None)) is not None: logger.debug(f"transmission: {trans}") params.transmission = trans / 100.0 if trans > 1.0 else trans - if (res := getattr(aaredb_params, 'targetresolution', None)) is not None: + if (res := getattr(aaredb_params, "targetresolution", None)) is not None: logger.debug(f"resolution: {res}") - logger.debug(f"requested dtz: {self.diffraction_geometry.calc_dtz_mm(res)} ") + logger.debug( + f"requested dtz: {self.diffraction_geometry.calc_dtz_mm(res)} " + ) new_res = 1 / ((1 / res) + 0.1) corrected_dtz = self.diffraction_geometry.calc_dtz_mm(new_res) logger.debug(f"corrected dtz: {corrected_dtz}") @@ -1756,8 +1852,8 @@ class AareDAQ: corrected_dtz = 108 params.dtz = round(corrected_dtz) - osc = getattr(aaredb_params, 'oscillation', None) - total = getattr(aaredb_params, 'totalrange', None) + osc = getattr(aaredb_params, "oscillation", None) + total = getattr(aaredb_params, "totalrange", None) if osc is not None: osc = abs(osc) @@ -1772,12 +1868,11 @@ class AareDAQ: params.incr_omega_deg = default_osc params.steps = round(abs(total) / default_osc) - #if file_prefix is not None: + # if file_prefix is not None: # params.file_prefix = file_prefix return params, file_prefix - @property def omega(self) -> float: return self.__devs.aerotech_omega @@ -1819,7 +1914,9 @@ class AareDAQ: # how/when beam-location was entered, rather than a separate flag that # must be set exactly on the state transition. beam_location = self.__cfg.state == BeamlineStateEnum.BeamLocation - logger.debug(f"zoom -> {val} (state={self.__cfg.state}, beam_location_presets={beam_location})") + logger.debug( + f"zoom -> {val} (state={self.__cfg.state}, beam_location_presets={beam_location})" + ) if beam_location: # Auto-exposure is too dark to see the beam at high zoom, so apply # the preset per-zoom gain/exposure (interpolated between the stored @@ -1830,7 +1927,9 @@ class AareDAQ: time.sleep(0.2) settings = self.__cfg.zoom_settings.get_camera_settings(val) self.__devs.samcam_settings = settings - logger.debug(f"applied beam-location preset for zoom {val}: gain={settings.gain}, exposure={settings.exposure}") + logger.debug( + f"applied beam-location preset for zoom {val}: gain={settings.gain}, exposure={settings.exposure}" + ) else: self.__devs.samcam_auto(AutoEnum.AUTO) self.__devs.zoom = val @@ -1842,7 +1941,9 @@ class AareDAQ: the beam-location preset (Redis), so it is re-applied on future zooms.""" zoom_value = self.zoom settings = self.samcam_settings - self.__cfg.save_zoom_camera_setting(zoom_value, settings, mode=ZoomModeEnum.BeamLocation) + self.__cfg.save_zoom_camera_setting( + zoom_value, settings, mode=ZoomModeEnum.BeamLocation + ) logger.info( f"Saved beam-location camera setting for zoom {zoom_value}: " f"gain={settings.gain}, exposure={settings.exposure}" @@ -1895,7 +1996,9 @@ class AareDAQ: def tweak_abr_meas_pos(self, c: AerotechCoordinate): self.__cfg.set_busy(BeamlineStateEnum.SampleAlignment) try: - new_meas_pos = AerotechCoordinate(at_mm=self.__cfg.abr_meas_pos.at_mm + c.at_mm) + 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 @@ -1907,7 +2010,9 @@ class AareDAQ: def save_abr_meas_pos(self): self.__cfg.set_busy(BeamlineStateEnum.SampleAlignment) try: - self.__cfg.abr_meas_pos = AerotechCoordinate(at_mm=self.__devs.aerotech_pos.at_mm) + self.__cfg.abr_meas_pos = AerotechCoordinate( + at_mm=self.__devs.aerotech_pos.at_mm + ) self.__devs.bec_worker.save_current_aerotech_position() self.__cfg.state_busy = False except Exception: @@ -1943,10 +2048,10 @@ class AareDAQ: def check_tell_mount_start_conditions(self) -> None: self.__devs.tell.validate_mount_start_conditions() - def _execute_dry(self, park:bool=True, unmount:bool=False): + def _execute_dry(self, park: bool = True, unmount: bool = False): self._create_mounting_service().dry(park=park, unmount=unmount) - def park_and_dry(self, park = True, unmount: bool = False): + def park_and_dry(self, park=True, unmount: bool = False): self.__cfg.try_set_busy(timeout=360) try: self.__set_state(BeamlineStateEnum.RobotSampleExchange) @@ -1962,10 +2067,10 @@ class AareDAQ: raise def tell_toggle_blower(self): - try: - self.__devs.tell.toggle_blower() - except Exception as e: - logger.error(f"Failed to turn off blower: {e}") + try: + self.__devs.tell.toggle_blower() + except Exception as e: + logger.error(f"Failed to turn off blower: {e}") def initialise_smargon(self): self.__cfg.try_set_busy(timeout=360) @@ -2006,7 +2111,6 @@ class AareDAQ: def sample(self, target: SampleShortInfo | None): self.__cfg.try_set_busy(timeout=360) try: - logger.debug(f"Mount target {target}") if target is None: @@ -2020,8 +2124,12 @@ class AareDAQ: current_sample = target or self.__cfg.current_sample sample_name = self._sample_mount_display_name(current_sample) if target is None: - raise UnmountingFailed(f"Failed to {operation_name.lower()} {sample_name}") - raise MountingFailed(f"Failed to {operation_name.lower()} {sample_name}") + raise UnmountingFailed( + f"Failed to {operation_name.lower()} {sample_name}" + ) + raise MountingFailed( + f"Failed to {operation_name.lower()} {sample_name}" + ) logger.info(f"Sample operation completed: {target}") self.__cfg.state_busy = False @@ -2034,8 +2142,10 @@ class AareDAQ: def list_loaded_pucks(self) -> List[PuckLoadedInfo]: return [] - def __auto_focus(self, settings: AutofocusSettings, settle_time_s: float = 1.0) -> float: - #TODO uses old code change + def __auto_focus( + self, settings: AutofocusSettings, settle_time_s: float = 1.0 + ) -> float: + # TODO uses old code change """ Scan smargon Z and find the position with maximum focus measure. @@ -2071,7 +2181,9 @@ class AareDAQ: def auto_exposure(self): self.__devs.samcam_auto(AutoEnum.ONCE) - def __setup_datacollection(self, request: RasterGridRequest | RotationScanRequest, screening: bool = False): + def __setup_datacollection( + self, request: RasterGridRequest | RotationScanRequest, screening: bool = False + ): request_omega = getattr(request, "omega_deg", None) if request_omega is None: request_omega = getattr(request, "start_omega_deg", None) @@ -2082,15 +2194,18 @@ class AareDAQ: if request.dtz is not None: if request.dtz < self.__cfg.cached_dtz_low: - raise ValueError(f"Requested DTZ {request.dtz} is less than low detector limit {self.__cfg.cached_dtz_low}") + raise ValueError( + f"Requested DTZ {request.dtz} is less than low detector limit {self.__cfg.cached_dtz_low}" + ) elif request.dtz > self.__cfg.cached_dtz_high: - raise ValueError(f"Requested DTZ {request.dtz} exceeds high detector limit {self.__cfg.cached_dtz_high}") - logger.info(f'requesting dtz to move to {request.dtz}') + raise ValueError( + f"Requested DTZ {request.dtz} exceeds high detector limit {self.__cfg.cached_dtz_high}" + ) + logger.info(f"requesting dtz to move to {request.dtz}") self.__cfg.dtz = request.dtz else: raise ValueError("No DTZ specified") - if self.sample is not None and self.sample.db_id is not None: sample_id = self.sample.db_id if screening: @@ -2104,31 +2219,31 @@ class AareDAQ: self.save_screenshot_db(sample_id, screenshot_name) if request.transmission is not None: - logger.info(f'requesting transmission to move to {request.transmission}') + logger.info(f"requesting transmission to move to {request.transmission}") self.__devs.transmission = request.transmission start_pos = getattr(request, "start", None) if start_pos is not None: - logger.info(f'requesting smargon to move to {start_pos}') + logger.info(f"requesting smargon to move to {start_pos}") self.__devs.smargon_pos = start_pos else: smargon_top_left = getattr(request, "smargon_top_left", None) has_valid_smargon_target = ( - smargon_top_left is not None - and getattr(smargon_top_left, "sh_mm", None) is not None + smargon_top_left is not None + and getattr(smargon_top_left, "sh_mm", None) is not None ) is_zero_placeholder = ( - has_valid_smargon_target - and abs(smargon_top_left.sh_mm.x) < 1e-6 - and abs(smargon_top_left.sh_mm.y) < 1e-6 - and abs(smargon_top_left.sh_mm.z) < 1e-6 - and abs(smargon_top_left.phi_deg) < 1e-6 - and abs(smargon_top_left.chi_deg) < 1e-6 + has_valid_smargon_target + and abs(smargon_top_left.sh_mm.x) < 1e-6 + and abs(smargon_top_left.sh_mm.y) < 1e-6 + and abs(smargon_top_left.sh_mm.z) < 1e-6 + and abs(smargon_top_left.phi_deg) < 1e-6 + and abs(smargon_top_left.chi_deg) < 1e-6 ) if has_valid_smargon_target and not is_zero_placeholder: - logger.info(f'requesting smargon to move to {smargon_top_left}') + logger.info(f"requesting smargon to move to {smargon_top_left}") self.__devs.set_smargon_pos( SmargonCoordinate( sh_mm=smargon_top_left.sh_mm, @@ -2137,15 +2252,19 @@ class AareDAQ: ) ) else: - logger.debug("No explicit smargon target in request; skipping pre-datacollection smargon move") + logger.debug( + "No explicit smargon target in request; skipping pre-datacollection smargon move" + ) self.__devs.smargon_wait(timeout=180) - #todo ADD TRANSMISSION - #if request.transmission is not None: + # todo ADD TRANSMISSION + # if request.transmission is not None: # self.__devs.transmission.wait() return - def _build_fake_rotation_result(self, request: RotationScanRequest) -> CompletedRotationScan: + def _build_fake_rotation_result( + self, request: RotationScanRequest + ) -> CompletedRotationScan: start_angle = 0.0 try: start_angle = float(self.omega) @@ -2157,7 +2276,9 @@ class AareDAQ: start_angle=start_angle, ) - def measure_raster(self, request: RasterGridRequest, auto_center: bool) -> CompletedRasterGrid: + def measure_raster( + self, request: RasterGridRequest, auto_center: bool + ) -> CompletedRasterGrid: """ Execute a raster scan. @@ -2204,10 +2325,11 @@ class AareDAQ: try: self.__jfjoch.wait_till_running(timeout=60.0) except Exception as e: - self._raise_if_critical_jfjoch_detector_error(e, command="wait_till_running") + self._raise_if_critical_jfjoch_detector_error( + e, command="wait_till_running" + ) raise try: - if request.screening: self.__devs.aerotech.screening_scan( rotation_deg=request.steps * request.incr_omega_deg, @@ -2227,7 +2349,9 @@ class AareDAQ: # 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.exp_time_s / float(request.steps) - pos_step = (request.end.sh_mm - request.start.sh_mm) * (1.0 / 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( @@ -2239,13 +2363,17 @@ class AareDAQ: self.__devs.aerotech_omega = omega_start if self.__cfg.simulated_detector: - logger.warning("Detector in simulation mode, returning fake zero rotation result.") + logger.warning( + "Detector in simulation mode, returning fake zero rotation result." + ) return self._build_fake_rotation_result(request) else: try: scan_result = self.__jfjoch.wait_till_done(60) except Exception as e: - self._raise_if_critical_jfjoch_detector_error(e, command="wait_till_done") + self._raise_if_critical_jfjoch_detector_error( + e, command="wait_till_done" + ) raise return CompletedRotationScan( request=copy.deepcopy(request), @@ -2284,16 +2412,23 @@ class AareDAQ: if result is None: logger.error("Rotation scan failed, no result returned") - raise DataCollectionException("Rotation scan failed, no result returned") + raise DataCollectionException( + "Rotation scan failed, no result returned" + ) self.__set_state(BeamlineStateEnum.SampleAlignment) self.__cfg.state_busy = False return result finally: try: - if self.__cfg.state_busy and self.__cfg.state != BeamlineStateEnum.Maintenance: + if ( + self.__cfg.state_busy + and self.__cfg.state != BeamlineStateEnum.Maintenance + ): self.__set_state(BeamlineStateEnum.SampleAlignment) except Exception as cleanup_error: - logger.exception(f"Failed to restore SampleAlignment after rotation: {cleanup_error}") + logger.exception( + f"Failed to restore SampleAlignment after rotation: {cleanup_error}" + ) finally: self.__cfg.state_busy = False @@ -2301,7 +2436,9 @@ class AareDAQ: def dtz(self) -> float: tmp = self.__cfg.dtz if tmp is None: - return cfg_get('daq.data_collection_settings.default_raster_scan_settings.dtz', 200) + return cfg_get( + "daq.data_collection_settings.default_raster_scan_settings.dtz", 200 + ) else: return tmp @@ -2314,12 +2451,16 @@ class AareDAQ: dtz_high = self.__cfg.cached_dtz_high if dtz_low is None or dtz_high is None: - logger.warning("DTZ limits not found in cache, refreshing hardware metadata") + logger.warning( + "DTZ limits not found in cache, refreshing hardware metadata" + ) try: self.refresh_detector_metadata_cache() except Exception as e: self.__cfg.state_busy = False - raise RuntimeError(f"DTZ limits unavailable and refresh failed: {e}") from e + raise RuntimeError( + f"DTZ limits unavailable and refresh failed: {e}" + ) from e dtz_low = self.__cfg.cached_dtz_low dtz_high = self.__cfg.cached_dtz_high @@ -2383,7 +2524,7 @@ class AareDAQ: smargon=self.__devs.smargon_pos, beam_size_mm=self.__cfg.beam_size_mm, aerotech=aerotech_pos - aerotech_pos_ref, - aerotech_meas=aerotech_pos + aerotech_meas=aerotech_pos, ) return sample_geom @@ -2406,7 +2547,9 @@ class AareDAQ: def get_beam_mark(self): return self.__cfg.get_beam_mark(self.__devs.zoom) - def ml_bounding_box(self, sample_id: int | None = None, filename: str | None = None) -> RasterGridRequest | None: + def ml_bounding_box( + self, sample_id: int | None = None, filename: str | None = None + ) -> RasterGridRequest | None: """ Request an ML-based bounding box for the sample. @@ -2428,7 +2571,7 @@ class AareDAQ: logger=logger, max_images=self.AUTO_RASTER_MAX_IMAGES, min_cell_size_mm=self.AUTO_RASTER_MAX_IMAGES, - skip_if_exceed_max_image_threshold=self.AUTO_RASTER_SKIP_IF_EXCEED_MAX_IMAGE_THRESHOLD + skip_if_exceed_max_image_threshold=self.AUTO_RASTER_SKIP_IF_EXCEED_MAX_IMAGE_THRESHOLD, ) self.__cfg.state_busy = False return r @@ -2436,7 +2579,9 @@ class AareDAQ: self.__cfg.state_busy = False raise - def face_detection(self, steps: int = 14, step_size: int = 15, face_min_ratio: float = 0.3) -> dict: + def face_detection( + self, steps: int = 14, step_size: int = 15, face_min_ratio: float = 0.3 + ) -> dict: """ Perform a face detection sequence by rotating the sample and using ML to find the flat face. @@ -2461,7 +2606,7 @@ class AareDAQ: self.__cfg.state_busy = False @log_timing(logger, "Auto loop center") - def auto_loop_center(self, sample:Optional[SampleShortInfo]=None) -> float: + def auto_loop_center(self, sample: Optional[SampleShortInfo] = None) -> float: """ Automatically center the loop using ML-based detection. This performs a multi-step sequence including rotation and centering. @@ -2480,9 +2625,11 @@ class AareDAQ: self.__cfg.try_set_busy(timeout=360) if sample is None: if self.sample is None: - raise LoopCenteringFailed("No sample available for loop centering." - "If you have mounted a manual sample," - "please add it in the manual sample panel.") + raise LoopCenteringFailed( + "No sample available for loop centering." + "If you have mounted a manual sample," + "please add it in the manual sample panel." + ) sample = self.sample if not self._execute_loop_centering(sample): @@ -2523,7 +2670,9 @@ class AareDAQ: """ self._screenshot_service.save_local(filename, settle_time_s) - def save_screenshot_db(self, sample_id: int, filename: str, settle_time_s: float = 0.2): + def save_screenshot_db( + self, sample_id: int, filename: str, settle_time_s: float = 0.2 + ): """ Capture a screenshot and upload it to the database for a specific sample. @@ -2534,7 +2683,9 @@ class AareDAQ: """ self._screenshot_service.save_to_db(sample_id, filename, settle_time_s) - def send_screenshot_db(self, filename: str | None = None, message: str | None = None) -> None: + def send_screenshot_db( + self, filename: str | None = None, message: str | None = None + ) -> None: sample = self.sample if sample is None: raise ValueError("No sample with a valid sample_id is mounted.") @@ -2545,7 +2696,9 @@ class AareDAQ: default_message=self._default_screenshot_message(sample.db_id), ) - def send_message_db(self, db_id:int, event_type:SampleEventType, comment: Optional[str] = None): + def send_message_db( + self, db_id: int, event_type: SampleEventType, comment: Optional[str] = None + ): self.__aare.send_sample_event(db_id, event_type, comment) @property @@ -2560,32 +2713,54 @@ class AareDAQ: # sample = copy.deepcopy(self.sample_spreadsheet) #switch to copy if too heavy! # sample.s = list(filter(lambda x: x.user == pgroup, sample.s)) # return sample - return SampleShortInfoList(s=[x for x in self.sample_spreadsheet.s if x.user == pgroup]) + return SampleShortInfoList( + s=[x for x in self.sample_spreadsheet.s if x.user == pgroup] + ) def get_beamline_default_raster_params(self) -> SimpleScanParameters: - default_exp_time_s = cfg_get("daq.data_collection_settings.default_raster_scan_settings.exp_time_s", 0.01) - default_transmission = cfg_get("daq.data_collection_settings.default_raster_scan_settings.transmission", 1.0) - default_dtz = cfg_get("daq.data_collection_settings.default_raster_scan_settings.dtz", 250) + default_exp_time_s = cfg_get( + "daq.data_collection_settings.default_raster_scan_settings.exp_time_s", 0.01 + ) + default_transmission = cfg_get( + "daq.data_collection_settings.default_raster_scan_settings.transmission", + 1.0, + ) + default_dtz = cfg_get( + "daq.data_collection_settings.default_raster_scan_settings.dtz", 250 + ) return SimpleScanParameters( dtz=default_dtz, exp_time_s=default_exp_time_s, - transmission=default_transmission + transmission=default_transmission, ) def get_beamline_default_rotation_params(self) -> SimpleScanParameters: - default_exp_time_s = cfg_get("daq.data_collection_settings.default_rotation_settings.exp_time_s", 0.01) - default_transmission = cfg_get("daq.data_collection_settings.default_rotation_settings.transmission", 1.0) - default_dtz = cfg_get("daq.data_collection_settings.default_rotation_settings.dtz", 250) - default_start_omega_deg = cfg_get("daq.data_collection_settings.default_rotation_settings.start_omega_deg", 0.0) - default_increment_omega_deg = cfg_get("daq.data_collection_settings.default_rotation_settings.incr_omega_deg", 0.2) - default_steps = cfg_get("daq.data_collection_settings.default_rotation_settings.steps", 1800) + default_exp_time_s = cfg_get( + "daq.data_collection_settings.default_rotation_settings.exp_time_s", 0.01 + ) + default_transmission = cfg_get( + "daq.data_collection_settings.default_rotation_settings.transmission", 1.0 + ) + default_dtz = cfg_get( + "daq.data_collection_settings.default_rotation_settings.dtz", 250 + ) + default_start_omega_deg = cfg_get( + "daq.data_collection_settings.default_rotation_settings.start_omega_deg", + 0.0, + ) + default_increment_omega_deg = cfg_get( + "daq.data_collection_settings.default_rotation_settings.incr_omega_deg", 0.2 + ) + default_steps = cfg_get( + "daq.data_collection_settings.default_rotation_settings.steps", 1800 + ) return SimpleScanParameters( dtz=default_dtz, exp_time_s=default_exp_time_s, transmission=default_transmission, start_omega_deg=default_start_omega_deg, incr_omega_deg=default_increment_omega_deg, - steps=default_steps + steps=default_steps, ) def get_auto_raster_params(self) -> SimpleScanParameters: @@ -2593,7 +2768,11 @@ class AareDAQ: if self.status.sample is None: return default_params - aaredb_params = self.status.sample.aaredb_params if hasattr(self.status.sample, "aaredb_params") else None + aaredb_params = ( + self.status.sample.aaredb_params + if hasattr(self.status.sample, "aaredb_params") + else None + ) if aaredb_params is None: return default_params @@ -2601,14 +2780,14 @@ class AareDAQ: params = default_params # Exposure - exp = getattr(aaredb_params, 'exposure', None) + exp = getattr(aaredb_params, "exposure", None) if exp is not None: params.exp_time_s = exp else: params.exp_time_s = 0.04 # Default # Transmission - trans = getattr(aaredb_params, 'transmission', None) + trans = getattr(aaredb_params, "transmission", None) if trans is not None: logger.debug(f"transmission: {trans}") params.transmission = trans / 100.0 if trans > 1.0 else trans @@ -2616,7 +2795,7 @@ class AareDAQ: params.transmission = 1.0 # Default # Resolution and DTZ - res = getattr(aaredb_params, 'targetresolution', None) + res = getattr(aaredb_params, "targetresolution", None) if res is not None: logger.debug(f"resolution: {res}") try: @@ -2635,17 +2814,21 @@ class AareDAQ: return params - def get_collection_params(self, prefer_smart: bool = False) -> tuple[SimpleScanParameters, str]: + def get_collection_params( + self, prefer_smart: bool = False + ) -> tuple[SimpleScanParameters, str]: spreadsheet_params, file_prefix = self.spreadsheet_params() logger.debug(f"spreadsheet_params: {spreadsheet_params}") smart_params = self.__cfg.auto_params - default_params = SimpleScanParameters(exp_time_s=0.04, dtz=110, incr_omega_deg=0.2) - #if file_prefix is not None: + default_params = SimpleScanParameters( + exp_time_s=0.04, dtz=110, incr_omega_deg=0.2 + ) + # if file_prefix is not None: # default_params.file_prefix = file_prefix - #self.__aare.send_msg_to_db(self.sample,event_type=SampleEventType(''), comment=f'smart_params: {smart_params}') + # self.__aare.send_msg_to_db(self.sample,event_type=SampleEventType(''), comment=f'smart_params: {smart_params}') if prefer_smart: if smart_params: - #if file_prefix: + # if file_prefix: # smart_params.file_prefix = file_prefix return smart_params, "smart_params" if spreadsheet_params: @@ -2654,13 +2837,18 @@ class AareDAQ: if spreadsheet_params: return spreadsheet_params, "spreadsheet_params" if smart_params: - #if file_prefix: + # if file_prefix: # smart_params.file_prefix = file_prefix return smart_params, "smart_params" return default_params, "defaults" - def _end_operation(self, start, operation:Optional[DAQOperation]=DAQOperation.AUTOMATION, error: bool = False) -> float: + def _end_operation( + self, + start, + operation: Optional[DAQOperation] = DAQOperation.AUTOMATION, + error: bool = False, + ) -> float: """ End an operation and return elapsed time. @@ -2717,19 +2905,18 @@ class AareDAQ: ) self._emit_automation_progress(progress) - formatted_date = datetime.now().strftime('%Y%m%d') + formatted_date = datetime.now().strftime("%Y%m%d") sample_prefix = "{}/{}/{:02d}/{}".format( - formatted_date, - sample.puck_name, - sample.pin, - sample.sample_name + formatted_date, sample.puck_name, sample.pin, sample.sample_name ) try: self._validate_automation_state(context="automation start") self.__cfg.try_set_busy(timeout=self.AUTOMATION_BUSY_TIMEOUT_S) - self._validate_automation_state(context="after acquiring automation busy state") + self._validate_automation_state( + context="after acquiring automation busy state" + ) logger.info("Cancelling any pending jfjoch operations") self.__jfjoch.cancel() self._set_progress_context( @@ -2737,22 +2924,34 @@ class AareDAQ: current_sample_name=getattr(sample, "sample_name", "") or "", ) - self._mark_progress_running(progress, WorkflowStateKind.MOUNT, "Mounting sample") + self._mark_progress_running( + progress, WorkflowStateKind.MOUNT, "Mounting sample" + ) if not self._execute_mount_and_prepare(sample): mount_error_message = self._last_mount_error_message or "Mount failed" - self._mark_progress_failed(progress, WorkflowStateKind.MOUNT, mount_error_message) + self._mark_progress_failed( + progress, WorkflowStateKind.MOUNT, mount_error_message + ) return self._end_operation(start, DAQOperation.MOUNT, error=True) self._validate_automation_state(context="after mount") - self._mark_progress_success(progress, WorkflowStateKind.MOUNT, "Mount complete") + self._mark_progress_success( + progress, WorkflowStateKind.MOUNT, "Mount complete" + ) self.__set_state(BeamlineStateEnum.SampleAlignment) - self._validate_automation_state(context="after transition to SampleAlignment") + self._validate_automation_state( + context="after transition to SampleAlignment" + ) logger.info(f"mounting done at {time.perf_counter() - start}") - self._mark_progress_running(progress, WorkflowStateKind.LOOP_CENTRE, "Centering sample") + self._mark_progress_running( + progress, WorkflowStateKind.LOOP_CENTRE, "Centering sample" + ) local_contact_config = self.get_local_contact_config() - mount_to_center_sleep_s = float(local_contact_config.mount_to_center_sleep_s) + mount_to_center_sleep_s = float( + local_contact_config.mount_to_center_sleep_s + ) if mount_to_center_sleep_s > 0: logger.info( @@ -2809,9 +3008,13 @@ class AareDAQ: self._validate_automation_state(context="after face detection") logger.info(f"Face Detection done at {time.perf_counter() - start}") - self._mark_progress_success(progress, WorkflowStateKind.LOOP_CENTRE, "Centering complete") + self._mark_progress_success( + progress, WorkflowStateKind.LOOP_CENTRE, "Centering complete" + ) - self._mark_progress_running(progress, WorkflowStateKind.RASTER, "Running raster") + self._mark_progress_running( + progress, WorkflowStateKind.RASTER, "Running raster" + ) hex_string = secrets.token_hex(3) raster_params = self.get_auto_raster_params() geom = self.sample_geometry @@ -2823,13 +3026,17 @@ class AareDAQ: n_x=1, n_y=1, dtz=raster_params.dtz, - grid_size_mm=Coordinate(x=geom.beam_size_mm.x * 0.5, y=geom.beam_size_mm.y * 0.5), + grid_size_mm=Coordinate( + x=geom.beam_size_mm.x * 0.5, y=geom.beam_size_mm.y * 0.5 + ), omega_deg=self.omega, transmission=raster_params.transmission, ) try: - raster_result = self._execute_raster_sequence(raster_grid, auto_center=True) + raster_result = self._execute_raster_sequence( + raster_grid, auto_center=True + ) except AutoRasterSampleSkipped as e: logger.warning( "Skipping sample during automation because auto-raster grid is too large", @@ -2851,7 +3058,9 @@ class AareDAQ: StepStatus.SKIPPED, "Data collection skipped because auto-raster was too large", ) - self._mark_progress_finished(progress, True, "Sample skipped: auto-raster too large") + self._mark_progress_finished( + progress, True, "Sample skipped: auto-raster too large" + ) return self._end_operation(start, DAQOperation.AUTOMATION, error=False) if raster_result is None: @@ -2862,17 +3071,25 @@ class AareDAQ: raster_request_log_context(raster_grid), ), ) - self._mark_progress_failed(progress, WorkflowStateKind.RASTER, "Raster failed") + self._mark_progress_failed( + progress, WorkflowStateKind.RASTER, "Raster failed" + ) return self._end_operation(start, DAQOperation.RASTER, error=True) self._validate_automation_state(context="after raster") - self._mark_progress_success(progress, WorkflowStateKind.RASTER, "Raster complete") + self._mark_progress_success( + progress, WorkflowStateKind.RASTER, "Raster complete" + ) self.__set_state(BeamlineStateEnum.DataCollection) - self._validate_automation_state(context="after transition to DataCollection") + self._validate_automation_state( + context="after transition to DataCollection" + ) logger.info(f"Raster scans completed at {time.perf_counter() - start}") - self._mark_progress_running(progress, WorkflowStateKind.DATA_COLLECTION, "Collecting data") + self._mark_progress_running( + progress, WorkflowStateKind.DATA_COLLECTION, "Collecting data" + ) params, source = self.get_collection_params(prefer_smart=False) logger.info(f"Using {source} for data collection: {params}") @@ -2892,23 +3109,35 @@ class AareDAQ: rotation_result = self._execute_rotation_sequence(rotation_request) if rotation_result is None: logger.error("Rotation result was None") - self._mark_progress_failed(progress, WorkflowStateKind.DATA_COLLECTION, "Collection failed") + self._mark_progress_failed( + progress, WorkflowStateKind.DATA_COLLECTION, "Collection failed" + ) return self._end_operation(start, DAQOperation.ROTATION, error=True) self._validate_automation_state(context="after data collection") - self._mark_progress_success(progress, WorkflowStateKind.DATA_COLLECTION, "Collection complete") + self._mark_progress_success( + progress, WorkflowStateKind.DATA_COLLECTION, "Collection complete" + ) logger.info(f"Rotation scan done at {time.perf_counter() - start}") self._validate_automation_state(context="automation end") except BECCommunicationError as e: - self._raise_if_critical_bec_error(e, command=getattr(e, "operation", None) or "bec") + self._raise_if_critical_bec_error( + e, command=getattr(e, "operation", None) or "bec" + ) except JFJochCommunicationError as e: - self._raise_if_critical_jfjoch_detector_error(e, command=e.endpoint or "unknown") + self._raise_if_critical_jfjoch_detector_error( + e, command=e.endpoint or "unknown" + ) raise - except (TransformationInvalidException, StateTransitionFailed, MaintenanceStateException) as e: + except ( + TransformationInvalidException, + StateTransitionFailed, + MaintenanceStateException, + ) as e: logger.error(f"Critical automation state error: {e}") if progress.current_step is not None: current_kind = next( @@ -2932,13 +3161,14 @@ class AareDAQ: raise except (BeamlineBusyTimeoutException, BeamlineBusyException) as e: - time_of_measure = abs(time.perf_counter() - start) if time_of_measure > self.AUTOMATION_BUSY_TIMEOUT_S: - logger.error(f"Error in measure due to Beamline Busy State timeout:" - f"Time of: {time_of_measure} is greater than timeout duration {self.AUTOMATION_BUSY_TIMEOUT_S}" - f"Error thrown: {e}") + logger.error( + f"Error in measure due to Beamline Busy State timeout:" + f"Time of: {time_of_measure} is greater than timeout duration {self.AUTOMATION_BUSY_TIMEOUT_S}" + f"Error thrown: {e}" + ) else: logger.error(f"Error in measure due to Beamline Busy State: {e}") @@ -2960,7 +3190,8 @@ class AareDAQ: operation=DAQOperation.AUTOMATION, sample=sample, error=e, - event_type=SampleEventType.FAILED) + event_type=SampleEventType.FAILED, + ) raise @@ -2983,7 +3214,8 @@ class AareDAQ: operation=DAQOperation.AUTOMATION, sample=sample, error=e, - event_type=SampleEventType.FAILED) + event_type=SampleEventType.FAILED, + ) raise Exception(f"Critical Error in automation: {e}") from e @@ -3002,7 +3234,7 @@ class AareDAQ: 4. If exception is raised during transformation, state is set to maintenance 5. If transformation goes OK, target state is set - Busy state will be cleared only, if exception is raised. """ + Busy state will be cleared only, if exception is raised.""" curr_state = self.__cfg.state @@ -3017,15 +3249,17 @@ class AareDAQ: if target == BeamlineStateEnum.Maintenance: self.__cfg.state = BeamlineStateEnum.Maintenance - logger.warning("Beamline entered Maintenance state during state transition request.") + logger.warning( + "Beamline entered Maintenance state during state transition request." + ) return elif target == curr_state: logger.debug( f"State already set to {target}", extra={"from_state": curr_state, "to_state": target}, ) - #TODO unify BeamlineStateEnum and BeamlineState and allow transition from same state to same state to recover motor positions! - #TODO or just allow BEC thingy + # TODO unify BeamlineStateEnum and BeamlineState and allow transition from same state to same state to recover motor positions! + # TODO or just allow BEC thingy return elif target != curr_state: self.__cfg.state = BeamlineStateEnum.Moving @@ -3189,7 +3423,9 @@ class AareDAQ: # the preset for the current zoom on entry. if target == BeamlineStateEnum.BeamLocation: self.__cfg.zoom_mode = ZoomModeEnum.BeamLocation - self.__devs.samcam_settings = self.__cfg.zoom_settings.get_camera_settings(self.__devs.zoom) + self.__devs.samcam_settings = ( + self.__cfg.zoom_settings.get_camera_settings(self.__devs.zoom) + ) elif self.__cfg.zoom_mode == ZoomModeEnum.BeamLocation: self.__cfg.zoom_mode = ZoomModeEnum.User logger.info( @@ -3231,11 +3467,15 @@ class AareDAQ: width = int(metadata.get("detector_width", 1)) height = int(metadata.get("detector_height", 1)) pixel_size_mm = float(metadata.get("pixel_size_mm", 0.15)) - detector_description = str(metadata.get("detector_description", "unavailable")) - detector_serial_number = str(metadata.get("detector_serial_number", "unavailable")) - energy=self.__devs.energy_kev - dtz=self.__devs.dtz - beam_center=self.__cfg.beam_center + detector_description = str( + metadata.get("detector_description", "unavailable") + ) + detector_serial_number = str( + metadata.get("detector_serial_number", "unavailable") + ) + energy = self.__devs.energy_kev + dtz = self.__devs.dtz + beam_center = self.__cfg.beam_center return DiffractionGeometry( energy_keV=energy, dtz_mm=dtz, @@ -3251,9 +3491,9 @@ class AareDAQ: logger.warning( f"Falling back to default diffraction geometry because cached detector metadata is unavailable: {e}" ) - energy=self.__devs.energy_kev - dtz=self.__devs.dtz - beam_center=self.__cfg.beam_center + energy = self.__devs.energy_kev + dtz = self.__devs.dtz + beam_center = self.__cfg.beam_center return DiffractionGeometry( energy_keV=energy, dtz_mm=dtz, @@ -3288,17 +3528,20 @@ class AareDAQ: dtz_max = self.__cfg.cached_dtz_high if ( - dtz_min is None or dtz_max is None - or (dtz_min > dtz_max) - or (dtz_min == dtz_max) - or (dtz_min == 0 and dtz_max == 0) + dtz_min is None + or dtz_max is None + or (dtz_min > dtz_max) + or (dtz_min == dtz_max) + or (dtz_min == 0 and dtz_max == 0) ): - logger.warning("DTZ limits missing from cache, using conservative defaults in beamline_status") - dtz_min = cfg_get('daq.hardware.default_detector_distance_minimum', 100) - dtz_max = cfg_get('daq.hardware.default_detector_distance_maximum', 1000) logger.warning( - f"using dtz_min {dtz_min} and dtz_max {dtz_max}" + "DTZ limits missing from cache, using conservative defaults in beamline_status" ) + dtz_min = cfg_get("daq.hardware.default_detector_distance_minimum", 100) + dtz_max = cfg_get( + "daq.hardware.default_detector_distance_maximum", 1000 + ) + logger.warning(f"using dtz_min {dtz_min} and dtz_max {dtz_max}") return BeamlineStatus( ring_current_mA=ring_current, @@ -3347,7 +3590,9 @@ class AareDAQ: aerotech_err = f"Cannot connect to Aerotech: {e}" return aerotech_ok, aerotech_err - def _safe_geom(self) -> tuple[SampleGeometryModel, bool, str | None, bool, str | None]: + def _safe_geom( + self, + ) -> tuple[SampleGeometryModel, bool, str | None, bool, str | None]: """ Return (geom, smargon_connected, smargon_error, aerotech_connected, aerotech_error) without raising. Uses a conservative fallback geometry if Smargon access fails. @@ -3355,7 +3600,13 @@ class AareDAQ: aerotech_connected, smargon_connected = True, True aerotech_error, smargon_error = None, None try: - return self.sample_geometry, smargon_connected, smargon_error, aerotech_connected, aerotech_error + return ( + self.sample_geometry, + smargon_connected, + smargon_error, + aerotech_connected, + aerotech_error, + ) except SmargonCommunicationError as e: logger.warning(f"Smargon error in _safe_geom: {e}") smargon_error = f"Cannot connect to Smargon: {e}" @@ -3386,13 +3637,19 @@ class AareDAQ: aerotech=Coordinate(x=0.0, y=0.0, z=0.0), aerotech_meas=Coordinate(x=0.0, y=0.0, z=0.0), ) - return fallback, smargon_connected, smargon_error, aerotech_connected, aerotech_error + return ( + fallback, + smargon_connected, + smargon_error, + aerotech_connected, + aerotech_error, + ) def _safe_beamline_status(self) -> BeamlineStatus: try: return self.beamline_status except Exception as e: - #TODO add error message to send to GUI to say problem + # TODO add error message to send to GUI to say problem logger.error(f"Failed to retrieve beamline status: {str(e)}") raise @@ -3438,29 +3695,31 @@ class AareDAQ: @property def status(self) -> DAQStatusModel: try: - #og_start = time.perf_counter() + # og_start = time.perf_counter() safe_sample, tell_ok, tell_err = self._safe_sample() - #logger.debug(f"Safe sample info call took {time.perf_counter() - og_start:.3f}s") - #start = time.perf_counter() - safe_geom, smargon_ok, smargon_err, aerotech_ok, aerotech_err = self._safe_geom() - #logger.debug(f"safe geom call took {time.perf_counter() - start:.3f}s") - #start = time.perf_counter() + # logger.debug(f"Safe sample info call took {time.perf_counter() - og_start:.3f}s") + # start = time.perf_counter() + safe_geom, smargon_ok, smargon_err, aerotech_ok, aerotech_err = ( + self._safe_geom() + ) + # logger.debug(f"safe geom call took {time.perf_counter() - start:.3f}s") + # start = time.perf_counter() safe_tell_state = self._safe_tell_state() - #logger.debug(f"Safe tell call call took {time.perf_counter() - start:.3f}s") - #start = time.perf_counter() + # logger.debug(f"Safe tell call call took {time.perf_counter() - start:.3f}s") + # start = time.perf_counter() safe_beamline_status = self._safe_beamline_status() - #logger.debug(f"Safe beamline status call took {time.perf_counter() - start:.3f}s") - #start = time.perf_counter() + # logger.debug(f"Safe beamline status call took {time.perf_counter() - start:.3f}s") + # start = time.perf_counter() safe_diffraction_geom = self._safe_diffraction_geometry() - #logger.debug(f"Safe diffraction geometry call took {time.perf_counter() - start:.3f}s") - #start = time.perf_counter() + # logger.debug(f"Safe diffraction geometry call took {time.perf_counter() - start:.3f}s") + # start = time.perf_counter() session_status = SessionStatus( - current_pgroup=self.__cfg.pgroup, - session=self.__cfg.session_state(0), # 0 is dummy session - staff=False - ) - #logger.debug(f"Safe session status call took {time.perf_counter() - start:.3f}s") - #start = time.perf_counter() + current_pgroup=self.__cfg.pgroup, + session=self.__cfg.session_state(0), # 0 is dummy session + staff=False, + ) + # logger.debug(f"Safe session status call took {time.perf_counter() - start:.3f}s") + # start = time.perf_counter() status = DAQStatusModel( state=self.state, busy=self.busy, @@ -3481,8 +3740,8 @@ class AareDAQ: aerotech_connected=aerotech_ok, aerotech_error=aerotech_err, ) - #logger.debug(f"Creating DAQStatusModel took {time.perf_counter() - start:.3f}s") - #logger.debug(f"returning status call took {time.perf_counter() - og_start:.3f}s") + # logger.debug(f"Creating DAQStatusModel took {time.perf_counter() - start:.3f}s") + # logger.debug(f"returning status call took {time.perf_counter() - og_start:.3f}s") return status except Exception as e: @@ -3534,11 +3793,15 @@ class AareDAQ: return [] return [str(item) for item in devices] - def bec_reinitialise_planner_and_position_devices(self, method: str = "auto") -> list[str]: + def bec_reinitialise_planner_and_position_devices( + self, method: str = "auto" + ) -> list[str]: self.__cfg.try_set_busy(timeout=360) try: self.__devs.bec_worker.load_user_macros() - return self.__devs.bec_worker.reinitialise_planner_and_position_devices(method=method) + return self.__devs.bec_worker.reinitialise_planner_and_position_devices( + method=method + ) finally: self.__cfg.state_busy = False @@ -3563,7 +3826,9 @@ class AareDAQ: finally: self.__cfg.state_busy = False - def fluorimeter_take_spectrum(self, fm: FluorescenceSpectrumParameterModel) -> FluorescenceSpectrumOutputModel: + def fluorimeter_take_spectrum( + self, fm: FluorescenceSpectrumParameterModel + ) -> FluorescenceSpectrumOutputModel: self.__cfg.try_set_busy(timeout=360) try: @@ -3585,5 +3850,7 @@ class AareDAQ: def get_local_contact_config(self) -> LocalContactConfigModel: return self.__cfg.get_local_contact_config() - def set_local_contact_config(self, config: LocalContactConfigModel) -> LocalContactConfigModel: + def set_local_contact_config( + self, config: LocalContactConfigModel + ) -> LocalContactConfigModel: return self.__cfg.set_local_contact_config(config) diff --git a/src/aare/daq/devices.py b/src/aare/daq/devices.py index 11113198..50e8d1fe 100644 --- a/src/aare/daq/devices.py +++ b/src/aare/daq/devices.py @@ -5,25 +5,24 @@ import time # - property to read device value # - 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 numpy as np +from aarecommon.config.beamline import cfg_get +from aarecommon.config.logger import setup_logger +from aarecommon.config.logger_events import log_timing +from aarecommon.math.coordinate import AerotechCoordinate, SmargonCoordinate +from aarecommon.models.beamline import MXBeamline +from aarecommon.models.models import SampleCameraSettings, StagePositionEnum from epics import PV -from aare.common.beamline import MXBeamline, cfg_get -from aare.common.coordinate import SmargonCoordinate, AerotechCoordinate -from aare.common.logger_config import setup_logger -from aare.common.logger_events import log_timing -from aare.common.models import SampleCameraSettings, StagePositionEnum -from aare.devices import smargon, aerotech -from aare.devices.area_detector import epicsAD, AutoEnum +from aare.devices import aerotech, smargon +from aare.devices.area_detector import AutoEnum, epicsAD +from aare.devices.bec_worker import BECClientWorker, DetectorCoverEnum from aare.devices.enum_pv import EnumPV from aare.devices.experimental_hutch_shutter import ExperimentalHutchShutter from aare.devices.my_motor import MyMotor from aare.devices.pss_state import PssState - -from aare.devices.set_get_pv import SetGetPV, PredefinedPV +from aare.devices.set_get_pv import PredefinedPV, SetGetPV from aare.devices.tell_client import make_tell_client -from aare.devices.bec_worker import BECClientWorker, DetectorCoverEnum from aare.devices.zmq_client import ZMQCameraClient logger = setup_logger("aareDAQ") @@ -43,50 +42,55 @@ class BeamlineDevices: # Personnel Safety System: gates whether the robot is allowed to move. self.pss = PssState(beamline=self._beamline) - #faster to define the dtz object here than in functions and then use + # faster to define the dtz object here than in functions and then use self.__dtz = self.bec_worker.dev.det_z - self.dtz_mod = cfg_get('daq.detector_distance_limit_modifier', 1.0) - #TODO convert epics pvs to BEC + self.dtz_mod = cfg_get("daq.detector_distance_limit_modifier", 1.0) + # TODO convert epics pvs to BEC self.__sample_cam = epicsAD(f"{BEAMLINE}-ES-MS:") - self.__front_light = PredefinedPV(name='front_light', - setpv=f"{BEAMLINE}-ES-FL:SET", - getpv=f"{BEAMLINE}-ES-FL:SET", - predefs={"off":1.49, - 'half':2.0, - 'max':3.0}, - timeout=10.0 - ) - self.__back_light = PredefinedPV(name='back_light', - setpv =f"{BEAMLINE}-ES-BL:SET", - getpv=f"{BEAMLINE}-ES-BL:SET", - predefs={"off": 0, - 'half': 0.98, - 'max': 1.2}, - timeout=10.0 - ) - + self.__front_light = PredefinedPV( + name="front_light", + setpv=f"{BEAMLINE}-ES-FL:SET", + getpv=f"{BEAMLINE}-ES-FL:SET", + predefs={"off": 1.49, "half": 2.0, "max": 3.0}, + timeout=10.0, + ) + self.__back_light = PredefinedPV( + name="back_light", + setpv=f"{BEAMLINE}-ES-BL:SET", + getpv=f"{BEAMLINE}-ES-BL:SET", + predefs={"off": 0, "half": 0.98, "max": 1.2}, + timeout=10.0, + ) # self.__front_light = self.bec_worker.dev.fl_bright # need wrapper on bec_worker layer # self.__back_light = self.bec_worker.dev.bl_bright #need wrapper on bec_worker layer - self.__back_light_pos = EnumPV(name = "back_light_pos", - setpv = f"{BEAMLINE}-ES-BL:POS-SET", - getpv = f"{BEAMLINE}-ES-BL:POS-GET", - timeout = 10.0) + self.__back_light_pos = EnumPV( + name="back_light_pos", + setpv=f"{BEAMLINE}-ES-BL:POS-SET", + getpv=f"{BEAMLINE}-ES-BL:POS-GET", + timeout=10.0, + ) self.__ringcurrent = self.bec_worker.ring_current - self.__zoom = SetGetPV(name = f"zoom", - setpv = f"{BEAMLINE}-ES-MS:ZOOM.VAL", - getpv = f"{BEAMLINE}-ES-MS:ZOOM.RBV") + self.__zoom = SetGetPV( + name=f"zoom", + setpv=f"{BEAMLINE}-ES-MS:ZOOM.VAL", + getpv=f"{BEAMLINE}-ES-MS:ZOOM.RBV", + ) - self.__cryojet_pos = EnumPV(name='cryojet_pos', - setpv = f"{BEAMLINE}-ES-CS:POS-SET", - getpv = f"{BEAMLINE}-ES-CS:POS-GET", - timeout = 10.0) + self.__cryojet_pos = EnumPV( + name="cryojet_pos", + setpv=f"{BEAMLINE}-ES-CS:POS-SET", + getpv=f"{BEAMLINE}-ES-CS:POS-GET", + timeout=10.0, + ) - self.__cryojet_x = MyMotor(f"{BEAMLINE}-ES-CS:TRX") #currently in is 5 out is 15? + self.__cryojet_x = MyMotor( + f"{BEAMLINE}-ES-CS:TRX" + ) # currently in is 5 out is 15? self.__cryojet_temperature_get = PV(f"{BEAMLINE}-ES-CS:TEMP_RBV") self.__cryojet_temperature_set = PV(f"{BEAMLINE}-ES-CS:TEMP.VAL") @@ -95,20 +99,20 @@ class BeamlineDevices: self.__transmission = SetGetPV( name="transmission", setpv=f"{BEAMLINE}-ES-BCFI:TRANSM-SET", - getpv=f"{BEAMLINE}-ES-BCFI:TRANSM-GET" + getpv=f"{BEAMLINE}-ES-BCFI:TRANSM-GET", ) else: self.__transmission = SetGetPV( - name = "transmission", - setpv = f"{BEAMLINE}-ES-SSFI:TRANSM-SET", - getpv = f"{BEAMLINE}-ES-SSFI:TRANSM-GET" + name="transmission", + setpv=f"{BEAMLINE}-ES-SSFI:TRANSM-SET", + getpv=f"{BEAMLINE}-ES-SSFI:TRANSM-GET", ) self.__fast_shutter = PV(f"{BEAMLINE}-ES-SHUTTER:SET") self.magnet_position_sensor = PV(f"{BEAMLINE}-ES-DFS:CBOX-CMP1") self.magnet_position_sensor_readout = PV(f"{BEAMLINE}-ES-DFS:CBOX-USER1") - #self.magnet_position_sensor_readout = PV(f"{BEAMLINE}-ES-DFS:CBOX-REFVAL1") + # self.magnet_position_sensor_readout = PV(f"{BEAMLINE}-ES-DFS:CBOX-REFVAL1") self.magnet_position_sensor_state = PV(f"{BEAMLINE}-ES-DFS:CBOX-STATE") def restart_bec_worker(self, simulated: bool = False) -> None: @@ -117,7 +121,9 @@ class BeamlineDevices: try: self.bec_worker.shutdown_client() except Exception as e: - logger.warning(f"Failed to shutdown previous BEC worker cleanly: {e}") + logger.warning( + f"Failed to shutdown previous BEC worker cleanly: {e}" + ) finally: beamline = MXBeamline.SIMULATED if simulated else self._beamline logger.info(f"Restarting BEC worker with simulated={simulated}") @@ -138,7 +144,7 @@ class BeamlineDevices: logger.info(f"Restarting Smargon controller with simulated={simulated}") self.__smargon = smargon.Smargon(beamline) -# Transmission + # Transmission @property def transmission(self) -> float: return self.__transmission.value @@ -151,7 +157,7 @@ class BeamlineDevices: logger.warning("Setting Transmission is untested") self.__transmission.move(value, wait=wait) -# Lamp light + # Lamp light @property def lamp_light(self) -> float: return self.__front_light.value @@ -175,7 +181,7 @@ class BeamlineDevices: def set_back_light(self, v: float, /, wait: bool = True): self.__back_light.move(v, wait=wait) -# Zoom + # Zoom @property def zoom(self) -> float: return self.__zoom.value @@ -187,7 +193,7 @@ class BeamlineDevices: def set_zoom(self, value: float, /, wait: bool = True): self.__zoom.move(value, wait=wait) -# Optics + # Optics @property def energy_kev(self) -> float: return self.bec_worker.check_current_energy() @@ -198,18 +204,18 @@ class BeamlineDevices: @property def flux(self) -> float: - #TODO FLUX + # TODO FLUX return self.transmission * self.full_flux @property def full_flux(self) -> float: - #TODO wire real flux - i0 needed - max_flux = cfg_get('daq.maximum_flux') + # TODO wire real flux - i0 needed + max_flux = cfg_get("daq.maximum_flux") if max_flux is None: return 0.0 return max_flux -# Cryojet + # Cryojet @property def cryojet_temp(self) -> float: return self.__cryojet_temperature_get.get() @@ -229,7 +235,7 @@ class BeamlineDevices: def cryojet_pos(self, value: StagePositionEnum): self.cryojet_pos_setter(value, wait=True) - def cryojet_pos_setter(self, value: StagePositionEnum, wait:bool=False): + def cryojet_pos_setter(self, value: StagePositionEnum, wait: bool = False): self.__cryojet_pos.move(value, wait=wait) # Shutter @@ -241,12 +247,12 @@ class BeamlineDevices: def shutter(self, opened: bool): self.__fast_shutter.put(opened) -# Sample camera + # Sample camera @property def samcam_settings(self) -> SampleCameraSettings: return SampleCameraSettings( gain=self.__sample_cam.gain_rbv.value, - exposure=self.__sample_cam.expo_rbv.value + exposure=self.__sample_cam.expo_rbv.value, ) @samcam_settings.setter @@ -263,10 +269,10 @@ class BeamlineDevices: """ return int(self.__sample_cam.uid.get()) -# Detector Z + # Detector Z @property def dtz(self) -> float: - return self.__dtz.read()['det_z']['value'] + return self.__dtz.read()["det_z"]["value"] @dtz.setter def dtz(self, value: float): @@ -274,11 +280,15 @@ class BeamlineDevices: def set_dtz(self, value: float, /, wait: bool = True): if value < self.__dtz.low_limit: - #raise ValueError(f"Detector distance cannot be less than {self.detector_distance_minimum} mm") - logger.warning(f"Requested detector distance {value} is less than minimum: {self.__dtz.low_limit}, setting to minimum") + # raise ValueError(f"Detector distance cannot be less than {self.detector_distance_minimum} mm") + logger.warning( + f"Requested detector distance {value} is less than minimum: {self.__dtz.low_limit}, setting to minimum" + ) value = self.__dtz.low_limit + self.dtz_mod if value > self.__dtz.high_limit: - logger.warning(f"Requested detector distance {value} is greater than maximum: {self.__dtz.high_limit}, setting to maximum") + logger.warning( + f"Requested detector distance {value} is greater than maximum: {self.__dtz.high_limit}, setting to maximum" + ) value = self.__dtz.high_limit - self.dtz_mod if wait: status = self.bec_worker.det_z(value, timeout=60) @@ -299,7 +309,7 @@ class BeamlineDevices: @property def aerotech_pos(self) -> AerotechCoordinate: - #TODO add u/omega???? + # TODO add u/omega???? return self.aerotech.get_position() @aerotech_pos.setter @@ -325,24 +335,24 @@ class BeamlineDevices: self.set_aerotech_omega(val, wait=True) 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, - ) + self, + val: float, + /, + wait: bool = True, + incremental: bool = False, + ): + target = AerotechCoordinate(omega_deg=val) + return self.aerotech.position( + target, + wait=wait, + incremental=incremental, + ) def aerotech_stop(self): - #TODO link cancel in GUI to cancel in aerotech if not already done + # TODO link cancel in GUI to cancel in aerotech if not already done pass -# Smargon goniometer + # Smargon goniometer @property def smargon_pos(self) -> SmargonCoordinate: return self.__smargon.readback @@ -366,11 +376,15 @@ class BeamlineDevices: def smargon_initialize(self): self.__smargon.initialize() + if __name__ == "__main__": - from aare.common.beamline import mx_beamline + from aarecommon.config.beamline import mx_beamline + beamline = mx_beamline() devs = BeamlineDevices(beamline) - devs.bec_worker.scilog_msg(message="Testing DAQ hijack of BEC sci log communication channel", color='blue') + devs.bec_worker.scilog_msg( + message="Testing DAQ hijack of BEC sci log communication channel", color="blue" + ) # devs.aerotech_pos = Coordinate(x=124.0, y=1.0, z=1.0) # print(devs.aerotech_pos) - # devs.aerotech_omega = 0.0 \ No newline at end of file + # devs.aerotech_omega = 0.0 diff --git a/src/aare/daq/mlbox.py b/src/aare/daq/mlbox.py index 3b7ac5ba..d1d59dd4 100644 --- a/src/aare/daq/mlbox.py +++ b/src/aare/daq/mlbox.py @@ -1,18 +1,21 @@ +import time from dataclasses import dataclass, field -from typing import Optional, Iterable +from typing import Iterable, Optional import cv2 import numpy as np -import time - +from aarecommon.config.logger import setup_logger +from aarecommon.config.logger_events import log_timing +from aarecommon.models.beamline import MXBeamline +from aarecommon.models.models import ( + BoundingBoxModel, + MLBoxModel, + MLBoxType, + MLOutputModel, +) from aarelcinfer_client.models import LatestPredictionModel -from aare.common.beamline import MXBeamline -from aare.common.models import MLBoxModel, MLOutputModel, MLBoxType, BoundingBoxModel -from aare.common.logger_config import setup_logger -from aare.common.logger_events import log_timing - -logger=setup_logger("aareDAQ") +logger = setup_logger("aareDAQ") @dataclass @@ -64,12 +67,14 @@ class MlBox: raise NotImplementedError(f"MLBox bundle mode not implemented for {bl}") elif bl == MXBeamline.X06DA: - from aare.common.aarelc_infer import AareLCInferWrapper + from aarecommon.lc_infer_wrapper import AareLCInferWrapper + self.__wrapper = AareLCInferWrapper(bl) elif bl == MXBeamline.X10SA: # Lazy import keeps module importable without external client installed - from aare.common.aarelc_infer import AareLCInferWrapper + from aarecommon.lc_infer_wrapper import AareLCInferWrapper + self.__wrapper = AareLCInferWrapper(bl) elif bl == MXBeamline.X06SA: @@ -86,7 +91,9 @@ class MlBox: encoded = np.frombuffer(jpeg_bytes, dtype=np.uint8) image = cv2.imdecode(encoded, cv2.IMREAD_COLOR) if image is None: - logger.warning("Failed to decode prediction bundle image: cv2.imdecode returned None") + logger.warning( + "Failed to decode prediction bundle image: cv2.imdecode returned None" + ) return None return image except Exception as e: @@ -111,7 +118,9 @@ class MlBox: def _prediction_score(predictions: MLOutputModel | None) -> tuple[int, float]: if predictions is None or not predictions.boxes: return (0, 0.0) - confs = [float(m.conf or 0.0) for m in predictions.boxes.values() if m is not None] + confs = [ + float(m.conf or 0.0) for m in predictions.boxes.values() if m is not None + ] return (len(confs), max(confs) if confs else 0.0) def _fetch_prediction_bundle(self): @@ -119,17 +128,23 @@ class MlBox: for attempt in range(1, self.RETRY_COUNT + 1): try: bundle = self.__wrapper.get_latest_prediction_bundle() - logger.debug(f"Fetched prediction bundle on attempt {attempt}/{self.RETRY_COUNT}") + logger.debug( + f"Fetched prediction bundle on attempt {attempt}/{self.RETRY_COUNT}" + ) return bundle except Exception as e: last_error = e - logger.warning(f"Prediction bundle fetch failed on attempt {attempt}/{self.RETRY_COUNT}: {e}") + logger.warning( + f"Prediction bundle fetch failed on attempt {attempt}/{self.RETRY_COUNT}: {e}" + ) if attempt < self.RETRY_COUNT: time.sleep(self.RETRY_SLEEP_S) raise last_error @staticmethod - def _extract_target_point(prediction: LatestPredictionModel | None) -> tuple[float, float] | None: + def _extract_target_point( + prediction: LatestPredictionModel | None, + ) -> tuple[float, float] | None: if prediction is None: return None @@ -150,7 +165,9 @@ class MlBox: if isinstance(raw, (list, tuple)) and len(raw) >= 2: return float(raw[0]), float(raw[1]) except Exception as e: - logger.warning(f"Failed to parse target_point from prediction metadata: {e}") + logger.warning( + f"Failed to parse target_point from prediction metadata: {e}" + ) return None @@ -165,7 +182,9 @@ class MlBox: if raw_focus is not None: focus = float(raw_focus) except Exception as e: - logger.warning(f"Failed to parse focus_score from prediction metadata: {e}") + logger.warning( + f"Failed to parse focus_score from prediction metadata: {e}" + ) return MLBundleMeta(target_point=target_point, focus=focus) @@ -186,7 +205,11 @@ class MlBox: for attempt in range(1, attempts + 1): bundle = self._fetch_prediction_bundle() ml_bundle = self._extract_bundle_candidate(bundle) - predictions, image, bundle_meta = ml_bundle.predictions, ml_bundle.image, ml_bundle.bundle_meta + predictions, image, bundle_meta = ( + ml_bundle.predictions, + ml_bundle.image, + ml_bundle.bundle_meta, + ) if best_image is None and image is not None: best_image = image @@ -219,9 +242,13 @@ class MlBox: if best_predictions is None: logger.info("No valid prediction bundle metadata was available") elif not best_predictions.boxes: - logger.info("Prediction bundle candidates contained no supported detections") + logger.info( + "Prediction bundle candidates contained no supported detections" + ) - return MLBundle(predictions=best_predictions, image=best_image, bundle_meta=best_meta) + return MLBundle( + predictions=best_predictions, image=best_image, bundle_meta=best_meta + ) def _get_latest_bundle_image(self) -> np.ndarray | None: bundle = self._fetch_prediction_bundle() @@ -237,8 +264,18 @@ class MlBox: # Extract (x1,y1,x2,y2) from MLBoxModel/BoundingBoxModel if not box0 or not box0.box or not box1 or not box1.box: raise ValueError("Invalid boxes passed to check_box_relation") - x10, y10, x20, y20 = box0.box.top_x, box0.box.top_y, box0.box.bottom_x, box0.box.bottom_y - x11, y11, x21, y21 = box1.box.top_x, box1.box.top_y, box1.box.bottom_x, box1.box.bottom_y + x10, y10, x20, y20 = ( + box0.box.top_x, + box0.box.top_y, + box0.box.bottom_x, + box0.box.bottom_y, + ) + x11, y11, x21, y21 = ( + box1.box.top_x, + box1.box.top_y, + box1.box.bottom_x, + box1.box.bottom_y, + ) left, top, overlap_x, overlap_y = False, False, False, False # Compute centers for robustness @@ -261,9 +298,14 @@ class MlBox: else: overlap_y = True - return {"left": left, "top": top, "overlap_x": overlap_x, "overlap_y": overlap_y} + return { + "left": left, + "top": top, + "overlap_x": overlap_x, + "overlap_y": overlap_y, + } - def box_relation(self, boxes: MLOutputModel, classes:list[str]|None = None): + def box_relation(self, boxes: MLOutputModel, classes: list[str] | None = None): pin: MLBoxModel | None = boxes.get_best_for_class(MLBoxType.PIN) if not pin or not classes: @@ -272,14 +314,18 @@ class MlBox: box_relative_to_pin: dict[str, bool] = {} for box_type in classes: # Ensure the entry exists and is an MLBoxModel - model = boxes.get(box_type) if hasattr(boxes, "get") else boxes.boxes.get(box_type) + model = ( + boxes.get(box_type) + if hasattr(boxes, "get") + else boxes.boxes.get(box_type) + ) if model: box_relative_to_pin = self.check_box_relation(pin, model) return box_relative_to_pin @staticmethod - def _check_overlap(a: MLBoxModel, b:MLBoxModel) -> float: + def _check_overlap(a: MLBoxModel, b: MLBoxModel) -> float: if not a or not b or not a.box or not b.box: logger.debug(f"Invalid box data for overlap check {a} or {b}") return 0.0 @@ -295,13 +341,19 @@ class MlBox: denominator = a_area + b_area - inter return inter / denominator if denominator > 0 else 0.0 - def box_filter_overlap(self, model:MLBoxModel, pin:MLBoxModel, overlap_parameter:float = 0.5) -> bool: + def box_filter_overlap( + self, model: MLBoxModel, pin: MLBoxModel, overlap_parameter: float = 0.5 + ) -> bool: if self._check_overlap(model, pin) >= overlap_parameter: - return True + return True return False - def _filter_predictions(self, predictions: MLOutputModel, overlap_with_pin: Optional[float] = None, - confidence_min: Optional[float] = None): + def _filter_predictions( + self, + predictions: MLOutputModel, + overlap_with_pin: Optional[float] = None, + confidence_min: Optional[float] = None, + ): pin = predictions.get_best_for_class(MLBoxType.PIN) keys_to_remove = [] @@ -357,7 +409,9 @@ class MlBox: return out if out.boxes else None @staticmethod - def _all_from_prediction_model(prediction: LatestPredictionModel | None) -> Optional[MLOutputModel]: + def _all_from_prediction_model( + prediction: LatestPredictionModel | None, + ) -> Optional[MLOutputModel]: if prediction is None or not getattr(prediction, "boxes", None): return None @@ -380,7 +434,9 @@ class MlBox: conf=conf, ) - logger.debug(f"Parsed {len(out.boxes)} supported detections from bundle metadata") + logger.debug( + f"Parsed {len(out.boxes)} supported detections from bundle metadata" + ) return out def get_best_detections(self, results) -> MLOutputModel | None: @@ -391,9 +447,9 @@ class MlBox: @staticmethod def get_preferred_class_box_with_confidence_threshold( - boxes: MLOutputModel, - preferred_class: Optional[Iterable[int] | int | MLBoxType] = None, - loop_preference_margin: float = 0.1 + boxes: MLOutputModel, + preferred_class: Optional[Iterable[int] | int | MLBoxType] = None, + loop_preference_margin: float = 0.1, ) -> Optional[MLBoxModel]: """ Get best box, preferring loops over pin even if pin has higher confidence, @@ -408,14 +464,21 @@ class MlBox: Best MLBoxModel according to preferences """ if preferred_class is None: - order = (MLBoxType.CRYSTAL, MLBoxType.LOOP_FACE, MLBoxType.LOOP_ALL, MLBoxType.PIN) + order = ( + MLBoxType.CRYSTAL, + MLBoxType.LOOP_FACE, + MLBoxType.LOOP_ALL, + MLBoxType.PIN, + ) else: if isinstance(preferred_class, MLBoxType): order = (preferred_class,) elif isinstance(preferred_class, int): order = (MLBoxType(preferred_class),) else: - order = tuple(MLBoxType(c) if isinstance(c, int) else c for c in preferred_class) + order = tuple( + MLBoxType(c) if isinstance(c, int) else c for c in preferred_class + ) # Get best box from each class best_boxes = {} @@ -451,18 +514,26 @@ class MlBox: return None @staticmethod - def get_preferred_class_box(boxes: MLOutputModel, - preferred_class: Optional[Iterable[int] | int | MLBoxType] = None) -> Optional[ - MLBoxModel]: + def get_preferred_class_box( + boxes: MLOutputModel, + preferred_class: Optional[Iterable[int] | int | MLBoxType] = None, + ) -> Optional[MLBoxModel]: if preferred_class is None: - order = (MLBoxType.CRYSTAL, MLBoxType.LOOP_FACE, MLBoxType.LOOP_ALL, MLBoxType.PIN) + order = ( + MLBoxType.CRYSTAL, + MLBoxType.LOOP_FACE, + MLBoxType.LOOP_ALL, + MLBoxType.PIN, + ) else: if isinstance(preferred_class, MLBoxType): order = (preferred_class,) elif isinstance(preferred_class, int): order = (MLBoxType(preferred_class),) else: - order = tuple(MLBoxType(c) if isinstance(c, int) else c for c in preferred_class) + order = tuple( + MLBoxType(c) if isinstance(c, int) else c for c in preferred_class + ) for cls in order: m = boxes.get_best_for_class(cls) @@ -470,12 +541,25 @@ class MlBox: return m return None - def predict(self, preferred_class = None, - overlap_with_pin: float | None = None, confidence_min: float | None = None, - return_image: bool = False, return_bundle_meta: bool = False - ) -> None | MLBoxModel | tuple[MLBoxModel | None, np.ndarray | None] | MLBoxPredictionResult: + def predict( + self, + preferred_class=None, + overlap_with_pin: float | None = None, + confidence_min: float | None = None, + return_image: bool = False, + return_bundle_meta: bool = False, + ) -> ( + None + | MLBoxModel + | tuple[MLBoxModel | None, np.ndarray | None] + | MLBoxPredictionResult + ): ml_bundle = self._collect_best_bundle() - predictions, image, bundle_meta = ml_bundle.predictions, ml_bundle.image, ml_bundle.bundle_meta + predictions, image, bundle_meta = ( + ml_bundle.predictions, + ml_bundle.image, + ml_bundle.bundle_meta, + ) if not predictions: if return_bundle_meta: return MLBoxPredictionResult( @@ -485,8 +569,11 @@ class MlBox: focus=bundle_meta.focus, ) return (None, image) if return_image else None - self._filter_predictions(predictions=predictions, overlap_with_pin=overlap_with_pin, - confidence_min=confidence_min) + self._filter_predictions( + predictions=predictions, + overlap_with_pin=overlap_with_pin, + confidence_min=confidence_min, + ) box = self.get_preferred_class_box(predictions, preferred_class) if return_bundle_meta: return MLBoxPredictionResult( @@ -498,12 +585,18 @@ class MlBox: return (box, image) if return_image else box def predict_best_no_filter( - self, - return_image: bool = False, - return_bundle_meta: bool = False - ) -> Optional[MLOutputModel] | tuple[Optional[MLOutputModel], np.ndarray | None] | MLBoxPredictionsResult: + self, return_image: bool = False, return_bundle_meta: bool = False + ) -> ( + Optional[MLOutputModel] + | tuple[Optional[MLOutputModel], np.ndarray | None] + | MLBoxPredictionsResult + ): ml_bundle = self._collect_best_bundle() - predictions, image, bundle_meta = ml_bundle.predictions, ml_bundle.image, ml_bundle.bundle_meta + predictions, image, bundle_meta = ( + ml_bundle.predictions, + ml_bundle.image, + ml_bundle.bundle_meta, + ) if return_bundle_meta: return MLBoxPredictionsResult( predictions=predictions, @@ -515,13 +608,23 @@ class MlBox: return predictions, image return predictions - def predict_all_best(self, - overlap_with_pin: float | None = None, - confidence_min: float | None = None, - return_image: bool = False, - return_bundle_meta: bool = False) -> Optional[MLOutputModel] | tuple[Optional[MLOutputModel], np.ndarray | None] | MLBoxPredictionsResult: + def predict_all_best( + self, + overlap_with_pin: float | None = None, + confidence_min: float | None = None, + return_image: bool = False, + return_bundle_meta: bool = False, + ) -> ( + Optional[MLOutputModel] + | tuple[Optional[MLOutputModel], np.ndarray | None] + | MLBoxPredictionsResult + ): ml_bundle = self._collect_best_bundle() - best, image, bundle_meta = ml_bundle.predictions, ml_bundle.image, ml_bundle.bundle_meta + best, image, bundle_meta = ( + ml_bundle.predictions, + ml_bundle.image, + ml_bundle.bundle_meta, + ) if not best: logger.debug(f"No best predictions from ML bundle: {best}") if return_bundle_meta: @@ -533,7 +636,9 @@ class MlBox: ) return (None, image) if return_image else None logger.debug(f"Best predictions from ML bundle: {best}") - self._filter_predictions(best, overlap_with_pin=overlap_with_pin, confidence_min=confidence_min) + self._filter_predictions( + best, overlap_with_pin=overlap_with_pin, confidence_min=confidence_min + ) logger.debug(f"Filtered best predictions from ML bundle: {best}") if return_bundle_meta: return MLBoxPredictionsResult( @@ -544,21 +649,34 @@ class MlBox: ) return (best, image) if return_image else best - def predict_all(self, - overlap_parameter: float | None = None, - confidence_filter: float | None = None, - return_image: bool = False) -> dict[str, list[MLBoxModel]] | tuple[dict[str, list[MLBoxModel]], np.ndarray | None]: + def predict_all( + self, + overlap_parameter: float | None = None, + confidence_filter: float | None = None, + return_image: bool = False, + ) -> ( + dict[str, list[MLBoxModel]] + | tuple[dict[str, list[MLBoxModel]], np.ndarray | None] + ): """ Return a dict keyed by '_' -> [MLBoxModel, ...] for each detection. The ordinal is the running count per class (1-based) in the order they appear after filtering. Example keys: 'Crystal_1', 'Crystal_2', 'Loop_face_1', 'Pin_1', ... """ ml_bundle = self._collect_best_bundle() - grouped, image, _bundle_meta = ml_bundle.predictions, ml_bundle.image, ml_bundle.bundle_meta + grouped, image, _bundle_meta = ( + ml_bundle.predictions, + ml_bundle.image, + ml_bundle.bundle_meta, + ) if not grouped: return ({}, image) if return_image else {} - self._filter_predictions(grouped, overlap_with_pin=overlap_parameter, confidence_min=confidence_filter) + self._filter_predictions( + grouped, + overlap_with_pin=overlap_parameter, + confidence_min=confidence_filter, + ) per_class_counter: dict[MLBoxType, int] = {} out: dict[str, list[MLBoxModel]] = {} @@ -576,12 +694,14 @@ class MlBox: return (out, image) if return_image else out - def predict_best_n_frames(self, - n_frames: int = 5, - preferred_class=None, - overlap_with_pin: float | None = None, - confidence_min: float | None = None, - return_image: bool = False) -> Optional[MLBoxModel] | tuple[Optional[MLBoxModel], np.ndarray | None]: + def predict_best_n_frames( + self, + n_frames: int = 5, + preferred_class=None, + overlap_with_pin: float | None = None, + confidence_min: float | None = None, + return_image: bool = False, + ) -> Optional[MLBoxModel] | tuple[Optional[MLBoxModel], np.ndarray | None]: """ Request bounding boxes for the next N bundle fetches and return the best available box. """ @@ -595,22 +715,26 @@ class MlBox: preferred_class=preferred_class, overlap_with_pin=overlap_with_pin, confidence_min=confidence_min, - return_image=True + return_image=True, ) if box is not None and (box.conf or 0.0) > best_conf: best_box = box best_image = image best_conf = box.conf or 0.0 - logger.info(f"ML detection successful at frame {frame_idx + 1}/{n_frames}: {box}") + logger.info( + f"ML detection successful at frame {frame_idx + 1}/{n_frames}: {box}" + ) if box is None: logger.debug(f"Frame {frame_idx + 1}/{n_frames}: no detection") except Exception as e: - logger.warning(f"Frame {frame_idx + 1}/{n_frames}: prediction error: {e}") + logger.warning( + f"Frame {frame_idx + 1}/{n_frames}: prediction error: {e}" + ) continue if best_box is None: logger.warning(f"No detections in any of {n_frames} frames") - return (best_box, best_image) if return_image else best_box \ No newline at end of file + return (best_box, best_image) if return_image else best_box diff --git a/src/aare/daq/operations/common/ml_bounding_box.py b/src/aare/daq/operations/common/ml_bounding_box.py index 7b42742c..57501291 100644 --- a/src/aare/daq/operations/common/ml_bounding_box.py +++ b/src/aare/daq/operations/common/ml_bounding_box.py @@ -4,21 +4,20 @@ from math import ceil, floor from typing import Callable import cv2 - -from aare.common.beamline import cfg_get -from aare.common.coordinate import Coordinate, SmargonCoordinate -from aare.common.exception_handler import AutoRasterSampleSkipped -from aare.common.logger_events import ( +from aarecommon.config.beamline import cfg_get +from aarecommon.config.logger_events import ( geom_log_context, log_ml_bundle_meta, merge_log_context, sample_log_context, ) -from aare.common.models import SampleShortInfo, MLBoxType -from aare.common.raster_grid import RasterGridRequest -from aare.common.sample_geometry import SampleGeometryModel -from aare.daq.mlbox import MLBoxPredictionResult, MLBoxPredictionsResult, MlBox +from aarecommon.errors.exception_handler import AutoRasterSampleSkipped +from aarecommon.math.coordinate import Coordinate, SmargonCoordinate +from aarecommon.math.sample_geometry import SampleGeometryModel +from aarecommon.models.models import MLBoxType, SampleShortInfo +from aarecommon.models.raster_grid import RasterGridRequest +from aare.daq.mlbox import MlBox, MLBoxPredictionResult, MLBoxPredictionsResult BoxTuple = tuple[float, float, float, float] @@ -27,6 +26,7 @@ BoxTuple = tuple[float, float, float, float] class MLRasterPlan: """Result of an auto-center ML detection: the raster grid request plus the raw loop boxes (at the current zoom) needed to drive zoom-to-fit.""" + grid_request: RasterGridRequest loop_all_box: BoxTuple | None loop_face_box: BoxTuple | None @@ -41,8 +41,12 @@ def _box_tuple(model) -> BoxTuple | None: def _box_extends_beyond(inner: BoxTuple, outer: BoxTuple) -> bool: - return (inner[0] < outer[0] or inner[1] < outer[1] - or inner[2] > outer[2] or inner[3] > outer[3]) + return ( + inner[0] < outer[0] + or inner[1] < outer[1] + or inner[2] > outer[2] + or inner[3] > outer[3] + ) def _box_union(a: BoxTuple, b: BoxTuple) -> BoxTuple: @@ -141,7 +145,9 @@ def get_ml_bounding_box( if filename is not None and bundle_image is not None: annotated_image = bundle_image.copy() - cv2.rectangle(annotated_image, (int(x1), int(y1)), (int(x2), int(y2)), (0, 255, 0), 2) + cv2.rectangle( + annotated_image, (int(x1), int(y1)), (int(x2), int(y2)), (0, 255, 0), 2 + ) upload_image(sample_id, filename, annotated_image) geom = sample_geometry @@ -165,7 +171,10 @@ def get_ml_bounding_box( ) return _box_to_raster_request( - x1=x1, y1=y1, x2=x2, y2=y2, + x1=x1, + y1=y1, + x2=x2, + y2=y2, sample=sample, sample_geometry=geom, logger=logger, @@ -179,7 +188,10 @@ def get_ml_bounding_box( def _box_to_raster_request( *, - x1: float, y1: float, x2: float, y2: float, + x1: float, + y1: float, + x2: float, + y2: float, sample: SampleShortInfo | None, sample_geometry: SampleGeometryModel, logger, @@ -202,7 +214,9 @@ def _box_to_raster_request( # of the n_y scan) can be padded more than the top. frac_x = float(cfg_get("daq.auto_raster.grid_padding_fraction_x", 0.15)) frac_y_top = float(cfg_get("daq.auto_raster.grid_padding_fraction_y", 0.15)) - frac_y_bottom = float(cfg_get("daq.auto_raster.grid_padding_fraction_y_bottom", frac_y_top)) + frac_y_bottom = float( + cfg_get("daq.auto_raster.grid_padding_fraction_y_bottom", frac_y_top) + ) pad_x = max(1, int(ceil(frac_x * n_x))) pad_y_top = max(1, int(ceil(frac_y_top * n_y))) pad_y_bottom = max(1, int(ceil(frac_y_bottom * n_y))) @@ -214,9 +228,15 @@ def _box_to_raster_request( "Padded auto-center raster grid", extra=merge_log_context( sample_log_context(sample), - {"sample_id": sample_id, "ml_image_name": filename, - "pad_cells_x": pad_x, "pad_cells_y_top": pad_y_top, - "pad_cells_y_bottom": pad_y_bottom, "n_x": n_x, "n_y": n_y}, + { + "sample_id": sample_id, + "ml_image_name": filename, + "pad_cells_x": pad_x, + "pad_cells_y_top": pad_y_top, + "pad_cells_y_bottom": pad_y_bottom, + "n_x": n_x, + "n_y": n_y, + }, ), ) @@ -336,43 +356,67 @@ def build_ml_raster_plan( focus=result.focus, ) - loop_all = predictions.get_best_for_class(MLBoxType.LOOP_ALL) if predictions else None - loop_face = predictions.get_best_for_class(MLBoxType.LOOP_FACE) if predictions else None + loop_all = ( + predictions.get_best_for_class(MLBoxType.LOOP_ALL) if predictions else None + ) + loop_face = ( + predictions.get_best_for_class(MLBoxType.LOOP_FACE) if predictions else None + ) # Grid box: prefer loop_face, else loop_all (matches the legacy (3, 0) order). grid_model = loop_face if loop_face is not None else loop_all if grid_model is None: logger.warning( "ML raster plan returned no loop detection", - extra={"sample_id": sample_id, "ml_image_name": filename, - "target_point": result.target_point}, + extra={ + "sample_id": sample_id, + "ml_image_name": filename, + "target_point": result.target_point, + }, ) if filename is not None and bundle_image is not None: upload_image(sample_id, f"{filename}_no_detection", bundle_image) return None - x1, y1, x2, y2 = (grid_model.box.top_x, grid_model.box.top_y, - grid_model.box.bottom_x, grid_model.box.bottom_y) + x1, y1, x2, y2 = ( + grid_model.box.top_x, + grid_model.box.top_y, + grid_model.box.bottom_x, + grid_model.box.bottom_y, + ) # Optionally extend the grid to cover crystals detected outside the loop box. if cfg_get("daq.auto_raster.include_crystal", False) and predictions is not None: for crystal in predictions.get_models_for_class(MLBoxType.CRYSTAL): - cbox = (crystal.box.top_x, crystal.box.top_y, crystal.box.bottom_x, crystal.box.bottom_y) + cbox = ( + crystal.box.top_x, + crystal.box.top_y, + crystal.box.bottom_x, + crystal.box.bottom_y, + ) if _box_extends_beyond(cbox, (x1, y1, x2, y2)): x1, y1, x2, y2 = _box_union((x1, y1, x2, y2), cbox) logger.info( "Extended ML raster grid to include crystal outside the loop box", - extra={"sample_id": sample_id, "ml_image_name": filename, - "crystal_box": cbox}, + extra={ + "sample_id": sample_id, + "ml_image_name": filename, + "crystal_box": cbox, + }, ) if filename is not None and bundle_image is not None: annotated_image = bundle_image.copy() - cv2.rectangle(annotated_image, (int(x1), int(y1)), (int(x2), int(y2)), (0, 255, 0), 2) + cv2.rectangle( + annotated_image, (int(x1), int(y1)), (int(x2), int(y2)), (0, 255, 0), 2 + ) upload_image(sample_id, filename, annotated_image) grid_request = _box_to_raster_request( - x1=x1, y1=y1, x2=x2, y2=y2, + x1=x1, + y1=y1, + x2=x2, + y2=y2, sample=sample, sample_geometry=sample_geometry, logger=logger, @@ -393,4 +437,4 @@ def build_ml_raster_plan( loop_face_box=_box_tuple(loop_face), image_width=image_width, image_height=image_height, - ) \ No newline at end of file + ) diff --git a/src/aare/daq/operations/common/runtime.py b/src/aare/daq/operations/common/runtime.py index 1f911b42..417a3bc6 100644 --- a/src/aare/daq/operations/common/runtime.py +++ b/src/aare/daq/operations/common/runtime.py @@ -1,8 +1,8 @@ from dataclasses import dataclass from typing import Protocol -from aare.common.models import DAQStatusModel, SampleShortInfo -from aare.common.sample_geometry import SampleGeometryModel +from aarecommon.math.sample_geometry import SampleGeometryModel +from aarecommon.models.models import DAQStatusModel, SampleShortInfo class SampleProvider(Protocol): @@ -36,4 +36,4 @@ class DAQRuntimeState: @property def status(self) -> DAQStatusModel: - return self.status_provider.status \ No newline at end of file + return self.status_provider.status diff --git a/src/aare/daq/operations/common/simulate_scan_result.py b/src/aare/daq/operations/common/simulate_scan_result.py index cef36566..062191c5 100644 --- a/src/aare/daq/operations/common/simulate_scan_result.py +++ b/src/aare/daq/operations/common/simulate_scan_result.py @@ -1,10 +1,9 @@ import copy +from aarecommon.models.raster_grid import CompletedRasterGridElem, RasterGridRequest +from aarecommon.models.rotation_scan import CompletedRotationScan, RotationScanRequest from jfjoch_client import ScanResult, ScanResultImagesInner -from aare.common.raster_grid import CompletedRasterGridElem, RasterGridRequest -from aare.common.rotation_scan import CompletedRotationScan, RotationScanRequest - def build_fake_scan_result( *, diff --git a/src/aare/daq/operations/face_detection/service.py b/src/aare/daq/operations/face_detection/service.py index dcc75d5b..5cda1d52 100644 --- a/src/aare/daq/operations/face_detection/service.py +++ b/src/aare/daq/operations/face_detection/service.py @@ -1,8 +1,8 @@ import time -from aare.common.coordinate import Coordinate, SmargonCoordinate -from aare.common.logger_events import log_duration, log_ml_bundle_meta -from aare.common.models import MLBoxModel, ZoomModeEnum +from aarecommon.config.logger_events import log_duration, log_ml_bundle_meta +from aarecommon.math.coordinate import Coordinate, SmargonCoordinate +from aarecommon.models.models import MLBoxModel, ZoomModeEnum import aare.daq.operations.face_detection.utils as fd from aare.daq.mlbox import MLBoxPredictionResult @@ -20,7 +20,9 @@ class FaceDetectionService: def _progress_emitter(self): emitter = self.ctx.services.face_detection_progress if emitter is None: - raise RuntimeError("FaceDetectionService requires services.face_detection_progress") + raise RuntimeError( + "FaceDetectionService requires services.face_detection_progress" + ) return emitter def _log_warning(self, message: str) -> None: @@ -76,15 +78,19 @@ class FaceDetectionService: self.ctx.deps.devs.smargon_wait(60) def run( - self, - *, - steps: int | None = None, - step_size: int | None = None, - face_min_ratio: float | None = None, + self, + *, + steps: int | None = None, + step_size: int | None = None, + face_min_ratio: float | None = None, ) -> FaceDetectionResult: steps = self.ctx.settings.steps if steps is None else steps step_size = self.ctx.settings.step_size if step_size is None else step_size - face_min_ratio = self.ctx.settings.face_min_ratio if face_min_ratio is None else face_min_ratio + face_min_ratio = ( + self.ctx.settings.face_min_ratio + if face_min_ratio is None + else face_min_ratio + ) try: self.ctx.deps.cfg.zoom_mode = ZoomModeEnum.LoopCenter self.ctx.deps.devs.lamp_light = 2.5 @@ -161,7 +167,9 @@ class FaceDetectionService: return FaceDetectionResult(success=True, payload=payload) total_detections = len(boxes_face) + len(boxes_loop) - face_ratio = len(boxes_face) / total_detections if total_detections > 0 else 0.0 + face_ratio = ( + len(boxes_face) / total_detections if total_detections > 0 else 0.0 + ) if boxes_face and face_ratio >= face_min_ratio: boxes = boxes_face @@ -175,10 +183,16 @@ class FaceDetectionService: ) else: boxes = boxes_face - self.logger.debug(f"using loop_face boxes (only source, {len(boxes_face)} entries)") + self.logger.debug( + f"using loop_face boxes (only source, {len(boxes_face)} entries)" + ) - best_fit_angle_area, area_params = fd.get_flat_face(boxes, start_angle, end_angle, True) - best_fit_angle_height, height_params = fd.get_flat_face(boxes, start_angle, end_angle, False) + best_fit_angle_area, area_params = fd.get_flat_face( + boxes, start_angle, end_angle, True + ) + best_fit_angle_height, height_params = fd.get_flat_face( + boxes, start_angle, end_angle, False + ) fit_results = { "Area": {"angle": best_fit_angle_area, "params": area_params}, "Height": {"angle": best_fit_angle_height, "params": height_params}, @@ -229,4 +243,4 @@ class FaceDetectionService: comment="Face detection sequence failed", ) finally: - self.ctx.deps.cfg.zoom_mode = ZoomModeEnum.User \ No newline at end of file + self.ctx.deps.cfg.zoom_mode = ZoomModeEnum.User diff --git a/src/aare/daq/operations/face_detection/utils.py b/src/aare/daq/operations/face_detection/utils.py index a3543ced..b931cf26 100644 --- a/src/aare/daq/operations/face_detection/utils.py +++ b/src/aare/daq/operations/face_detection/utils.py @@ -1,25 +1,30 @@ import json -import time -from typing import Tuple, Dict, List, Optional -import numpy as np -from scipy.optimize import curve_fit, OptimizeWarning -import warnings -import statistics import math -from aare.common.logger_config import setup_logger +import statistics +import time +import warnings +from typing import Dict, List, Optional, Tuple + +import numpy as np +from aarecommon.config.logger import setup_logger +from scipy.optimize import OptimizeWarning, curve_fit logger = setup_logger("aareDAQ") + def box_height_from_tuple(box: tuple[float, float, float, float]) -> float: x1, y1, x2, y2 = box return abs(y2 - y1) + def box_area_from_tuple(box: tuple[float, float, float, float]) -> float: x1, y1, x2, y2 = box - return abs(y2 - y1) * abs(x2 -x1) + return abs(y2 - y1) * abs(x2 - x1) -def prepare_samples(boxes_by_angle: dict[int, tuple[float,float,float,float]], - area = False) -> List[Tuple[float, float]]: + +def prepare_samples( + boxes_by_angle: dict[int, tuple[float, float, float, float]], area=False +) -> List[Tuple[float, float]]: # angles in degrees -> (theta_rad, height) samples = [] for deg, box in boxes_by_angle.items(): @@ -27,10 +32,16 @@ def prepare_samples(boxes_by_angle: dict[int, tuple[float,float,float,float]], samples.append((float(deg), v)) return samples -def cos_model(theta_deg: float | np.ndarray, A: float, B: float, phi_rad: float, C: float): + +def cos_model( + theta_deg: float | np.ndarray, A: float, B: float, phi_rad: float, C: float +): return A + B * np.cos(C * np.deg2rad(theta_deg) - phi_rad) -def mad_filter(samples: List[Tuple[float, float]], k: float = 3.5) -> List[Tuple[float, float]]: + +def mad_filter( + samples: List[Tuple[float, float]], k: float = 3.5 +) -> List[Tuple[float, float]]: if not samples: return samples ys = [y for _, y in samples] @@ -38,22 +49,24 @@ def mad_filter(samples: List[Tuple[float, float]], k: float = 3.5) -> List[Tuple mad = statistics.median([abs(y - med) for y in ys]) or 1.0 return [(d, y) for d, y in samples if abs(y - med) <= k * mad] + def samples_to_json(samples): output_data = { - 'timestamp': time.ctime(), - 'scan_results': samples, - 'total_results': len(samples) + "timestamp": time.ctime(), + "scan_results": samples, + "total_results": len(samples), } - with open("cos_test.json", 'w') as f: + with open("cos_test.json", "w") as f: json.dump(output_data, f, indent=2) return + def fit_metrics(y_true: np.ndarray, y_pred: np.ndarray) -> tuple[float, float, float]: resid = y_true - y_pred rmse = float(np.sqrt(np.mean(resid**2))) mae = float(np.mean(np.abs(resid))) # R² with protection against zero variance - ss_tot = float(np.sum((y_true - np.mean(y_true))**2)) + ss_tot = float(np.sum((y_true - np.mean(y_true)) ** 2)) r2 = float(1.0 - np.sum(resid**2) / ss_tot) if ss_tot > 0 else float("nan") return rmse, mae, r2 @@ -62,8 +75,15 @@ def fit_cosine(samples: List[Tuple[float, float]]) -> dict: samples = mad_filter(samples, k=3.5) if len(samples) < 3: A = sum(y for _, y in samples) / max(1, len(samples)) - return {"A": A, "B": 1.0, "phi_rad": 0.0, "C": 1.0, - "rmse": None, "mae": None, "r2": None} + return { + "A": A, + "B": 1.0, + "phi_rad": 0.0, + "C": 1.0, + "rmse": None, + "mae": None, + "r2": None, + } degs = np.array([d for d, _ in samples], dtype=float) ys = np.array([y for _, y in samples], dtype=float) @@ -89,7 +109,9 @@ def fit_cosine(samples: List[Tuple[float, float]]) -> dict: try: with warnings.catch_warnings(record=True) as w: warnings.simplefilter("always", OptimizeWarning) - popt, _ = curve_fit(cos_model, degs, ys, p0=[A0, B0, phi0, C0], bounds=bounds, maxfev=10000) + popt, _ = curve_fit( + cos_model, degs, ys, p0=[A0, B0, phi0, C0], bounds=bounds, maxfev=10000 + ) if w: logger.debug(f"curve_fit OptimizeWarning: {w[-1].message}") logger.info(f"fit cosine: {popt}") @@ -97,28 +119,47 @@ def fit_cosine(samples: List[Tuple[float, float]]) -> dict: B = max(0.0, B) yhat = cos_model(degs, A, B, phi, C) rmse, mae, r2 = fit_metrics(ys, yhat) - return {"A": A, "B": B, "phi_rad": phi, "C": C, - "rmse": rmse, "mae": mae, "r2": r2} + return { + "A": A, + "B": B, + "phi_rad": phi, + "C": C, + "rmse": rmse, + "mae": mae, + "r2": r2, + } except Exception as e: # Fallback to initial logger.info(f"error in curve fit {e}") yhat0 = cos_model(degs, A0, max(0.0, B0), phi0, C0) rmse0, mae0, r2_0 = fit_metrics(ys, yhat0) - return {"A": float(A0), "B": max(0.0, float(B0)), "phi_rad": float(phi0), "C": float(C0), - "rmse": rmse0, "mae": mae0, "r2": r2_0} + return { + "A": float(A0), + "B": max(0.0, float(B0)), + "phi_rad": float(phi0), + "C": float(C0), + "rmse": rmse0, + "mae": mae0, + "r2": r2_0, + } + def get_samples_out(boxes): samples_out = [] for deg, box in boxes.items(): h = box_height_from_tuple(box) a = box_area_from_tuple(box) - samples_out.append({"angle_deg": float(deg), "height": float(h), "area": float(a)}) + samples_out.append( + {"angle_deg": float(deg), "height": float(h), "area": float(a)} + ) samples_out.sort(key=lambda x: x["angle_deg"]) return samples_out -def choose_best_fit(fits_by_name: Dict[str, Dict]) -> Tuple[Optional[float], Optional[Dict], Optional[str]]: +def choose_best_fit( + fits_by_name: Dict[str, Dict], +) -> Tuple[Optional[float], Optional[Dict], Optional[str]]: def key(entry: Dict): params = entry.get("params") or {} @@ -155,12 +196,15 @@ def choose_best_fit(fits_by_name: Dict[str, Dict]) -> Tuple[Optional[float], Opt best_angle = best_fit.get("angle") best_angle = float(best_angle) if isinstance(best_angle, (int, float)) else None + return (best_angle, best_fit, best_name) - return (best_angle, - best_fit, - best_name) -def get_flat_face(boxes: dict[int, tuple[float,float,float,float]], start_angle:int, end_angle:int, area:bool = False) -> tuple[int, dict]: +def get_flat_face( + boxes: dict[int, tuple[float, float, float, float]], + start_angle: int, + end_angle: int, + area: bool = False, +) -> tuple[int, dict]: samples = prepare_samples(boxes, area=area) parameters = fit_cosine(samples) @@ -176,15 +220,27 @@ def get_flat_face(boxes: dict[int, tuple[float,float,float,float]], start_angle: return 0 return max(search_grid, key=lambda d: cos_model(d, A, B, phi, C)) - best_fit_angle= safe_best_angle(parameters) + best_fit_angle = safe_best_angle(parameters) return best_fit_angle, parameters -def chose_best_angle(boxes: dict[int, tuple[float,float,float,float]], fit_results) -> int: + +def chose_best_angle( + boxes: dict[int, tuple[float, float, float, float]], fit_results +) -> int: measured_angles = list(boxes.keys()) choose_best_fit(fit_results) - candidates = [a for a in (fit_results["Area"]["angle"], fit_results["Height"]["angle"]) if isinstance(a, (int, float))] + candidates = [ + a + for a in (fit_results["Area"]["angle"], fit_results["Height"]["angle"]) + if isinstance(a, (int, float)) + ] if measured_angles and candidates: - chosen = min(candidates, key=lambda a: min(abs(((a - m + 180) % 360) - 180) for m in measured_angles)) + chosen = min( + candidates, + key=lambda a: min( + abs(((a - m + 180) % 360) - 180) for m in measured_angles + ), + ) else: chosen = candidates[0] if candidates else 0 - return chosen \ No newline at end of file + return chosen diff --git a/src/aare/daq/operations/loop_centering/analyzer.py b/src/aare/daq/operations/loop_centering/analyzer.py index 66d6eada..56af4fee 100644 --- a/src/aare/daq/operations/loop_centering/analyzer.py +++ b/src/aare/daq/operations/loop_centering/analyzer.py @@ -1,6 +1,6 @@ -from aare.common.coordinate import Coordinate, SmargonCoordinate -from aare.common.logger_events import log_ml_bundle_meta -from aare.common.models import MLBoxType +from aarecommon.config.logger_events import log_ml_bundle_meta +from aarecommon.math.coordinate import Coordinate, SmargonCoordinate +from aarecommon.models.models import MLBoxType from aare.daq.operations.loop_centering.models import ( AngleAnalysis, @@ -87,10 +87,14 @@ class LoopCenteringAnalyzer: geom = self.ctx.runtime.sample_geometry return SmargonCoordinate( - sh_mm=geom.picture_to_smargon(Coordinate(x=target_point[0], y=target_point[1])) + sh_mm=geom.picture_to_smargon( + Coordinate(x=target_point[0], y=target_point[1]) + ) ) - def _interpret_ml_loop_centre_box(self, boxes) -> tuple[SmargonCoordinate | None, int | None, list[int]]: + def _interpret_ml_loop_centre_box( + self, boxes + ) -> tuple[SmargonCoordinate | None, int | None, list[int]]: if boxes is None: return None, None, [] @@ -188,4 +192,4 @@ class LoopCenteringAnalyzer: calculated_target=analysis.calculated_target, predicted_target=analysis.predicted_target, ) - return analysis \ No newline at end of file + return analysis diff --git a/src/aare/daq/operations/loop_centering/models.py b/src/aare/daq/operations/loop_centering/models.py index dc2d8f51..487c6bfd 100644 --- a/src/aare/daq/operations/loop_centering/models.py +++ b/src/aare/daq/operations/loop_centering/models.py @@ -1,6 +1,7 @@ from dataclasses import dataclass, field -from aare.common.coordinate import SmargonCoordinate +from aarecommon.math.coordinate import SmargonCoordinate + from aare.daq.mlbox import MlBox from aare.daq.operations.common.models import ( BaseOperationContext, @@ -47,4 +48,4 @@ class LoopCenteringSettings: class LoopCenteringContext( BaseOperationContext[LoopCenteringDependencies, LoopCenteringSettings] ): - pass \ No newline at end of file + pass diff --git a/src/aare/daq/operations/loop_centering/service.py b/src/aare/daq/operations/loop_centering/service.py index c4d0f3a7..b19b8b70 100644 --- a/src/aare/daq/operations/loop_centering/service.py +++ b/src/aare/daq/operations/loop_centering/service.py @@ -1,10 +1,9 @@ import time import traceback -from aare.common.exception_handler import LoopCenteringFailed -from aare.common.logger_events import log_duration -from aare.common.models import LoopCenteringResult, MLBoxType -from aare.devices.area_detector import AutoEnum +from aarecommon.config.logger_events import log_duration +from aarecommon.errors.exception_handler import LoopCenteringFailed +from aarecommon.models.models import LoopCenteringResult, MLBoxType from aare.daq.operations.loop_centering.analyzer import LoopCenteringAnalyzer from aare.daq.operations.loop_centering.models import ( @@ -12,6 +11,7 @@ from aare.daq.operations.loop_centering.models import ( AttemptSummary, LoopCenteringContext, ) +from aare.devices.area_detector import AutoEnum class LoopCenteringService: @@ -45,7 +45,11 @@ class LoopCenteringService: sample_id=sample_id, ) - if analysis.has_valid_target and not analysis.ignore_only and analysis.final_target is not None: + if ( + analysis.has_valid_target + and not analysis.ignore_only + and analysis.final_target is not None + ): time_to_move_smargon = time.perf_counter() self.ctx.deps.devs.smargon_pos = analysis.final_target self.ctx.deps.devs.smargon_wait(60) @@ -63,7 +67,9 @@ class LoopCenteringService: event=f"alc_move_zoom_{zoom_value:.0f}_angle_{angle}", ) else: - self.logger.debug(f"Skipping move at angle {angle}; classes={analysis.classes}") + self.logger.debug( + f"Skipping move at angle {angle}; classes={analysis.classes}" + ) if sample_id is not None and analysis.selected_class in ( MLBoxType.LOOP_ALL.value, @@ -106,10 +112,14 @@ class LoopCenteringService: self.logger.info( f"submitting to db loop center sequence for sample {sample_id}, zoom={zoom_value}" ) - self.ctx.services.screenshots.save_to_db(sample_id, "pre_alc", wait_screenshot_sleep_sec) + self.ctx.services.screenshots.save_to_db( + sample_id, "pre_alc", wait_screenshot_sleep_sec + ) for attempt_number in range(1, max_attempts + 1): - self.logger.info(f"Starting ALC attempt {attempt_number}/{max_attempts}") + self.logger.info( + f"Starting ALC attempt {attempt_number}/{max_attempts}" + ) attempt = AttemptSummary(attempt_number=attempt_number) for angle in angles: @@ -123,10 +133,14 @@ class LoopCenteringService: attempt.first_pass.append(result) for c in result.classes: - found_classes_count[int(c)] = found_classes_count.get(int(c), 0) + 1 + found_classes_count[int(c)] = ( + found_classes_count.get(int(c), 0) + 1 + ) if attempt_number == 1: - valid_seen_first_pass = any(result.has_valid_target for result in attempt.first_pass) + valid_seen_first_pass = any( + result.has_valid_target for result in attempt.first_pass + ) if not valid_seen_first_pass: failure_reason = ( "ALC failed on attempt 1: no loop face, loop all or crystal " @@ -145,11 +159,15 @@ class LoopCenteringService: attempt.correction_pass.append(result) for c in result.classes: - found_classes_count[int(c)] = found_classes_count.get(int(c), 0) + 1 + found_classes_count[int(c)] = ( + found_classes_count.get(int(c), 0) + 1 + ) attempts.append(attempt) - valid_seen_correction = any(result.has_valid_target for result in attempt.correction_pass) + valid_seen_correction = any( + result.has_valid_target for result in attempt.correction_pass + ) self.logger.debug( f"ALC attempt {attempt_number}: valid_seen_correction={valid_seen_correction}" @@ -158,7 +176,9 @@ class LoopCenteringService: if valid_seen_correction: if sample_id is not None: self.logger.info(f"sample {sample_id} centered") - self.ctx.services.traces.append_smargon_trace(sample_id=sample_id, event="alc_success") + self.ctx.services.traces.append_smargon_trace( + sample_id=sample_id, event="alc_success" + ) return LoopCenteringResult(success=True) failure_reason = f"ALC exceeded max attempts ({max_attempts})" @@ -203,4 +223,4 @@ class LoopCenteringService: success=False, comment=alc_comment, error=e, - ) \ No newline at end of file + ) diff --git a/src/aare/daq/operations/mounting/models.py b/src/aare/daq/operations/mounting/models.py index 18fc8ad9..252da7d7 100644 --- a/src/aare/daq/operations/mounting/models.py +++ b/src/aare/daq/operations/mounting/models.py @@ -1,7 +1,8 @@ from dataclasses import dataclass -from aare.common.coordinate import AerotechCoordinate -from aare.common.models import SampleShortInfo +from aarecommon.math.coordinate import AerotechCoordinate +from aarecommon.models.models import SampleShortInfo + from aare.daq.operations.common.models import BeamlineDependencies @@ -32,4 +33,4 @@ class MountingResult: @property def is_error(self) -> bool: - return self.error is not None \ No newline at end of file + return self.error is not None diff --git a/src/aare/daq/operations/mounting/service.py b/src/aare/daq/operations/mounting/service.py index 505afffd..614de690 100644 --- a/src/aare/daq/operations/mounting/service.py +++ b/src/aare/daq/operations/mounting/service.py @@ -1,15 +1,15 @@ import time -from aare.common.exception_handler import ( +from aarecommon.errors.exception_handler import ( CriticalTellException, DoorSafetyError, MountingFailed, TellCommunicationError, UnmountingFailed, ) -from aare.devices.tell_client import TellEventValueEnum -from aare.daq.operations.mounting.models import MountingContext, MountingResult +from aare.daq.operations.mounting.models import MountingContext, MountingResult +from aare.devices.tell_client import TellEventValueEnum DRY_AFTER_FAIL_COUNT = 3 STOP_AFTER_FAIL_COUNT = 5 @@ -43,7 +43,9 @@ class MountingService: self.logger.error( "Goniometer didn't reach position based on magnet position sensor readout" ) - raise Exception("Goniometer is not in position based on magnet position sensor readout") + raise Exception( + "Goniometer is not in position based on magnet position sensor readout" + ) def _handle_consecutive_mount_failure(self) -> None: count = self.ctx.deps.cfg.increment_mount_failure_streak() @@ -57,13 +59,17 @@ class MountingService: self.logger.exception(f"Failed to dry after mount failure: {e}") if count >= STOP_AFTER_FAIL_COUNT: - self.logger.error(f"Mount failed {count} times in a row, stopping automation") + self.logger.error( + f"Mount failed {count} times in a row, stopping automation" + ) self.logger.error("Unmounting sample and drying") try: self._unmount_current_sample(timeout=60.0) self.dry(park=True) except Exception as e: - self.logger.exception(f"Failed to clean up after repeated mount failure: {e}") + self.logger.exception( + f"Failed to clean up after repeated mount failure: {e}" + ) raise MountingFailed( f"Mount failed {count} times in a row, stopping automation.", @@ -84,9 +90,13 @@ class MountingService: if isinstance(value, str): self.ctx.deps.devs.tell.check_command_ok() - self.logger.error(f"{self.ctx.deps.devs.tell.get_result(self.ctx.deps.devs.tell._last_cmd_id)}") + self.logger.error( + f"{self.ctx.deps.devs.tell.get_result(self.ctx.deps.devs.tell._last_cmd_id)}" + ) self.logger.error(f"Unexpected string response from Tell mount: {value}") - raise CriticalTellException(f"Critical error in TELL mount: unexpected response '{value}'") + raise CriticalTellException( + f"Critical error in TELL mount: unexpected response '{value}'" + ) if value is None: self.logger.error("Tell mount returned no response (None)") @@ -242,4 +252,4 @@ class MountingService: did_unmount_previous=False, error=e, comment=str(e), - ) \ No newline at end of file + ) diff --git a/src/aare/daq/operations/raster/models.py b/src/aare/daq/operations/raster/models.py index 389c26e2..e6fab589 100644 --- a/src/aare/daq/operations/raster/models.py +++ b/src/aare/daq/operations/raster/models.py @@ -1,7 +1,8 @@ from dataclasses import dataclass from typing import TYPE_CHECKING -from aare.common.raster_grid import RasterGridRequest +from aarecommon.models.raster_grid import RasterGridRequest + from aare.daq.mlbox import MlBox from aare.daq.operations.common.models import ( BaseOperationContext, @@ -36,4 +37,4 @@ class RasterContext(BaseOperationContext[RasterDependencies, RasterSettings]): @dataclass class RasterBoundingBoxResult: request: RasterGridRequest | None - uploaded_filename: str | None = None \ No newline at end of file + uploaded_filename: str | None = None diff --git a/src/aare/daq/operations/raster/service.py b/src/aare/daq/operations/raster/service.py index bc915783..c99ee63d 100644 --- a/src/aare/daq/operations/raster/service.py +++ b/src/aare/daq/operations/raster/service.py @@ -3,24 +3,37 @@ import time from math import ceil, floor import cv2 -from aareDB import SampleEventType -from jfjoch_client.exceptions import NotFoundException - -from aare.common.coordinate import AerotechCoordinate, Coordinate, SmargonCoordinate -from aare.common.exception_handler import AutoRasterSampleSkipped, RasterScanException -from aare.common.find_xtal import raster_highest_score, get_xtal_size, get_best_res, \ - get_best_b_factor, compute_crystal_score_array -from aare.common.logger_events import ( +from aarecommon.config.beamline import cfg_get +from aarecommon.config.logger_events import ( geom_log_context, log_ml_bundle_meta, merge_log_context, raster_request_log_context, sample_log_context, ) -from aare.common.beamline import cfg_get -from aare.common.models import BeamlineStateEnum, MLBoxType -from aare.common.raster_grid import CompletedRasterGrid, CompletedRasterGridElem, RasterGridRequest, grid_to_image_id -from aare.common.simulate_raster import generate_no_beam_scan_result +from aarecommon.errors.exception_handler import ( + AutoRasterSampleSkipped, + RasterScanException, +) +from aarecommon.math.coordinate import AerotechCoordinate, Coordinate, SmargonCoordinate +from aarecommon.math.find_xtal import ( + compute_crystal_score_array, + get_best_b_factor, + get_best_res, + get_xtal_size, + raster_highest_score, +) +from aarecommon.math.raster_grid import grid_to_image_id +from aarecommon.math.simulate_raster import generate_no_beam_scan_result +from aarecommon.models.models import BeamlineStateEnum, MLBoxType +from aarecommon.models.raster_grid import ( + CompletedRasterGrid, + CompletedRasterGridElem, + RasterGridRequest, +) +from aareDB import SampleEventType +from jfjoch_client.exceptions import NotFoundException + from aare.daq.mlbox import MLBoxPredictionResult from aare.daq.operations.common.ml_bounding_box import ( MLRasterPlan, @@ -92,7 +105,9 @@ class RasterService: return try: - diffraction_image = self.ctx.deps.jfjoch.get_diffraction_image(image_id, show_spots=True, show_res_est=True, show_beam_center=True) + diffraction_image = self.ctx.deps.jfjoch.get_diffraction_image( + image_id, show_spots=True, show_res_est=True, show_beam_center=True + ) except NotFoundException: self.logger.warning( "JFJoch diffraction preview image was not found after raster; continuing without upload", @@ -125,7 +140,9 @@ class RasterService: except Exception: pass - self.ctx.deps.aare.upload_jpg(sample_id, filename, diffraction_image, message=comment) + self.ctx.deps.aare.upload_jpg( + sample_id, filename, diffraction_image, message=comment + ) def ml_bounding_box( self, @@ -162,11 +179,18 @@ class RasterService: ) @staticmethod - def _box_touches_frame_edge(box, width: int | None, height: int | None, margin: int = 2) -> bool: + def _box_touches_frame_edge( + box, width: int | None, height: int | None, margin: int = 2 + ) -> bool: if width is None or height is None: return False x1, y1, x2, y2 = box - return x1 <= margin or y1 <= margin or x2 >= width - margin or y2 >= height - margin + return ( + x1 <= margin + or y1 <= margin + or x2 >= width - margin + or y2 >= height - margin + ) def _zoom_to_fit_box(self, plan: MLRasterPlan) -> tuple | None: """Pick the box that drives zoom-to-fit: loop_all if detected and not @@ -224,20 +248,22 @@ class RasterService: self.ctx.deps.devs.samcam_auto(AutoEnum.ONCE) def auto_center_line_scan_top_left( - self, - *, - omega_deg: float, - file_prefix: str | None, - grid_size_mm: Coordinate, - default_n_y: int = 50, - y_retarget_threshold_mm: float | None = None, - y_padding_fraction_each_side: float | None = None, + self, + *, + omega_deg: float, + file_prefix: str | None, + grid_size_mm: Coordinate, + default_n_y: int = 50, + y_retarget_threshold_mm: float | None = None, + y_padding_fraction_each_side: float | None = None, ) -> tuple[SmargonCoordinate, int]: geom = self.ctx.sample_geometry beam_x_pxl = geom.beam_location_pxl.x beam_y_pxl = geom.beam_location_pxl.y - line_scan_centre = geom.picture_to_smargon(Coordinate(x=beam_x_pxl, y=beam_y_pxl)) + line_scan_centre = geom.picture_to_smargon( + Coordinate(x=beam_x_pxl, y=beam_y_pxl) + ) n_y = default_n_y prediction_result: MLBoxPredictionResult = self.ctx.deps.mlbox.predict( @@ -259,9 +285,7 @@ class RasterService: prediction_type = getattr(prediction_box, "type", None) if prediction_type == MLBoxType.LOOP_FACE: - y_padding_fraction_each_side = ( - local_contact_config.line_scan_loop_face_y_padding_fraction_each_side - ) + y_padding_fraction_each_side = local_contact_config.line_scan_loop_face_y_padding_fraction_each_side elif prediction_type == MLBoxType.LOOP_ALL: y_padding_fraction_each_side = ( local_contact_config.line_scan_loop_all_y_padding_fraction_each_side @@ -328,9 +352,17 @@ class RasterService: ), ) - if prediction_box is not None and prediction_box.box is not None and grid_size_mm.y > 0: + if ( + prediction_box is not None + and prediction_box.box is not None + and grid_size_mm.y > 0 + ): box_height_pxl = abs(prediction_box.box.bottom_y - prediction_box.box.top_y) - padded_height_mm = box_height_pxl * geom.pixel_in_mm * (1.0 + 2.0 * y_padding_fraction_each_side) + padded_height_mm = ( + box_height_pxl + * geom.pixel_in_mm + * (1.0 + 2.0 * y_padding_fraction_each_side) + ) n_y = max(1, int(ceil(padded_height_mm / grid_size_mm.y))) self.logger.info( "Computed second auto-center raster y size from ML box height", @@ -358,7 +390,8 @@ class RasterService: "file_prefix": file_prefix, "omega_deg": omega_deg, "default_n_y": default_n_y, - "has_prediction_box": prediction_box is not None and prediction_box.box is not None, + "has_prediction_box": prediction_box is not None + and prediction_box.box is not None, }, ), ) @@ -433,7 +466,9 @@ class RasterService: if not self.ctx.deps.cfg.simulated_detector: self.ctx.deps.jfjoch.wait_till_running(timeout=60.0) else: - self.logger.info("Simulated detector mode enabled; faking jfjoch intilalisation.") + self.logger.info( + "Simulated detector mode enabled; faking jfjoch intilalisation." + ) self.logger.debug( f"Starting grid scan with {request.n_x}x{request.n_y} points, exp time {request.exp_time_s}s" @@ -470,7 +505,9 @@ class RasterService: y = None if self.ctx.deps.cfg.simulated_detector: - self.logger.info("Simulated detector mode enabled; using fake raster result.") + self.logger.info( + "Simulated detector mode enabled; using fake raster result." + ) scan_result = generate_no_beam_scan_result(request) com = raster_highest_score(scan_result.images) target_coor = com.get_com_mm(request) @@ -491,7 +528,9 @@ class RasterService: {"exp_time_s": request.exp_time_s}, ), ) - raise RasterScanException("JFJoch returned no ScanResult for raster") + raise RasterScanException( + "JFJoch returned no ScanResult for raster" + ) com = raster_highest_score(scan_result.images) if com is None: @@ -502,7 +541,9 @@ class RasterService: x = ((request.n_x - 1) * request.grid_size_mm.x) / 2.0 y = ((request.n_y - 1) * request.grid_size_mm.y) / 2.0 - target_coor_offset = self.ctx.sample_geometry.smargon_nudge(Coordinate(x=x, y=y)) + target_coor_offset = self.ctx.sample_geometry.smargon_nudge( + Coordinate(x=x, y=y) + ) self.logger.info( "Calculated raster centre offset", @@ -515,16 +556,24 @@ class RasterService: "centre_offset_z_mm": target_coor_offset.z, "grid_half_width_x_mm": x, "grid_half_height_y_mm": y, - "top_left_x_mm": getattr(request.smargon_top_left.sh_mm, "x", None), - "top_left_y_mm": getattr(request.smargon_top_left.sh_mm, "y", None), - "top_left_z_mm": getattr(request.smargon_top_left.sh_mm, "z", None), + "top_left_x_mm": getattr( + request.smargon_top_left.sh_mm, "x", None + ), + "top_left_y_mm": getattr( + request.smargon_top_left.sh_mm, "y", None + ), + "top_left_z_mm": getattr( + request.smargon_top_left.sh_mm, "z", None + ), }, ), ) else: self.logger.info("Calcualted COM is not None, procedding") target_coor = com.get_com_mm(request) - target_coor_offset = self.ctx.sample_geometry.smargon_nudge(target_coor) + target_coor_offset = self.ctx.sample_geometry.smargon_nudge( + target_coor + ) self.logger.info( "Calculated raster centre offset", extra=merge_log_context( @@ -536,15 +585,22 @@ class RasterService: "centre_offset_z_mm": target_coor_offset.z, "grid_half_width_x_mm": x, "grid_half_height_y_mm": y, - "top_left_x_mm": getattr(request.smargon_top_left.sh_mm, "x", None), - "top_left_y_mm": getattr(request.smargon_top_left.sh_mm, "y", None), - "top_left_z_mm": getattr(request.smargon_top_left.sh_mm, "z", None), + "top_left_x_mm": getattr( + request.smargon_top_left.sh_mm, "x", None + ), + "top_left_y_mm": getattr( + request.smargon_top_left.sh_mm, "y", None + ), + "top_left_z_mm": getattr( + request.smargon_top_left.sh_mm, "z", None + ), }, ), ) - - self.logger.info(f"moving Smargon to grid centre offset {target_coor_offset}") + self.logger.info( + f"moving Smargon to grid centre offset {target_coor_offset}" + ) target_smargon = SmargonCoordinate( sh_mm=request.smargon_top_left.sh_mm + target_coor_offset, phi_deg=request.smargon_top_left.phi_deg, @@ -569,9 +625,15 @@ class RasterService: ), ) - sample_id = self.ctx.sample.db_id if self.ctx.sample and self.ctx.sample.db_id is not None else None + sample_id = ( + self.ctx.sample.db_id + if self.ctx.sample and self.ctx.sample.db_id is not None + else None + ) if sample_id: - self.logger.debug("moving to XtalSnapshot to take a screenshot of the sample") + self.logger.debug( + "moving to XtalSnapshot to take a screenshot of the sample" + ) if self.ctx.services.state is None: raise RuntimeError("RasterService requires services.state") self.ctx.services.state.set_state(BeamlineStateEnum.XtalSnapshot) @@ -591,7 +653,9 @@ class RasterService: raster_request=request, geom=self.ctx.sample_geometry, com=None, - beam_mark_pxl=self.ctx.deps.cfg.get_beam_mark(self.ctx.deps.devs.zoom), + beam_mark_pxl=self.ctx.deps.cfg.get_beam_mark( + self.ctx.deps.devs.zoom + ), ) else: self.ctx.deps.aare.ingest_gridscan( @@ -600,16 +664,18 @@ class RasterService: raster_request=request, geom=self.ctx.sample_geometry, com=None, - beam_mark_pxl=self.ctx.deps.cfg.get_beam_mark(self.ctx.deps.devs.zoom), + beam_mark_pxl=self.ctx.deps.cfg.get_beam_mark( + self.ctx.deps.devs.zoom + ), ) if com is not None and com.max_image is not None: - diffraction_image_filename = ( - f"{self.ctx.sample.db_id}_best_diffraction_from_raster_image_{com.max_image}" - ) + diffraction_image_filename = f"{self.ctx.sample.db_id}_best_diffraction_from_raster_image_{com.max_image}" diffraction_image_id = com.max_image else: - diffraction_image_filename = f"{self.ctx.sample.db_id}_diffraction_image_near_grid_scan" + diffraction_image_filename = ( + f"{self.ctx.sample.db_id}_diffraction_image_near_grid_scan" + ) diffraction_image_id = self._grid_image_id_from_centre_offset( x_mm=x, y_mm=y, @@ -648,9 +714,17 @@ class RasterService: }, ), ) - self.ctx.deps.cfg.crystal_size = get_xtal_size(self.ctx.deps.cfg.crystal_size, compute_crystal_score_array(scan_result.images), r=request) - self.ctx.deps.cfg.last_best_res = get_best_res(result_list= scan_result.images) - self.ctx.deps.cfg.last_best_b_factor = get_best_b_factor(result_list=scan_result.images) + self.ctx.deps.cfg.crystal_size = get_xtal_size( + self.ctx.deps.cfg.crystal_size, + compute_crystal_score_array(scan_result.images), + r=request, + ) + self.ctx.deps.cfg.last_best_res = get_best_res( + result_list=scan_result.images + ) + self.ctx.deps.cfg.last_best_b_factor = get_best_b_factor( + result_list=scan_result.images + ) return CompletedRasterGridElem( request=copy.deepcopy(request), result=scan_result, @@ -668,7 +742,9 @@ class RasterService: ) raise - def execute_auto_center(self, request: RasterGridRequest) -> CompletedRasterGrid | None: + def execute_auto_center( + self, request: RasterGridRequest + ) -> CompletedRasterGrid | None: sample = self.ctx.sample if sample is None: raise Exception("Sample must be mounted to auto center") @@ -719,7 +795,9 @@ class RasterService: ) self.ctx.deps.devs.aerotech_omega = geom.omega_deg + 90.0 time.sleep(0.2) - plan = self.ml_raster_plan(sample.db_id, f"ml_{geom.omega_deg + 90.0:.2f}deg") + plan = self.ml_raster_plan( + sample.db_id, f"ml_{geom.omega_deg + 90.0:.2f}deg" + ) if plan is None: self.logger.error( @@ -750,9 +828,15 @@ class RasterService: "ml_n_y": r.n_y, "ml_grid_size_x_mm": getattr(r.grid_size_mm, "x", None), "ml_grid_size_y_mm": getattr(r.grid_size_mm, "y", None), - "ml_top_left_x_mm": getattr(getattr(r.smargon_top_left, "sh_mm", None), "x", None), - "ml_top_left_y_mm": getattr(getattr(r.smargon_top_left, "sh_mm", None), "y", None), - "ml_top_left_z_mm": getattr(getattr(r.smargon_top_left, "sh_mm", None), "z", None), + "ml_top_left_x_mm": getattr( + getattr(r.smargon_top_left, "sh_mm", None), "x", None + ), + "ml_top_left_y_mm": getattr( + getattr(r.smargon_top_left, "sh_mm", None), "y", None + ), + "ml_top_left_z_mm": getattr( + getattr(r.smargon_top_left, "sh_mm", None), "z", None + ), }, ), ) @@ -772,7 +856,9 @@ class RasterService: self.ctx.deps.jfjoch.measure_raster(grid, status) self.logger.info("detector initialised") else: - self.logger.info("Simulated detector mode enabled; using fake raster result.") + self.logger.info( + "Simulated detector mode enabled; using fake raster result." + ) if self.ctx.services.state is None: raise RuntimeError("RasterService requires services.state") @@ -810,7 +896,9 @@ class RasterService: grid.n_x = 1 grid.file_prefix = f"{old_prefix}_{grid.omega_deg}deg" geom = self.ctx.sample_geometry - grid.grid_size_mm = Coordinate(x=geom.beam_size_mm.x, y=geom.beam_size_mm.y * 0.25) + grid.grid_size_mm = Coordinate( + x=geom.beam_size_mm.x, y=geom.beam_size_mm.y * 0.25 + ) grid.smargon_top_left, grid.n_y = self.auto_center_line_scan_top_left( omega_deg=grid.omega_deg, file_prefix=grid.file_prefix, @@ -841,7 +929,9 @@ class RasterService: self.ctx.deps.jfjoch.measure_raster(grid, status) self.logger.info("detector initialised") else: - self.logger.info("Simulated detector mode enabled; using fake raster result.") + self.logger.info( + "Simulated detector mode enabled; using fake raster result." + ) if self.ctx.services.state is None: raise RuntimeError("RasterService requires services.state") @@ -856,4 +946,4 @@ class RasterService: ) res2 = self.execute(grid) - return CompletedRasterGrid(r=[res1, res2]) \ No newline at end of file + return CompletedRasterGrid(r=[res1, res2]) diff --git a/src/aare/daq/operations/rotation/service.py b/src/aare/daq/operations/rotation/service.py index 6f5691b9..d6dfdcce 100644 --- a/src/aare/daq/operations/rotation/service.py +++ b/src/aare/daq/operations/rotation/service.py @@ -1,11 +1,11 @@ import copy import time +from aarecommon.math.coordinate import SmargonCoordinate +from aarecommon.models.models import BeamlineStateEnum +from aarecommon.models.rotation_scan import CompletedRotationScan, RotationScanRequest from aareDB import SampleEventType -from aare.common.coordinate import SmargonCoordinate -from aare.common.models import BeamlineStateEnum -from aare.common.rotation_scan import CompletedRotationScan, RotationScanRequest from aare.daq.operations.common.simulate_scan_result import build_fake_rotation_result from aare.daq.operations.rotation.models import RotationContext @@ -66,7 +66,9 @@ class RotationService: if request.start is not None and request.end is not None: smargon_time_step = request.exp_time_s / float(request.steps) - pos_step = (request.end.sh_mm - request.start.sh_mm) * (1.0 / 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.ctx.deps.devs.smargon.target = SmargonCoordinate( @@ -74,11 +76,15 @@ class RotationService: ) time.sleep(smargon_time_step) - self.ctx.deps.devs.aerotech.wait_till_done(timeout=int(round(total_time + 60, 0))) + self.ctx.deps.devs.aerotech.wait_till_done( + timeout=int(round(total_time + 60, 0)) + ) self.ctx.deps.devs.aerotech_omega = omega_start if self.ctx.deps.cfg.simulated_detector: - self.logger.warning("Detector in simulation mode, returning fake zero rotation result.") + self.logger.warning( + "Detector in simulation mode, returning fake zero rotation result." + ) return build_fake_rotation_result(request, start_angle=float(omega_start)) scan_result = self.ctx.deps.jfjoch.wait_till_done(60) @@ -124,9 +130,7 @@ class RotationService: idx, (", " + ", ".join(bits)) if bits else "" ) filename = f"{sample.db_id}_screening_diffraction_{image_id}" - self.ctx.deps.aare.upload_jpg( - sample.db_id, filename, jpg, message=comment - ) + self.ctx.deps.aare.upload_jpg(sample.db_id, filename, jpg, message=comment) def run(self, request: RotationScanRequest) -> CompletedRotationScan: sample = self.ctx.sample @@ -137,7 +141,11 @@ class RotationService: self.ctx.services.datacollection.prepare(request) self._prepare_detector(request) - if sample is not None and sample.db_id is not None and self.ctx.services.events is not None: + if ( + sample is not None + and sample.db_id is not None + and self.ctx.services.events is not None + ): self.ctx.services.events.send(sample.db_id, SampleEventType.COLLECTING) if self.ctx.services.state is not None: @@ -158,7 +166,9 @@ class RotationService: # For screening runs, ingest the per-wedge diffraction images. Done # before the COLLECTED event so they pin to the same sample event the # run is bound to (the COLLECTING event), matching the preview above. - self.logger.debug(f"Ingesting screening diffraction images: {request.screening}") + self.logger.debug( + f"Ingesting screening diffraction images: {request.screening}" + ) if request.screening and not self.ctx.deps.cfg.simulated_detector: self._ingest_screening_diffraction(sample, result.result) @@ -170,7 +180,9 @@ class RotationService: sample=sample, result=result.result, geom=self.ctx.sample_geometry, - beam_mark_pxl=self.ctx.deps.cfg.get_beam_mark(self.ctx.deps.devs.zoom), + beam_mark_pxl=self.ctx.deps.cfg.get_beam_mark( + self.ctx.deps.devs.zoom + ), ) - return result \ No newline at end of file + return result diff --git a/src/aare/daq/operations/screenshot/service.py b/src/aare/daq/operations/screenshot/service.py index 5906aa49..18561a21 100644 --- a/src/aare/daq/operations/screenshot/service.py +++ b/src/aare/daq/operations/screenshot/service.py @@ -1,13 +1,13 @@ +import time from dataclasses import dataclass from datetime import datetime from pathlib import Path -import time from typing import Protocol import cv2 import numpy as np +from aarecommon.models.models import SampleShortInfo -from aare.common.models import SampleShortInfo from aare.daq.aaredb import AareWrapper from aare.daq.mlbox import MlBox from aare.devices.mx_lib import clean_filename @@ -68,7 +68,9 @@ class ScreenshotService: self.logger.debug(f"saving screenshot {filename} from inference image") cv2.imwrite(f"{self.output_dir}/{filename}.jpg", bgr_image) - def save_to_db(self, sample_id: int, filename: str, settle_time_s: float = 0.2) -> None: + def save_to_db( + self, sample_id: int, filename: str, settle_time_s: float = 0.2 + ) -> None: time.sleep(settle_time_s) sample = self._current_sample() bgr_image = self._get_inference_image() @@ -100,9 +102,13 @@ class ScreenshotService: if filename: safe_filename = clean_filename(filename) if not pgroup: - raise ValueError("No active pgroup set; cannot save screenshot to photos directory.") + raise ValueError( + "No active pgroup set; cannot save screenshot to photos directory." + ) - photos_dir = Path(self.photos_root) / pgroup / "raw" / "photos" / str(sample_id) + photos_dir = ( + Path(self.photos_root) / pgroup / "raw" / "photos" / str(sample_id) + ) photos_dir.mkdir(parents=True, exist_ok=True) photo_path = photos_dir / f"{safe_filename}.jpeg" cv2.imwrite(str(photo_path), bgr_image) @@ -113,10 +119,12 @@ class ScreenshotService: final_message = (message or "").strip() or default_message def _upload() -> None: - self.aare.upload_image(sample_id, upload_name, bgr_image, message=final_message) + self.aare.upload_image( + sample_id, upload_name, bgr_image, message=final_message + ) self.run_noncritical( _upload, description=f"send screenshot '{upload_name}'", sample=sample, - ) \ No newline at end of file + ) diff --git a/src/aare/daq/server.py b/src/aare/daq/server.py index 28fec1ac..451f5b2a 100644 --- a/src/aare/daq/server.py +++ b/src/aare/daq/server.py @@ -1,50 +1,60 @@ import asyncio import hmac import io -import os, time +import json +import os import random +import time from contextlib import asynccontextmanager from typing import AsyncGenerator, Optional -import json + import cv2 import urllib3 import uvicorn +from aarecommon.config.beamline import mx_beamline +from aarecommon.config.logger import get_uvicorn_logging_config, setup_logger +from aarecommon.errors.codes import AareErrorCode, export_error_codes_grouped +from aarecommon.errors.exception_handler import ( + MaintenanceStateException, + SampleException, + UserRightsException, +) +from aarecommon.math.coordinate import AerotechCoordinate, Coordinate, SmargonCoordinate +from aarecommon.math.sample_geometry import SampleGeometryModel +from aarecommon.models.auth import BatonRequestStatus, BatonStatus +from aarecommon.models.automation import AutomationProgress +from aarecommon.models.models import ( + AutofocusSettings, + BeamlineSettingsModel, + BeamlineStateEnum, + CryojetSettingsModel, + CrystalSize, + DAQStatusModel, + FluorescenceSpectrumOutputModel, + FluorescenceSpectrumParameterModel, + RecoveryActionRequest, + SampleCameraSettings, + SampleShortInfo, + SampleShortInfoList, + SessionStatus, + SimpleScanParameters, + TokenData, +) +from aarecommon.models.raster_grid import CompletedRasterGrid, RasterGridRequest +from aarecommon.models.rotation_scan import CompletedRotationScan, RotationScanRequest from aareDB import SampleEventType - -from aare.common.coordinate import AerotechCoordinate -from aare.common.auth_models import BatonStatus, BatonRequestStatus -from aare.common.coordinate import SmargonCoordinate, Coordinate -from aare.common.error_codes import export_error_codes_grouped, AareErrorCode -from aare.common.logger_config import setup_logger, get_uvicorn_logging_config -from aare.common.models import SampleShortInfo, DAQStatusModel, BeamlineStateEnum, BeamlineSettingsModel, \ - SampleShortInfoList, SessionStatus, SampleCameraSettings, AutofocusSettings, TokenData, \ - CryojetSettingsModel, SimpleScanParameters, CrystalSize, FluorescenceSpectrumParameterModel, \ - FluorescenceSpectrumOutputModel, RecoveryActionRequest -from aare.common.automation_models import AutomationProgress -from aare.common.raster_grid import RasterGridRequest, CompletedRasterGrid -from aare.common.rotation_scan import RotationScanRequest, CompletedRotationScan -from aare.common.sample_geometry import SampleGeometryModel -from fastapi import FastAPI, Depends, Request -from fastapi.concurrency import run_in_threadpool -from fastapi import HTTPException +from fastapi import Depends, FastAPI, HTTPException, Request from fastapi import status as api_status +from fastapi.concurrency import run_in_threadpool from fastapi.security import OAuth2PasswordBearer, OAuth2PasswordRequestForm from starlette.responses import StreamingResponse from aare.daq import auth -from aare.common.beamline import mx_beamline from aare.daq.config import BeamlineConfig -from aare.daq.daq import AareDAQ from aare.daq.config_model import LocalContactConfigModel - +from aare.daq.daq import AareDAQ from aare.daq.server_exception_handler import register_exception_handlers -from aare.common.exception_handler import ( - SampleException, - UserRightsException, - MaintenanceStateException, - ) - logger = setup_logger("aareDAQ") # OAuth2 setup @@ -73,6 +83,7 @@ _automation_progress_state: dict = { } _automation_progress_state_lock = asyncio.Lock() + @asynccontextmanager async def lifespan(application: FastAPI): """ @@ -121,9 +132,11 @@ async def lifespan(application: FastAPI): # Shutdown: add cleanup here if needed logger.info(f"Worker {os.getpid()} shutting down.") + app = FastAPI(lifespan=lifespan) register_exception_handlers(app) + def _required_recovery_code() -> str: """ Get the required recovery confirmation code from environment variables. @@ -143,6 +156,7 @@ def _required_recovery_code() -> str: ) return code + def _validate_recovery_code(confirmation_code: str) -> None: """ Validate the provided recovery confirmation code against the server configuration. @@ -162,6 +176,7 @@ def _validate_recovery_code(confirmation_code: str) -> None: detail="Invalid confirmation code.", ) + def _sample_is_mounted() -> bool: """ Check if a sample is currently mounted on the beamline. @@ -174,6 +189,7 @@ def _sample_is_mounted() -> bool: except Exception: return False + def _push_face_detection_progress(payload: dict) -> None: """ Update the global face detection state with new progress information. @@ -191,6 +207,7 @@ def _push_face_detection_progress(payload: dict) -> None: except Exception as e: logger.warning(f"Failed to update face detection progress: {e}") + def _get_automation_progress_state() -> dict: """ Read automation progress state from shared config/Redis storage. @@ -204,6 +221,7 @@ def _get_automation_progress_state() -> dict: logger.warning(f"Failed to read automation progress from Redis: {e}") return {"seq": 0, "progress": None} + def _push_automation_progress(progress: AutomationProgress) -> None: """ Update the shared automation progress state. @@ -222,9 +240,8 @@ def _push_automation_progress(progress: AutomationProgress) -> None: f"current_step={state.get('progress', {}).get('current_step')}" ) except Exception as e: - logger.warning( - f"Failed to update automation progress: {type(e).__name__}: {e}" - ) + logger.warning(f"Failed to update automation progress: {type(e).__name__}: {e}") + async def face_detection_event_stream() -> AsyncGenerator[str, None]: """ @@ -245,6 +262,7 @@ async def face_detection_event_stream() -> AsyncGenerator[str, None]: except asyncio.CancelledError: return + async def automation_progress_event_stream() -> AsyncGenerator[str, None]: """ Generator for Server-Sent Events (SSE) of automation progress. @@ -264,6 +282,7 @@ async def automation_progress_event_stream() -> AsyncGenerator[str, None]: except asyncio.CancelledError: return + @app.post("/token") async def login(request: Request): """ @@ -284,6 +303,7 @@ async def login(request: Request): data = await run_in_threadpool(auth.authenticate_user, cfg, username) return {"access_token": data, "token_type": "bearer"} + @app.get("/meta/error-codes") async def meta_error_codes() -> dict[str, dict[str, str]]: """ @@ -292,6 +312,7 @@ async def meta_error_codes() -> dict[str, dict[str, str]]: """ return export_error_codes_grouped() + @app.get("/status") async def status(token: str = Depends(oauth2_scheme)) -> DAQStatusModel: """ @@ -325,14 +346,18 @@ async def status(token: str = Depends(oauth2_scheme)) -> DAQStatusModel: full.sample = sample if in_ro and sample_view_allowed else None full.box = getattr(full, "box", None) if in_ro else None full.last_best_res = getattr(full, "last_best_res", None) if in_ro else None - full.last_best_b_factor = getattr(full, "last_best_b_factor", None) if in_ro else None - full.crystal_size = getattr( - full, "crystal_size", CrystalSize(x=0, y=0, z=0) - ) if in_ro else CrystalSize(x=0, y=0, z=0) + full.last_best_b_factor = ( + getattr(full, "last_best_b_factor", None) if in_ro else None + ) + full.crystal_size = ( + getattr(full, "crystal_size", CrystalSize(x=0, y=0, z=0)) + if in_ro + else CrystalSize(x=0, y=0, z=0) + ) full.session = SessionStatus( current_pgroup=cfg.pgroup, session=cfg.session_state(data.session), - staff=data.staff + staff=data.staff, ) if data.staff: @@ -341,9 +366,11 @@ async def status(token: str = Depends(oauth2_scheme)) -> DAQStatusModel: own_gui = cfg.get_gui_session(data.session) full.open_guis = [own_gui] if own_gui is not None else [] return full + + @app.get("/beamline/geometry") async def sample_geometry( - token: str = Depends(oauth2_scheme), + token: str = Depends(oauth2_scheme), ) -> SampleGeometryModel: """ Get the sample geometry information. @@ -375,6 +402,7 @@ async def omega(val: float, token: str = Depends(oauth2_scheme)): daq.omega = val return "OK" + @app.put("/beamline/omega_rel") async def omega(val: float, token: str = Depends(oauth2_scheme)): """ @@ -392,6 +420,7 @@ async def omega(val: float, token: str = Depends(oauth2_scheme)): daq.omega_rel(val) return "OK" + @app.put("/beamline/front_light") async def front_light(val: float, token: str = Depends(oauth2_scheme)): """ @@ -409,6 +438,7 @@ async def front_light(val: float, token: str = Depends(oauth2_scheme)): daq.front_light = val return "OK" + @app.put("/beamline/back_light") async def back_light(val: float, token: str = Depends(oauth2_scheme)): """ @@ -426,6 +456,7 @@ async def back_light(val: float, token: str = Depends(oauth2_scheme)): daq.back_light = val return "OK" + @app.put("/beamline/zoom") async def zoom(val: float, token: str = Depends(oauth2_scheme)): """ @@ -463,7 +494,9 @@ async def mono_pitch_scan(plot: bool = False, token: str = Depends(oauth2_scheme @app.put("/beamline/change_energy") -async def change_energy(value: float, plot: bool = False, token: str = Depends(oauth2_scheme)): +async def change_energy( + value: float, plot: bool = False, token: str = Depends(oauth2_scheme) +): """ Change monochromator energy. Staff only. @@ -501,7 +534,9 @@ async def smargon(val: SmargonCoordinate, token: str = Depends(oauth2_scheme)): @app.post("/beamline/tweak_abr_meas_pos") -async def tweak_abr_meas_pos(val: AerotechCoordinate, token: str = Depends(oauth2_scheme)): +async def tweak_abr_meas_pos( + val: AerotechCoordinate, token: str = Depends(oauth2_scheme) +): """ Tweak the Aerotech measurement position. Staff only. @@ -533,6 +568,7 @@ async def save_abr_meas_pos(token: str = Depends(oauth2_scheme)): daq.save_abr_meas_pos() return "OK" + @app.post("/beamline/save_beam_location_camera_setting") async def save_beam_location_camera_setting(token: str = Depends(oauth2_scheme)): """ @@ -564,6 +600,7 @@ async def anneal(time_s: float, token: str = Depends(oauth2_scheme)): daq.anneal(time_s) return "OK" + @app.post("/smargon/initialize") async def initialise_smargon(token: str = Depends(oauth2_scheme)) -> dict: """ @@ -597,6 +634,7 @@ async def bec_load_user_macros(token: str = Depends(oauth2_scheme)) -> dict: "message": "BEC user macros loaded.", } + @app.get("/bec/user_macros") async def bec_list_all_user_macros(token: str = Depends(oauth2_scheme)) -> list: """ @@ -606,6 +644,7 @@ async def bec_list_all_user_macros(token: str = Depends(oauth2_scheme)) -> list: auth.check_jwt_staff_only(data) return daq.bec_list_all_user_macros() + @app.get("/bec/devices") async def bec_list_all_devices(token: str = Depends(oauth2_scheme)) -> list: """ @@ -615,10 +654,11 @@ async def bec_list_all_devices(token: str = Depends(oauth2_scheme)) -> list: auth.check_jwt_staff_only(data) return daq.bec_list_all_devices() + @app.post("/bec/reinitialise_planner_and_position_devices") async def bec_reinitialise_planner_and_position_devices( - method: str = "auto", - token: str = Depends(oauth2_scheme), + method: str = "auto", + token: str = Depends(oauth2_scheme), ) -> dict: """ Reinitialise BEC planner and position devices. Staff only. @@ -640,6 +680,7 @@ async def bec_reinitialise_planner_and_position_devices( "message": "BEC planner and position devices reinitialised.", } + @app.post("/bec/save_current_bs_pos") async def bec_save_current_bs_pos(token: str = Depends(oauth2_scheme)) -> dict: """ @@ -653,6 +694,7 @@ async def bec_save_current_bs_pos(token: str = Depends(oauth2_scheme)) -> dict: "message": "Saved current BEC beamstop work position.", } + @app.post("/bec/save_current_collimator_pos") async def bec_save_current_collimator_pos(token: str = Depends(oauth2_scheme)) -> dict: """ @@ -666,8 +708,11 @@ async def bec_save_current_collimator_pos(token: str = Depends(oauth2_scheme)) - "message": "Saved current BEC collimator work position.", } + @app.post("/bec/save_current_aerotech_position") -async def bec_save_current_aerotech_position(token: str = Depends(oauth2_scheme)) -> dict: +async def bec_save_current_aerotech_position( + token: str = Depends(oauth2_scheme), +) -> dict: """ Save the current BEC aerotech work position and reload device config. Staff only. """ @@ -679,6 +724,7 @@ async def bec_save_current_aerotech_position(token: str = Depends(oauth2_scheme) "message": "Saved current BEC aerotech work position and reloaded devices.", } + def initialise_aerotech(self): self.__cfg.try_set_busy(timeout=360) try: @@ -689,6 +735,7 @@ def initialise_aerotech(self): logger.error(f"Failed to initialise Aerotech: {e}") raise + def detector_take_pedestal(self): self.__cfg.try_set_busy(timeout=360) try: @@ -699,6 +746,7 @@ def detector_take_pedestal(self): logger.error(f"Failed to take detector pedestal: {e}") raise + def initialise_detector(self): self.__cfg.try_set_busy(timeout=360) try: @@ -709,6 +757,7 @@ def initialise_detector(self): logger.error(f"Failed to initialise detector: {e}") raise + @app.get("/local_contact/simulation_state") async def local_contact_simulation_state(token: str = Depends(oauth2_scheme)) -> dict: """ @@ -728,6 +777,7 @@ async def local_contact_device_state(token: str = Depends(oauth2_scheme)) -> dic auth.check_jwt_staff_only(data) return daq.get_local_contact_device_state() + @app.get("/local_contact/links") async def local_contact_links(token: str = Depends(oauth2_scheme)) -> dict: """ @@ -737,11 +787,12 @@ async def local_contact_links(token: str = Depends(oauth2_scheme)) -> dict: auth.check_jwt_staff_only(data) return daq.get_local_contact_links() + @app.post("/local_contact/simulate/{device}") async def local_contact_set_simulation( - device: str, - enabled: bool, - token: str = Depends(oauth2_scheme), + device: str, + enabled: bool, + token: str = Depends(oauth2_scheme), ) -> dict: """ Enable or disable runtime simulation for a backend device and restart its wrapper. Staff only. @@ -752,10 +803,11 @@ async def local_contact_set_simulation( result["message"] = f"{device} simulation set to {enabled}." return result + @app.post("/local_contact/restart/{device}") async def local_contact_restart_device( - device: str, - token: str = Depends(oauth2_scheme), + device: str, + token: str = Depends(oauth2_scheme), ) -> dict: """ Restart a Local Contact backend wrapper. Staff only. @@ -783,9 +835,10 @@ async def local_contact_restart_device( result["message"] = f"{device} backend restarted." return result + @app.post("/local_contact/resync/detector_metadata") async def local_contact_resync_detector_metadata( - token: str = Depends(oauth2_scheme), + token: str = Depends(oauth2_scheme), ) -> dict: """ Refresh cached detector metadata and DTZ limits. Staff only. @@ -800,8 +853,11 @@ async def local_contact_resync_detector_metadata( "payload": payload, } + @app.get("/local_contact/config") -async def local_contact_config(token: str = Depends(oauth2_scheme)) -> LocalContactConfigModel: +async def local_contact_config( + token: str = Depends(oauth2_scheme), +) -> LocalContactConfigModel: """ Return Local Contact config values. Staff only. """ @@ -812,8 +868,8 @@ async def local_contact_config(token: str = Depends(oauth2_scheme)) -> LocalCont @app.put("/local_contact/config") async def local_contact_set_config( - payload: LocalContactConfigModel, - token: str = Depends(oauth2_scheme), + payload: LocalContactConfigModel, + token: str = Depends(oauth2_scheme), ) -> LocalContactConfigModel: """ Update Local Contact config values. Staff only. @@ -839,6 +895,7 @@ async def goto_abr_meas_pos(token: str = Depends(oauth2_scheme)): daq.goto_abr_meas_pos() return "OK" + @app.post("/beam_mark/add") async def mark_beam(x: float, y: float, token: str = Depends(oauth2_scheme)): """ @@ -874,6 +931,7 @@ async def clear_beam_mark(token: str = Depends(oauth2_scheme)): daq.clear_mark_beam() return "OK" + @app.post("/beamline/beam_center") async def beam_center(x: float, y: float, token: str = Depends(oauth2_scheme)): """ @@ -892,6 +950,7 @@ async def beam_center(x: float, y: float, token: str = Depends(oauth2_scheme)): daq.beam_center = (x, y) return "OK" + @app.post("/beamline/beam_size_mm") async def beam_size_mm(x: float, y: float, token: str = Depends(oauth2_scheme)): """ @@ -910,6 +969,7 @@ async def beam_size_mm(x: float, y: float, token: str = Depends(oauth2_scheme)): daq.beam_size_mm = Coordinate(x=x, y=y) return "OK" + @app.put("/beamline/samcam") async def samcam_settings(s: SampleCameraSettings, token: str = Depends(oauth2_scheme)): """ @@ -927,6 +987,7 @@ async def samcam_settings(s: SampleCameraSettings, token: str = Depends(oauth2_s daq.samcam_settings = s return "OK" + @app.put("/beamline/autoexposure") async def samcam_autoexposure(token: str = Depends(oauth2_scheme)): """ @@ -943,6 +1004,7 @@ async def samcam_autoexposure(token: str = Depends(oauth2_scheme)): daq.auto_exposure() return "OK" + @app.post("/samcam/autofocus") async def samcam_autofocus(s: AutofocusSettings, token: str = Depends(oauth2_scheme)): """ @@ -960,6 +1022,7 @@ async def samcam_autofocus(s: AutofocusSettings, token: str = Depends(oauth2_sch daq.auto_focus(s) return "OK" + @app.post("/beamline/shutter") async def shutter(val: bool, token: str = Depends(oauth2_scheme)): """ @@ -1026,7 +1089,7 @@ async def sample(token: str = Depends(oauth2_scheme)) -> SampleShortInfo: dewar_name="", db_id=-1, pin=s.pin, - location=s.location + location=s.location, ) @@ -1041,15 +1104,16 @@ async def park_and_dry(token: str = Depends(oauth2_scheme)): Returns: Dictionary with status and message. """ - #TODO: add unmount toggle, and combine aprk_and_dry and tell_dry to one command that takes park and unmount as signals + # TODO: add unmount toggle, and combine aprk_and_dry and tell_dry to one command that takes park and unmount as signals auth.check_jwt_rw(cfg, auth.parse_token(token)) logger.debug("Executing unmount, dry and park") - daq.park_and_dry(park=True, unmount = True) + daq.park_and_dry(park=True, unmount=True) return { "ok": True, "message": "TELL has been dried and parked", } + @app.post("/tell/dry") async def tell_dry(token: str = Depends(oauth2_scheme)) -> dict: """ @@ -1064,12 +1128,13 @@ async def tell_dry(token: str = Depends(oauth2_scheme)) -> dict: data = auth.parse_token(token) auth.check_jwt_rw(data) logger.debug("Executing dry and return to dewar") - daq.park_and_dry(park=False, unmount = False) + daq.park_and_dry(park=False, unmount=False) return { "ok": True, "message": "TELL dry cycle completed.", } + @app.post("/tell/toggle_blower") async def tell_toggle_blower(token: str = Depends(oauth2_scheme)) -> dict: """ @@ -1089,8 +1154,11 @@ async def tell_toggle_blower(token: str = Depends(oauth2_scheme)) -> dict: "message": "TELL blower toggled.", } + @app.post("/sample/mount") -async def mount(dbid: int, token: str = Depends(oauth2_scheme), reference: bool = False): +async def mount( + dbid: int, token: str = Depends(oauth2_scheme), reference: bool = False +): """ Mount a sample from the spreadsheet onto the goniometer. @@ -1128,6 +1196,7 @@ async def mount(dbid: int, token: str = Depends(oauth2_scheme), reference: bool return "OK" + @app.post("/sample/unmount") async def unmount(token: str = Depends(oauth2_scheme)): """ @@ -1160,6 +1229,7 @@ async def manual(s: SampleShortInfo, token: str = Depends(oauth2_scheme)): daq.create_sample(s) logger.debug(f"DB ID after creating {s.db_id}") + @app.post("/sample/resync") async def sample_resync(token: str = Depends(oauth2_scheme)) -> dict: """ @@ -1195,6 +1265,7 @@ def get_spreadsheet(data: TokenData) -> SampleShortInfoList: else: return cfg.spreadsheet_pgroup(data.pgroups) + def get_reference_tools() -> SampleShortInfoList: """ Get the list of reference tools. @@ -1204,6 +1275,7 @@ def get_reference_tools() -> SampleShortInfoList: """ return cfg.reference_tools + async def reference_tools_event_stream() -> AsyncGenerator[str, None]: """ Generator for SSE of reference tools updates. @@ -1218,6 +1290,7 @@ async def reference_tools_event_stream() -> AsyncGenerator[str, None]: except asyncio.CancelledError: return + async def spreadsheet_event_stream(data: TokenData) -> AsyncGenerator[str, None]: """ Generator for SSE of sample spreadsheet updates. @@ -1255,10 +1328,11 @@ async def spreadsheet_sse(token: str = Depends(oauth2_scheme)): "Cache-Control": "no-cache", "Connection": "keep-alive", "Access-Control-Allow-Origin": "*", - "Access-Control-Allow-Headers": "Cache-Control" - } + "Access-Control-Allow-Headers": "Cache-Control", + }, ) + @app.get("/sse/reference_tools") async def reference_tools_sse(): """ @@ -1274,10 +1348,11 @@ async def reference_tools_sse(): "Cache-Control": "no-cache", "Connection": "keep-alive", "Access-Control-Allow-Origin": "*", - "Access-Control-Allow-Headers": "Cache-Control" - } + "Access-Control-Allow-Headers": "Cache-Control", + }, ) + @app.get("/sample/spreadsheet") async def spreadsheet(token: str = Depends(oauth2_scheme)) -> SampleShortInfoList: """ @@ -1291,6 +1366,7 @@ async def spreadsheet(token: str = Depends(oauth2_scheme)) -> SampleShortInfoLis """ return get_spreadsheet(auth.parse_token(token)) + @app.get("/sample/reference_tools") async def reference_tools(token: str = Depends(oauth2_scheme)) -> SampleShortInfoList: """ @@ -1353,6 +1429,7 @@ async def sample_alignment(token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.state = BeamlineStateEnum.SampleAlignment + @app.post("/state/beam_location") async def beam_location(token: str = Depends(oauth2_scheme)): """ @@ -1364,6 +1441,7 @@ async def beam_location(token: str = Depends(oauth2_scheme)): auth.check_jwt_staff(cfg, auth.parse_token(token)) daq.state = BeamlineStateEnum.BeamLocation + @app.post("/state/beamstop_alignment") async def beamstop_alignment(token: str = Depends(oauth2_scheme)): """ @@ -1379,6 +1457,7 @@ async def beamstop_alignment(token: str = Depends(oauth2_scheme)): daq.state = BeamlineStateEnum.BeamstopAlignment return "OK" + @app.post("/state/flux_measurement") async def flux_measurement(token: str = Depends(oauth2_scheme)): """ @@ -1394,6 +1473,7 @@ async def flux_measurement(token: str = Depends(oauth2_scheme)): daq.state = BeamlineStateEnum.FluxMeasurement return "OK" + @app.post("/state/data_collection") async def data_collection(token: str = Depends(oauth2_scheme)): """ @@ -1440,6 +1520,7 @@ async def xray_fluorescence(token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.state = BeamlineStateEnum.XrayFluorescence + @app.post("/state/xtal_snapshot") async def beam_location(token: str = Depends(oauth2_scheme)): """ @@ -1451,6 +1532,7 @@ async def beam_location(token: str = Depends(oauth2_scheme)): auth.check_jwt_staff(cfg, auth.parse_token(token)) daq.state = BeamlineStateEnum.XtalSnapshot + @app.post("/state/maintenance") async def maintenance(token: str = Depends(oauth2_scheme)) -> str: """ @@ -1467,8 +1549,11 @@ async def maintenance(token: str = Depends(oauth2_scheme)) -> str: logger.warning("Beamline state set to Maintenance via protected endpoint.") return "OK" + @app.post("/access/take_over_beamline") -async def take_over_beamline(payload: RecoveryActionRequest, token: str = Depends(oauth2_scheme)) -> str: +async def take_over_beamline( + payload: RecoveryActionRequest, token: str = Depends(oauth2_scheme) +) -> str: """ Take over the beamline session. Staff only. @@ -1489,8 +1574,11 @@ async def take_over_beamline(payload: RecoveryActionRequest, token: str = Depend ) return "OK" + @app.post("/state/free_beamline") -async def free_beamline(payload: RecoveryActionRequest, token: str = Depends(oauth2_scheme)) -> str: +async def free_beamline( + payload: RecoveryActionRequest, token: str = Depends(oauth2_scheme) +) -> str: """ Clear the beamline busy flag. Staff only. @@ -1511,8 +1599,11 @@ async def free_beamline(payload: RecoveryActionRequest, token: str = Depends(oau ) return "OK" + @app.post("/recovery/recover_beamline") -async def recover_beamline(payload: RecoveryActionRequest, token: str = Depends(oauth2_scheme)) -> dict: +async def recover_beamline( + payload: RecoveryActionRequest, token: str = Depends(oauth2_scheme) +) -> dict: """ Recover the beamline from an error state. Staff only. @@ -1552,8 +1643,12 @@ async def recover_beamline(payload: RecoveryActionRequest, token: str = Depends( "previous_busy": prev_busy, "new_state": BeamlineStateEnum.Maintenance.name, } + + @app.post("/recovery/unmount_sample") -async def recovery_unmount_sample(payload: RecoveryActionRequest, token: str = Depends(oauth2_scheme)) -> dict: +async def recovery_unmount_sample( + payload: RecoveryActionRequest, token: str = Depends(oauth2_scheme) +) -> dict: """ Forcefully unmount a sample during recovery. Staff only. @@ -1610,9 +1705,14 @@ async def recovery_unmount_sample(payload: RecoveryActionRequest, token: str = D "message": "Recovery unmount completed.", } + # Scans @app.post("/scan/raster") -async def raster(val: RasterGridRequest, auto_center: bool = False, token: str = Depends(oauth2_scheme)) -> CompletedRasterGrid: +async def raster( + val: RasterGridRequest, + auto_center: bool = False, + token: str = Depends(oauth2_scheme), +) -> CompletedRasterGrid: """ Execute a raster scan. @@ -1629,7 +1729,9 @@ async def raster(val: RasterGridRequest, auto_center: bool = False, token: str = @app.post("/scan/rotation") -async def rotation(val: RotationScanRequest, token: str = Depends(oauth2_scheme)) -> CompletedRotationScan: +async def rotation( + val: RotationScanRequest, token: str = Depends(oauth2_scheme) +) -> CompletedRotationScan: """ Execute a rotation scan. @@ -1643,6 +1745,7 @@ async def rotation(val: RotationScanRequest, token: str = Depends(oauth2_scheme) auth.check_jwt_rw(cfg, auth.parse_token(token)) return daq.measure_rotation(val) + @app.post("/scan/auto") async def auto(s: SampleShortInfo, token: str = Depends(oauth2_scheme)): """ @@ -1672,7 +1775,9 @@ async def auto(s: SampleShortInfo, token: str = Depends(oauth2_scheme)): @app.post("/scan/smart_params") -async def set_smart_params(p: SimpleScanParameters, token: str = Depends(oauth2_scheme)) -> str: +async def set_smart_params( + p: SimpleScanParameters, token: str = Depends(oauth2_scheme) +) -> str: """ Set the 'smart' scan parameters. @@ -1687,6 +1792,7 @@ async def set_smart_params(p: SimpleScanParameters, token: str = Depends(oauth2_ cfg.auto_params = p return "OK" + @app.post("/scan/cancel") async def cancel(token: str = Depends(oauth2_scheme)): """ @@ -1699,6 +1805,7 @@ async def cancel(token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.cancel() + # ALC routines @app.post("/alc/center_loop") async def alc_center_loop(token: str = Depends(oauth2_scheme)) -> str: @@ -1718,7 +1825,9 @@ async def alc_center_loop(token: str = Depends(oauth2_scheme)) -> str: @app.post("/alc/ml_bounding_box") -async def alc_ml_bounding_box(token: str = Depends(oauth2_scheme)) -> RasterGridRequest | None: +async def alc_ml_bounding_box( + token: str = Depends(oauth2_scheme), +) -> RasterGridRequest | None: """ Request an ML-based bounding box for the sample. @@ -1731,8 +1840,11 @@ async def alc_ml_bounding_box(token: str = Depends(oauth2_scheme)) -> RasterGrid auth.check_jwt_rw(cfg, auth.parse_token(token)) return daq.ml_bounding_box() + @app.post("/face_detection/run") -async def face_detection_run(steps: int, step_size: int, token: str = Depends(oauth2_scheme)) -> dict: +async def face_detection_run( + steps: int, step_size: int, token: str = Depends(oauth2_scheme) +) -> dict: """ Trigger the face detection sequence. @@ -1746,16 +1858,19 @@ async def face_detection_run(steps: int, step_size: int, token: str = Depends(oa """ logger.debug(f"Face detection run: {steps} steps, {step_size} step size") auth.check_jwt_rw(cfg, auth.parse_token(token)) - _push_face_detection_progress({ - "running": True, - "status": "starting", - "samples": [], - "height_fit": {}, - "area_fit": {}, - }) + _push_face_detection_progress( + { + "running": True, + "status": "starting", + "samples": [], + "height_fit": {}, + "area_fit": {}, + } + ) result = daq.face_detection(steps=steps, step_size=step_size) return result + @app.get("/sse/face_detection") async def sse_face_detection(token: str = Depends(oauth2_scheme)): """ @@ -1775,10 +1890,11 @@ async def sse_face_detection(token: str = Depends(oauth2_scheme)): "Cache-Control": "no-cache", "Connection": "keep-alive", "Access-Control-Allow-Origin": "*", - "Access-Control-Allow-Headers": "Cache-Control" - } + "Access-Control-Allow-Headers": "Cache-Control", + }, ) + @app.get("/sse/automation_progress") async def sse_automation_progress(token: str = Depends(oauth2_scheme)): """ @@ -1799,7 +1915,7 @@ async def sse_automation_progress(token: str = Depends(oauth2_scheme)): "Connection": "keep-alive", "Access-Control-Allow-Origin": "*", "Access-Control-Allow-Headers": "Cache-Control", - } + }, ) @@ -1844,6 +1960,7 @@ async def set_pgroup(val: str, token: str = Depends(oauth2_scheme)) -> str: cfg.pgroup = val return "OK" + @app.delete("/access/pgroup") async def del_pgroup(token: str = Depends(oauth2_scheme)) -> str: """ @@ -1859,6 +1976,7 @@ async def del_pgroup(token: str = Depends(oauth2_scheme)) -> str: cfg.pgroup = None return "OK" + @app.put("/beamline/commissioning_mode") async def set_commissioning_mode(val: bool, token: str = Depends(oauth2_scheme)) -> str: """ @@ -1875,6 +1993,7 @@ async def set_commissioning_mode(val: bool, token: str = Depends(oauth2_scheme)) cfg.commissioning_mode = val return "OK" + @app.get("/admin/gui_sessions") async def get_gui_sessions(token: str = Depends(oauth2_scheme)) -> list[dict]: """ @@ -1887,9 +2006,9 @@ async def get_gui_sessions(token: str = Depends(oauth2_scheme)) -> list[dict]: @app.post("/admin/gui_sessions/{session_id}/request_close") async def request_gui_close( - session_id: int, - grace_seconds: int = 60, - token: str = Depends(oauth2_scheme), + session_id: int, + grace_seconds: int = 60, + token: str = Depends(oauth2_scheme), ) -> dict: """ Staff-only request for a remote GUI to close gracefully. @@ -1917,8 +2036,8 @@ async def request_gui_close( @app.delete("/admin/gui_sessions/{session_id}") async def force_remove_gui_session( - session_id: int, - token: str = Depends(oauth2_scheme), + session_id: int, + token: str = Depends(oauth2_scheme), ) -> dict: """ Staff-only hard removal of a GUI session from Redis. @@ -1939,10 +2058,11 @@ async def force_remove_gui_session( "message": "GUI session force removed.", } + @app.post("/admin/gui_sessions/{session_id}/interaction") async def update_gui_interaction( - session_id: int, - token: str = Depends(oauth2_scheme), + session_id: int, + token: str = Depends(oauth2_scheme), ) -> dict: """ GUI-side activity heartbeat. @@ -1955,7 +2075,9 @@ async def update_gui_interaction( detail="Cannot update interaction for another session.", ) - payload = cfg.update_gui_interaction(session=session_id, last_interaction_ts=time.time()) + payload = cfg.update_gui_interaction( + session=session_id, last_interaction_ts=time.time() + ) if payload is None: raise HTTPException( status_code=api_status.HTTP_404_NOT_FOUND, @@ -1964,6 +2086,7 @@ async def update_gui_interaction( return {"ok": True} + @app.post("/access/end_session") async def end_session(token: str = Depends(oauth2_scheme)) -> str: """ @@ -2003,8 +2126,10 @@ 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: """ @@ -2041,6 +2166,7 @@ async def baton_request(token: str = Depends(oauth2_scheme)) -> dict: data = auth.parse_token(token) return auth.request_baton(cfg, data) + @app.post("/baton/respond") async def baton_respond(accept: bool, token: str = Depends(oauth2_scheme)) -> dict: """ @@ -2124,15 +2250,15 @@ async def baton_check_timeout(token: str = Depends(oauth2_scheme)) -> dict: elapsed = time.time() - pending.created_at if elapsed < pending.timeout_seconds: - return { - "pending": True, - "remaining_seconds": pending.timeout_seconds - elapsed - } + return {"pending": True, "remaining_seconds": pending.timeout_seconds - elapsed} return auth.request_baton(cfg, data) + @app.put("/access/allow_non_staff_request_from_staff") -async def set_allow_non_staff_request_from_staff(val: bool, token: str = Depends(oauth2_scheme)) -> str: +async def set_allow_non_staff_request_from_staff( + val: bool, token: str = Depends(oauth2_scheme) +) -> str: """ Set whether non-staff users can request the baton from staff members. Staff only. @@ -2148,6 +2274,7 @@ async def set_allow_non_staff_request_from_staff(val: bool, token: str = Depends 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. @@ -2196,10 +2323,11 @@ async def sse_baton(token: str = Depends(oauth2_scheme)): "Cache-Control": "no-cache", "Connection": "keep-alive", "Access-Control-Allow-Origin": "*", - "Access-Control-Allow-Headers": "Cache-Control" - } + "Access-Control-Allow-Headers": "Cache-Control", + }, ) + @app.get("/beamline/settings") async def get_settings(token: str = Depends(oauth2_scheme)) -> BeamlineSettingsModel: """ @@ -2229,8 +2357,11 @@ async def put_settings(s: BeamlineSettingsModel, token: str = Depends(oauth2_sch auth.check_jwt_staff(cfg, auth.parse_token(token)) cfg.settings = s + @app.get("/beamline/cryo_settings") -async def get_cryo_settings(token: str = Depends(oauth2_scheme)) -> CryojetSettingsModel: +async def get_cryo_settings( + token: str = Depends(oauth2_scheme), +) -> CryojetSettingsModel: """ Get the current cryojet settings. Staff only. @@ -2244,8 +2375,11 @@ async def get_cryo_settings(token: str = Depends(oauth2_scheme)) -> CryojetSetti auth.check_jwt_staff(cfg, auth.parse_token(token)) return cfg.cryojet_settings + @app.put("/beamline/cryo_settings") -async def put_cryo_settings(s: CryojetSettingsModel, token: str = Depends(oauth2_scheme)): +async def put_cryo_settings( + s: CryojetSettingsModel, token: str = Depends(oauth2_scheme) +): """ Update the cryojet settings. Staff only. @@ -2257,6 +2391,7 @@ async def put_cryo_settings(s: CryojetSettingsModel, token: str = Depends(oauth2 auth.check_jwt_staff(cfg, auth.parse_token(token)) cfg.cryojet_settings = s + @app.put("/access/all_pgroups") async def get_all_pgroups(token: str = Depends(oauth2_scheme)): """ @@ -2273,7 +2408,7 @@ async def get_all_pgroups(token: str = Depends(oauth2_scheme)): base_path = "/sls/mx/data/" - now=time.monotonic() + now = time.monotonic() cached = _all_pgroups_cache.get(base_path) if cached is not None: items, ts = cached @@ -2288,10 +2423,10 @@ async def get_all_pgroups(token: str = Depends(oauth2_scheme)): name = entry.name # quick name filter: p + 5 digits, and directory if ( - len(name) == 6 and - name[0] == "p" and - name[1:].isdigit() and - entry.is_dir(follow_symlinks=False) + len(name) == 6 + and name[0] == "p" + and name[1:].isdigit() + and entry.is_dir(follow_symlinks=False) ): names.append(name) names.sort(key=lambda d: int(d[1:])) @@ -2301,8 +2436,11 @@ async def get_all_pgroups(token: str = Depends(oauth2_scheme)): logger.error(f"Failed to list pgroups: {e}") return [] + @app.post("/fluorimeter/spectrum") -async def fluorimeter_spectrum(input: FluorescenceSpectrumParameterModel, token: str = Depends(oauth2_scheme)) -> FluorescenceSpectrumOutputModel: +async def fluorimeter_spectrum( + input: FluorescenceSpectrumParameterModel, token: str = Depends(oauth2_scheme) +) -> FluorescenceSpectrumOutputModel: """ Request a fluorescence spectrum measurement. @@ -2316,8 +2454,11 @@ async def fluorimeter_spectrum(input: FluorescenceSpectrumParameterModel, token: auth.check_jwt_rw(cfg, auth.parse_token(token)) return daq.fluorimeter_take_spectrum(input) + @app.post("/fluorimeter/start") -async def fluorimeter_start(erase: bool = False, token: str = Depends(oauth2_scheme)) -> str: +async def fluorimeter_start( + erase: bool = False, token: str = Depends(oauth2_scheme) +) -> str: """ Start the fluorimeter measurement. @@ -2332,6 +2473,7 @@ async def fluorimeter_start(erase: bool = False, token: str = Depends(oauth2_sch daq.fluorimeter_start(erase) return "OK" + @app.post("/fluorimeter/stop") async def fluorimeter_stop(token: str = Depends(oauth2_scheme)) -> str: """ @@ -2347,6 +2489,7 @@ async def fluorimeter_stop(token: str = Depends(oauth2_scheme)) -> str: daq.fluorimeter_stop() return "OK" + @app.get("/fluorimeter/status") async def fluorimeter_status(token: str = Depends(oauth2_scheme)) -> int | None: """ @@ -2361,6 +2504,7 @@ async def fluorimeter_status(token: str = Depends(oauth2_scheme)) -> int | None: auth.check_jwt_ro(cfg, auth.parse_token(token)) return daq.fluorimeter_status() + @app.get("/fluorimeter/data") async def fluorimeter_data(token: str = Depends(oauth2_scheme)) -> list[int] | None: """ @@ -2375,8 +2519,11 @@ async def fluorimeter_data(token: str = Depends(oauth2_scheme)) -> list[int] | N auth.check_jwt_ro(cfg, auth.parse_token(token)) return daq.fluorimeter_data() + @app.get("/fluorimeter/background") -async def fluorimeter_background(token: str = Depends(oauth2_scheme)) -> list[int] | None: +async def fluorimeter_background( + token: str = Depends(oauth2_scheme), +) -> list[int] | None: """ Get the fluorimeter background data. @@ -2389,6 +2536,7 @@ async def fluorimeter_background(token: str = Depends(oauth2_scheme)) -> list[in auth.check_jwt_ro(cfg, auth.parse_token(token)) return daq.fluorimeter_background() + # Optional: SSE stream for live spectra while acquiring async def fluorimeter_stream() -> AsyncGenerator[str, None]: """ @@ -2405,11 +2553,33 @@ async def fluorimeter_stream() -> AsyncGenerator[str, None]: b = daq.fluorimeter_background() # Ensure JSON-serializable lists; avoid numpy truthiness status = "stopped" if s == 0 else "running" - logger.debug(f"fluorimeter status: {status}. got data from fluorimeter: {d[0]} and background: {b[0]}") - data = d.tolist() if hasattr(d, "tolist") else (list(d) if isinstance(d, (tuple, list)) else ([] if d is None else [d])) - bkg = b.tolist() if hasattr(b, "tolist") else (list(b) if isinstance(b, (tuple, list)) else ([] if b is None else [b])) - logger.debug(f"fluorimeter data: {data[0]}. fluorimeter background: {bkg[0]}") - payload = json.dumps({"status": s, "data": data, "background": bkg}, separators=(",", ":")) + logger.debug( + f"fluorimeter status: {status}. got data from fluorimeter: {d[0]} and background: {b[0]}" + ) + data = ( + d.tolist() + if hasattr(d, "tolist") + else ( + list(d) + if isinstance(d, (tuple, list)) + else ([] if d is None else [d]) + ) + ) + bkg = ( + b.tolist() + if hasattr(b, "tolist") + else ( + list(b) + if isinstance(b, (tuple, list)) + else ([] if b is None else [b]) + ) + ) + logger.debug( + f"fluorimeter data: {data[0]}. fluorimeter background: {bkg[0]}" + ) + payload = json.dumps( + {"status": s, "data": data, "background": bkg}, separators=(",", ":") + ) yield f"data: {payload}\n\n" if s == 0: await asyncio.sleep(0.2) @@ -2420,6 +2590,7 @@ async def fluorimeter_stream() -> AsyncGenerator[str, None]: logger.error(f">>> Fluorimeter stream cancelled: {e}") return + @app.get("/sse/fluorimeter") async def sse_fluorimeter(token: str = Depends(oauth2_scheme)): """ @@ -2444,15 +2615,16 @@ async def sse_fluorimeter(token: str = Depends(oauth2_scheme)): "Cache-Control": "no-cache", "Connection": "keep-alive", "Access-Control-Allow-Origin": "*", - "Access-Control-Allow-Headers": "Cache-Control" - } + "Access-Control-Allow-Headers": "Cache-Control", + }, ) + @app.post("/samcam/send_screenshot_db") async def send_screenshot_db( - filename: str | None = None, - message: str | None = None, - token: str = Depends(oauth2_scheme), + filename: str | None = None, + message: str | None = None, + token: str = Depends(oauth2_scheme), ) -> str: """ Capture a screenshot and send it to the database. @@ -2470,6 +2642,7 @@ async def send_screenshot_db( daq.send_screenshot_db(filename=filename, message=message) return "OK" + async def send_message_db( db_id: int, event_type: SampleEventType, @@ -2481,6 +2654,7 @@ async def send_message_db( daq.send_screenshot_db(db_id=db_id, event_type=event_type, comment=comment) return "OK" + @app.get("/camera/source") async def get_camera_source(token: str = Depends(oauth2_scheme)): """Get the current camera image source.""" @@ -2495,6 +2669,7 @@ async def set_camera_source(use_zmq: bool, token: str = Depends(oauth2_scheme)): daq._AareDAQ__devs.set_camera_source(use_zmq) return {"source": daq._AareDAQ__devs.samcam_source} + @app.post("/state/maintenance") async def maintenance(token: str = Depends(oauth2_scheme)) -> str: """ @@ -2509,18 +2684,27 @@ async def maintenance(token: str = Depends(oauth2_scheme)) -> str: auth.check_jwt_staff(cfg, auth.parse_token(token)) cfg.state = BeamlineStateEnum.Maintenance logger.warning("Beamline state set to Maintenance via protected endpoint.") - raise MaintenanceStateException("Beamline was set to Maintenance. Automation must stop.") + raise MaintenanceStateException( + "Beamline was set to Maintenance. Automation must stop." + ) def main(): # Remove in production! - #urllib3.disable_warnings() + # urllib3.disable_warnings() # Run the application using uvicorn - uvicorn.run("aare.daq.server:app", host="127.0.0.1", port=5210, workers=4, proxy_headers=False, log_config=get_uvicorn_logging_config()) + uvicorn.run( + "aare.daq.server:app", + host="127.0.0.1", + port=5210, + workers=4, + proxy_headers=False, + log_config=get_uvicorn_logging_config(), + ) if __name__ == "__main__": - #start_image_stats_receiver(zmq_url="tcp://129.129.110.12:9089") + # start_image_stats_receiver(zmq_url="tcp://129.129.110.12:9089") main() - #stop_image_stats_receiver() \ No newline at end of file + # stop_image_stats_receiver() diff --git a/src/aare/daq/server_exception_handler.py b/src/aare/daq/server_exception_handler.py index 39a48d07..bcaef31f 100644 --- a/src/aare/daq/server_exception_handler.py +++ b/src/aare/daq/server_exception_handler.py @@ -23,40 +23,42 @@ All handlers produce the unified response body defined in §3: "context": {"endpoint": "/state", "operation": "GET"} } """ + from __future__ import annotations import json from typing import Any +from aarecommon.config.logger import setup_logger +from aarecommon.errors.codes import AareErrorCode, code_for_exception_class +from aarecommon.errors.exception_handler import ( + AareAuthError, + AareException, + AareUserError, + AuthenticationException, + AutomationError, +) from fastapi import HTTPException from fastapi import status as api_status from starlette.requests import Request from starlette.responses import JSONResponse -from aare.common.logger_config import setup_logger -from aare.common.error_codes import AareErrorCode, code_for_exception_class -from aare.common.exception_handler import ( - AareException, - AutomationError, - AareUserError, - AareAuthError, - AuthenticationException, -) - logger = setup_logger("aareDAQ") # Attributes that live on exceptions for routing/transport, NOT for the # context payload. -_NON_CONTEXT_ATTRS = frozenset({ - "args", - "message", - "critical", - "code", - "headers", - "status_code", - "exception", # BECCommunicationError stores a Python exception here -- not JSON-safe -}) +_NON_CONTEXT_ATTRS = frozenset( + { + "args", + "message", + "critical", + "code", + "headers", + "status_code", + "exception", # BECCommunicationError stores a Python exception here -- not JSON-safe + } +) def _extract_context(exc: Exception) -> dict[str, Any]: @@ -83,8 +85,12 @@ def _extract_context(exc: Exception) -> dict[str, Any]: return ctx -def _error_body(exc: Exception, *, code_override: str | None = None, - critical_override: bool | None = None) -> dict[str, Any]: +def _error_body( + exc: Exception, + *, + code_override: str | None = None, + critical_override: bool | None = None, +) -> dict[str, Any]: """Build the unified error response body for any exception. ``code_override`` is used by the auth handler to surface the finer-grained @@ -119,7 +125,9 @@ def register_exception_handlers(app) -> None: """ @app.exception_handler(AutomationError) - async def _handle_automation_error(request: Request, exc: AutomationError) -> JSONResponse: + async def _handle_automation_error( + request: Request, exc: AutomationError + ) -> JSONResponse: body = _error_body(exc) log = logger.error if exc.critical else logger.warning log( @@ -188,7 +196,9 @@ def register_exception_handlers(app) -> None: ) @app.exception_handler(HTTPException) - async def _handle_http_exception(request: Request, exc: HTTPException) -> JSONResponse: + async def _handle_http_exception( + request: Request, exc: HTTPException + ) -> JSONResponse: """HTTPException isn't in the AareException hierarchy but FastAPI raises it internally (e.g. validation, route-not-found) and routes raise it manually. Adapt to the new body shape so clients see one schema.""" @@ -210,7 +220,9 @@ def register_exception_handlers(app) -> None: "code": legacy_code, "exception_class": "HTTPException", "message": legacy_message, - "context": {k: v for k, v in detail.items() if k not in ("code", "message")}, + "context": { + k: v for k, v in detail.items() if k not in ("code", "message") + }, }, headers=exc.headers, ) diff --git a/src/aare/daq/spreadsheetupdater.py b/src/aare/daq/spreadsheetupdater.py index 8f27c407..dda5acc5 100644 --- a/src/aare/daq/spreadsheetupdater.py +++ b/src/aare/daq/spreadsheetupdater.py @@ -1,28 +1,35 @@ -import os import json -import websocket +import os import time + +from aarecommon.config.beamline import mx_beamline +from aarecommon.models.models import DewarAddress, SampleShortInfo, SampleShortInfoList from aareDB.models import PuckWithTellPosition -from aare.common.models import SampleShortInfoList, SampleShortInfo, DewarAddress -from aare.common.beamline import mx_beamline + from aare.daq.config import BeamlineConfig beamline = mx_beamline() config = BeamlineConfig(bl=beamline) + def get_ws_headers() -> list[str]: password = os.getenv("AAREDB_SHARED_PASSWORD") if not password: raise ValueError("The AAREDB_SHARED_PASSWORD environment variable is not set.") return [f"X-Shared-Password: {password}"] + def set_spreadsheet_in_redis(spreadsheet): """ Store the spreadsheet in Redis using your BeamlineConfig instance. """ - print('[REDIS][DEBUG] Data to write:', json.dumps(spreadsheet, indent=4)) # Pretty-print the data - print('[REDIS][INFO] Writing spreadsheet to Redis...') - config.__client.set(f"{config._BeamlineConfig__bl}:spreadsheet", json.dumps(spreadsheet)) + print( + "[REDIS][DEBUG] Data to write:", json.dumps(spreadsheet, indent=4) + ) # Pretty-print the data + print("[REDIS][INFO] Writing spreadsheet to Redis...") + config.__client.set( + f"{config._BeamlineConfig__bl}:spreadsheet", json.dumps(spreadsheet) + ) def on_message(ws, message): @@ -32,7 +39,9 @@ def on_message(ws, message): try: data = json.loads(message) - pucks_data = data["samples"] if isinstance(data, dict) and "samples" in data else data + pucks_data = ( + data["samples"] if isinstance(data, dict) and "samples" in data else data + ) pucks = [PuckWithTellPosition(**item) for item in pucks_data] print(f"INFO: Received pucks: {[p.puck_name for p in pucks]}") @@ -46,7 +55,10 @@ def on_message(ws, message): # Reference tool if isinstance(p.tell_position, str) and p.tell_position.startswith("X"): - dewar_address = DewarAddress(segment="X", pos=int(p.tell_position[1:]) if len(p.tell_position) > 1 else 1) + dewar_address = DewarAddress( + segment="X", + pos=int(p.tell_position[1:]) if len(p.tell_position) > 1 else 1, + ) target_list = reference_short_infos # Normal puck like "A1", "B5", "C3" @@ -57,7 +69,9 @@ def on_message(ws, message): target_list = normal_short_infos else: - print(f"[WARN] Skipping unknown tell_position format: {p.tell_position}") + print( + f"[WARN] Skipping unknown tell_position format: {p.tell_position}" + ) continue info = SampleShortInfo( @@ -69,13 +83,13 @@ def on_message(ws, message): user=s.pgroup, pin=s.position, location=dewar_address, - priority=s.priority,#getattr(s, "priority", 1.0), + priority=s.priority, # getattr(s, "priority", 1.0), comment=getattr(s, "comments", None), mount_count=s.mount_count or 0, screening_count=s.screening_count or 0, raster_count=s.raster_count or 0, rotation_count=s.rotation_count or 0, - aaredb_params=s.data_collection_parameters + aaredb_params=s.data_collection_parameters, ) target_list.append(info) @@ -103,6 +117,7 @@ def on_message(ws, message): except Exception as exc: print("[WS][ERROR] Failed to parse or convert message:", exc) + def on_error(ws, error): """ Handle WebSocket errors. @@ -136,15 +151,18 @@ def main(): while True: try: import ssl + import websocket # FORCE a clean context context = ssl.create_default_context(ssl.Purpose.SERVER_AUTH) - context.load_verify_locations(cafile="/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem") + context.load_verify_locations( + cafile="/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem" + ) beamline_name = beamline.value.lower() context.load_cert_chain( certfile=f"/etc/ssl/certs/secrets/mx-{beamline_name}-queue-01_from_dmz-01.crt", - keyfile=f"/etc/ssl/certs/secrets/mx-{beamline_name}-queue-01_from_dmz-01.key" + keyfile=f"/etc/ssl/certs/secrets/mx-{beamline_name}-queue-01_from_dmz-01.key", ) # Explicitly set the SNI hostname to match NGINX server_name @@ -152,7 +170,7 @@ def main(): ssl_opt = { "context": context, "server_hostname": "mx-aaredb-dmz-01.psi.ch", - "check_hostname": True + "check_hostname": True, } ws = websocket.WebSocketApp( @@ -172,5 +190,6 @@ def main(): time.sleep(5) + if __name__ == "__main__": - main() \ No newline at end of file + main() diff --git a/src/aare/daq/tell_state_machine.py b/src/aare/daq/tell_state_machine.py index 8ed298c8..eb1eec90 100644 --- a/src/aare/daq/tell_state_machine.py +++ b/src/aare/daq/tell_state_machine.py @@ -3,7 +3,7 @@ from __future__ import annotations import re from datetime import datetime, timezone -from aare.common.tell_models import TellActivityEnum, TellPhaseEnum, TellStateModel +from aarecommon.models.tell import TellActivityEnum, TellPhaseEnum, TellStateModel MOUNT_CMD_RE = re.compile(r'^mount\("([^"]+)",\s*(\d+),\s*(\d+)') UNMOUNT_CMD_RE = re.compile(r"^unmount\(") @@ -72,7 +72,9 @@ def _update_state( return state.model_copy(update=updates) -def _ready_state(state: TellStateModel, event_name: str, event_value: str | None) -> TellStateModel: +def _ready_state( + state: TellStateModel, event_name: str, event_value: str | None +) -> TellStateModel: return _update_state( state, activity=TellActivityEnum.IDLE, @@ -129,7 +131,9 @@ def advance_tell_state( return _update_state( state, activity=TellActivityEnum.MOUNTING, - message=f"Mounting {sample_position}" if sample_position else "Mounting sample", + message=f"Mounting {sample_position}" + if sample_position + else "Mounting sample", event_name=event_name, event_value=value, operation="mount", @@ -154,7 +158,11 @@ def advance_tell_state( if UNMOUNT_STATUS_RE.match(value): next_operation = "mount" if state.operation == "mount" else "unmount" - next_phase = TellPhaseEnum.AUTO_UNMOUNT if state.operation == "mount" else TellPhaseEnum.PREPARING + next_phase = ( + TellPhaseEnum.AUTO_UNMOUNT + if state.operation == "mount" + else TellPhaseEnum.PREPARING + ) return _update_state( state, @@ -283,7 +291,11 @@ def advance_tell_state( else TellPhaseEnum.COMPLETE ) next_operation = "mount" if state.operation == "mount" else "unmount" - next_message = "Returning sample to puck" if state.operation == "mount" else "Sample returned to puck" + next_message = ( + "Returning sample to puck" + if state.operation == "mount" + else "Sample returned to puck" + ) return _update_state( state, activity=TellActivityEnum.UNMOUNTING, @@ -292,7 +304,9 @@ def advance_tell_state( event_value=value, operation=next_operation, phase=next_phase, - mount_success=True if state.operation == "unmount" else state.mount_success, + mount_success=True + if state.operation == "unmount" + else state.mount_success, mount_error="" if state.operation == "unmount" else state.mount_error, ) @@ -352,7 +366,9 @@ def advance_tell_state( if value == "Busy": return _update_state( state, - message=state.message if state.operation not in {None, "idle"} else "Busy", + message=state.message + if state.operation not in {None, "idle"} + else "Busy", event_name=event_name, event_value=value, ) @@ -361,4 +377,4 @@ def advance_tell_state( state, event_name=event_name, event_value=value, - ) \ No newline at end of file + ) diff --git a/src/aare/daq/tellupdater.py b/src/aare/daq/tellupdater.py index 7ce91a3f..264aef7d 100644 --- a/src/aare/daq/tellupdater.py +++ b/src/aare/daq/tellupdater.py @@ -1,29 +1,32 @@ -import os import json +import os +import ssl import threading +import time from collections import deque from datetime import datetime, timezone from typing import Any, cast import requests -import websocket import sseclient -import time -import ssl -from aareDB.models import PuckWithTellPosition +import websocket +from aarecommon.config.beamline import mx_beamline +from aarecommon.config.logger import setup_logger +from aarecommon.errors.exception_handler import ( + TellCommunicationError, + TellConnectionException, +) +from aarecommon.models.beamline import MXBeamline +from aarecommon.models.tell import TellStateModel from aareDB.exceptions import ApiException +from aareDB.models import PuckWithTellPosition from pydantic import BaseModel, ValidationError from requests import RequestException -from aare.devices.tell_client import TellClient - -from aare.common.exception_handler import TellCommunicationError, TellConnectionException -from aare.common.logger_config import setup_logger -from aare.common.tell_models import TellStateModel from aare.daq.aaredb import AareWrapper from aare.daq.config import BeamlineConfig from aare.daq.tell_state_machine import advance_tell_state, initial_tell_state -from aare.common.beamline import MXBeamline, mx_beamline +from aare.devices.tell_client import TellClient logger = setup_logger(__name__) @@ -47,9 +50,7 @@ TRACKED_TELL_EVENT_VALUES = { } TRACKED_MOTION_EVENT_VALUES = { - "Motion Task": { - "dry" - }, + "Motion Task": {"dry"}, } TRACKED_MOTION_SYNC_EVENTS = { "Motion Sync": { @@ -61,12 +62,7 @@ TRACKED_MOTION_SYNC_EVENTS = { "Sample get on Puck", }, } -TRACKED_STATE_EVENTS = { - "state": { - "Ready", - "Busy" - } -} +TRACKED_STATE_EVENTS = {"state": {"Ready", "Busy"}} latest_tell_events = {} tell_event_history = deque(maxlen=25) @@ -132,7 +128,9 @@ def _get_redis_context() -> tuple[Any | None, str | None]: redis_client = getattr(config, "_BeamlineConfig__client", None) beamline_key = getattr(config, "_BeamlineConfig__bl", None) if redis_client is None or beamline_key is None: - logger.error("[REDIS] BeamlineConfig internals unavailable; skipping TELL redis write") + logger.error( + "[REDIS] BeamlineConfig internals unavailable; skipping TELL redis write" + ) return None, None return cast(Any, redis_client), str(beamline_key) @@ -238,6 +236,7 @@ def get_dmz_client_cert_paths(bl: MXBeamline) -> tuple[str, str]: f"{cert_dir}/mx-{beamline}-queue-01_from_dmz-01.key", ) + def listen_to_sse(): logger.info("[SSE] listener started") if not tell_client.url: @@ -249,11 +248,13 @@ def listen_to_sse(): while True: try: logger.info(f"[SSE] Connecting to {sse_url}") - if sse_url.startswith("https://mx-aaredb-dmz-01"): #"https://mx-db-01" + if sse_url.startswith("https://mx-aaredb-dmz-01"): # "https://mx-db-01" # mTLS path - #ca_root = "/etc/ssl/certs/secrets/mx-db-01_Full_Chain_CA.pem" + # 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 = requests.get( + sse_url, stream=True, cert=cert_pair, verify=ca_root + ) response.raise_for_status() client = sseclient.SSEClient(response) else: @@ -270,7 +271,9 @@ def listen_to_sse(): # print(f"event = {event.event} with data: {event.data}") on_sse_event(event) except SSE_RECOVERABLE_EXCEPTIONS: - logger.exception("[SSE] Recoverable failure while consuming TELL event stream") + logger.exception( + "[SSE] Recoverable failure while consuming TELL event stream" + ) logger.info(f"[SSE] Reconnecting in {SSE_RECONNECT_DELAY_S} seconds") time.sleep(SSE_RECONNECT_DELAY_S) @@ -298,11 +301,13 @@ def compare_and_report_change(old, new): return joined, left + def ws_update_samples_info(pucks): """Send sample info to TELL robot.""" tell_client.set_samples_info(pucks) logger.info(f"Payload sent to TELL: {pucks}") + def handle_tell_change_event(): """Fetch the latest detected pucks from TELL and update the database.""" try: @@ -314,6 +319,7 @@ def handle_tell_change_event(): except (ApiException, json.JSONDecodeError, KeyError, TypeError, ValueError): logger.exception("[SSE] Failed to update puck state after TELL change") + def on_sse_event(event): global current_tell_state @@ -334,6 +340,7 @@ def on_sse_event(event): logger.info("[SSE] Dewar content changed; updating") handle_tell_change_event() + def on_message(ws, message): try: data = json.loads(message) @@ -350,12 +357,15 @@ def on_message(ws, message): except Exception: logger.exception("[WS] Failed to process message") + def on_error(ws, error): logger.error(f"[WS] Error: {error}") + def on_close(ws, close_status_code, close_msg): logger.info(f"[WS] Closed status={close_status_code} message={close_msg}") + def on_open(ws): logger.info("[WS] WebSocket opened") global current_pucks @@ -378,13 +388,12 @@ def main(): context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) # 2. Load the CA to verify the NGINX server's identity - context.load_verify_locations(cafile="/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem") + context.load_verify_locations( + cafile="/etc/ssl/certs/secrets/mx-aaredb-dmz-01_Full_Chain_CA.pem" + ) # 3. Load the Client Certificate and Key (mTLS) - context.load_cert_chain( - certfile=cert_file, - keyfile=key_file - ) + context.load_cert_chain(certfile=cert_file, keyfile=key_file) # Optional: Ensure hostname matching is active (recommended) context.check_hostname = True @@ -408,6 +417,7 @@ def main(): logger.info(f"[WS] Reconnecting in {SSE_RECONNECT_DELAY_S} seconds") time.sleep(SSE_RECONNECT_DELAY_S) + if __name__ == "__main__": # Configuration beamline = mx_beamline() diff --git a/src/aare/daq/workflows.py b/src/aare/daq/workflows.py index 1bb8e409..485b3785 100644 --- a/src/aare/daq/workflows.py +++ b/src/aare/daq/workflows.py @@ -1,19 +1,21 @@ -from aare.common.models import SampleCameraSettings -from aare.common.logger_config import setup_logger +import time +from aarecommon.config.logger import setup_logger +from aarecommon.models.models import SampleCameraSettings + +from aare.daq.config import ABR_OMEGA_MOUNT, ABR_POS_MOUNT, BeamlineConfig +from aare.daq.devices import BeamlineDevices from aare.devices.area_detector import AutoEnum from aare.devices.bec_worker import BeamlineState -from aare.daq.devices import BeamlineDevices -from aare.daq.config import BeamlineConfig, ABR_POS_MOUNT, ABR_OMEGA_MOUNT - -import time logger = setup_logger("aareDAQ") SAFE_POSITION = 600 -def _move_detector_to_safe_position_if_needed(devs: BeamlineDevices, cfg: BeamlineConfig): +def _move_detector_to_safe_position_if_needed( + devs: BeamlineDevices, cfg: BeamlineConfig +): safe_position = cfg.dtz_safe_position current_dtz = devs.dtz @@ -39,17 +41,19 @@ def _move_detector_to_safe_position_if_needed(devs: BeamlineDevices, cfg: Beamli def common_2rse(devs: BeamlineDevices, cfg: BeamlineConfig): _move_detector_to_safe_position_if_needed(devs, cfg) devs.bec_worker.move_to(BeamlineState.ROBOT_SAMPLE_EXCHANGE) - print('move smargon home') + print("move smargon home") devs.smargon_move_home() - print('try to move aerotech') + print("try to move aerotech") devs.aerotech_pos = ABR_POS_MOUNT devs.aerotech_omega = ABR_OMEGA_MOUNT + def rse2se(devs: BeamlineDevices, cfg: BeamlineConfig): devs.samcam_auto(AutoEnum.AUTO) common_2rse(devs, cfg) devs.samcam_auto(AutoEnum.ONCE) + def m2se(devs: BeamlineDevices, cfg: BeamlineConfig): devs.samcam_auto(AutoEnum.AUTO) devs.bec_worker.move_to(BeamlineState.MANUAL_SAMPLE_EXCHANGE) @@ -59,11 +63,13 @@ def m2se(devs: BeamlineDevices, cfg: BeamlineConfig): devs.smargon_move_home() devs.samcam_auto(AutoEnum.ONCE) + def m2sa(devs: BeamlineDevices, cfg: BeamlineConfig): devs.samcam_auto(AutoEnum.AUTO) devs.bec_worker.move_to(BeamlineState.SAMPLE_ALIGNMENT) devs.samcam_auto(AutoEnum.ONCE) + def sa2se(devs: BeamlineDevices, cfg: BeamlineConfig): try: devs.samcam_auto(AutoEnum.AUTO) @@ -80,64 +86,80 @@ def sa2se(devs: BeamlineDevices, cfg: BeamlineConfig): devs.smargon_move_home() raise + def sa2rse(devs: BeamlineDevices, cfg: BeamlineConfig): common_2rse(devs, cfg) -def sa2xtal_snapshot(devs:BeamlineDevices, cfg:BeamlineConfig): + +def sa2xtal_snapshot(devs: BeamlineDevices, cfg: BeamlineConfig): logger.info("SA to Xtal Snapshot") devs.samcam_settings = SampleCameraSettings(gain=0.0, exposure=0.001) logger.debug(f"new sam cam settings: {devs.samcam_settings}") devs.bec_worker.move_to(BeamlineState.XTAL_SNAPSHOT) -def dc2xtal_snapshot(devs:BeamlineDevices, cfg:BeamlineConfig): + +def dc2xtal_snapshot(devs: BeamlineDevices, cfg: BeamlineConfig): logger.debug("DC to Xtal Snapshot") devs.samcam_settings = SampleCameraSettings(gain=0.0, exposure=0.001) logger.debug(f"new sam cam settings: {devs.samcam_settings}") devs.bec_worker.move_to(BeamlineState.XTAL_SNAPSHOT) -def xtal_snapshot2dc(devs:BeamlineDevices, cfg:BeamlineConfig): + +def xtal_snapshot2dc(devs: BeamlineDevices, cfg: BeamlineConfig): devs.samcam_auto(AutoEnum.AUTO) devs.bec_worker.move_to(BeamlineState.DATA_COLLECTION) devs.samcam_auto(AutoEnum.ONCE) -def xtal_snapshot2sa(devs:BeamlineDevices, cfg:BeamlineConfig): + +def xtal_snapshot2sa(devs: BeamlineDevices, cfg: BeamlineConfig): devs.samcam_auto(AutoEnum.AUTO) devs.bec_worker.move_to(BeamlineState.SAMPLE_ALIGNMENT) devs.samcam_auto(AutoEnum.ONCE) -def xtal_snapshot2rse(devs:BeamlineDevices, cfg:BeamlineConfig): + +def xtal_snapshot2rse(devs: BeamlineDevices, cfg: BeamlineConfig): common_2rse(devs, cfg) -def xtal_snapshot2se(devs:BeamlineDevices, cfg:BeamlineConfig): + +def xtal_snapshot2se(devs: BeamlineDevices, cfg: BeamlineConfig): sa2se(devs, cfg) -def xtal_snapshot2xrf(devs:BeamlineDevices, cfg:BeamlineConfig): + +def xtal_snapshot2xrf(devs: BeamlineDevices, cfg: BeamlineConfig): devs.bec_worker.move_to(BeamlineState.FLUX_MEASUREMENT) + def xtal_snapshot2bl(devs: BeamlineDevices, cfg: BeamlineConfig): devs.bec_worker.move_to(BeamlineState.BEAM_VISUALISATION) + def xtal_snapshot2dh(devs: BeamlineDevices, cfg: BeamlineConfig): - xtal_snapshot2sa(devs,cfg) - sa2dh(devs,cfg) + xtal_snapshot2sa(devs, cfg) + sa2dh(devs, cfg) + def dc2rse(devs: BeamlineDevices, cfg: BeamlineConfig): common_2rse(devs, cfg) + def se2sa(devs: BeamlineDevices, cfg: BeamlineConfig): devs.samcam_auto(AutoEnum.AUTO) devs.bec_worker.move_to(BeamlineState.SAMPLE_ALIGNMENT) devs.aerotech_pos = cfg.abr_meas_pos devs.samcam_auto(AutoEnum.ONCE) + def rse2sa(devs: BeamlineDevices, cfg: BeamlineConfig): se2sa(devs, cfg) + def sa2dc(devs: BeamlineDevices, cfg: BeamlineConfig): start = time.perf_counter() det_z_pos = cfg.dtz status = devs.bec_worker.det_z(value=det_z_pos) - logger.info(f"start detector move to {cfg.dtz} at {time.perf_counter() - start:.2f}") + logger.info( + f"start detector move to {cfg.dtz} at {time.perf_counter() - start:.2f}" + ) devs.samcam_auto(AutoEnum.AUTO) logger.info(f"sam cam auto exposure in {time.perf_counter() - start:.2f}") devs.bec_worker.move_to(BeamlineState.DATA_COLLECTION) @@ -147,6 +169,7 @@ def sa2dc(devs: BeamlineDevices, cfg: BeamlineConfig): devs.samcam_auto(AutoEnum.ONCE) logger.info(f"sam cam auto exposure in {time.perf_counter() - start:.2f}") + def dc2sa(devs: BeamlineDevices, cfg: BeamlineConfig): status = _move_detector_to_safe_position_if_needed(devs, cfg) devs.samcam_auto(AutoEnum.AUTO) @@ -156,75 +179,94 @@ def dc2sa(devs: BeamlineDevices, cfg: BeamlineConfig): status.wait(timeout=60) devs.samcam_auto(AutoEnum.ONCE) + def sa2xrf(devs: BeamlineDevices, cfg: BeamlineConfig): """sample alignment to XrfCollection""" devs.samcam_auto(AutoEnum.AUTO) devs.bec_worker.move_to(BeamlineState.FLUX_MEASUREMENT) devs.samcam_auto(AutoEnum.ONCE) + def xrf2sa(devs: BeamlineDevices, cfg: BeamlineConfig): """XrfCollection to sample alignment""" devs.samcam_auto(AutoEnum.AUTO) devs.bec_worker.move_to(BeamlineState.SAMPLE_ALIGNMENT) devs.samcam_auto(AutoEnum.ONCE) + def common2flux_measurement(devs: BeamlineDevices, cfg: BeamlineConfig): devs.bec_worker.move_to(BeamlineState.FLUX_MEASUREMENT) + def sa2flux_measurement(devs: BeamlineDevices, cfg: BeamlineConfig): devs.bec_worker.move_to(BeamlineState.FLUX_MEASUREMENT) + def flux_measurement2sa(devs: BeamlineDevices, cfg: BeamlineConfig): devs.bec_worker.move_to(BeamlineState.SAMPLE_ALIGNMENT) + def flux_measurement2ba(devs: BeamlineDevices, cfg: BeamlineConfig): devs.bec_worker.move_to(BeamlineState.BEAMSTOP_ALIGNMENT) + def flux_measurement2bl(devs: BeamlineDevices, cfg: BeamlineConfig): devs.bec_worker.move_to(BeamlineState.BEAM_VISUALISATION) + def ba2flux_measurement(devs: BeamlineDevices, cfg: BeamlineConfig): devs.bec_worker.move_to(BeamlineState.FLUX_MEASUREMENT) + def bl2flux_measurement(devs: BeamlineDevices, cfg: BeamlineConfig): devs.bec_worker.move_to(BeamlineState.FLUX_MEASUREMENT) + def sa2ba(devs: BeamlineDevices, cfg: BeamlineConfig): """Sample alignment to BeamstopAlignment""" devs.bec_worker.move_to(BeamlineState.BEAMSTOP_ALIGNMENT) devs.samcam_auto(AutoEnum.ONCE) + def ba2sa(devs: BeamlineDevices, cfg: BeamlineConfig): devs.bec_worker.move_to(BeamlineState.SAMPLE_ALIGNMENT) devs.samcam_auto(AutoEnum.ONCE) + def sa2bl(devs: BeamlineDevices, cfg: BeamlineConfig): devs.bec_worker.move_to(BeamlineState.BEAM_VISUALISATION) devs.samcam_settings.exposure = 0.001 - #TODO SAMCAM SETTINGS SHOULD BE DOEN via zoom settings + # TODO SAMCAM SETTINGS SHOULD BE DOEN via zoom settings + def bl2sa(devs: BeamlineDevices, cfg: BeamlineConfig): devs.bec_worker.move_to(BeamlineState.SAMPLE_ALIGNMENT) devs.samcam_auto(AutoEnum.ONCE) + def bl2ba(devs: BeamlineDevices, cfg: BeamlineConfig): devs.bec_worker.move_to(BeamlineState.BEAMSTOP_ALIGNMENT) devs.samcam_auto(AutoEnum.ONCE) + def ba2bl(devs: BeamlineDevices, cfg: BeamlineConfig): devs.bec_worker.move_to(BeamlineState.BEAM_VISUALISATION) devs.samcam_settings.exposure = 0.001 + def sa2dh(devs: BeamlineDevices, cfg: BeamlineConfig): common2dh(devs, cfg) + def common2dh(devs: BeamlineDevices, cfg: BeamlineConfig): devs.samcam_auto(AutoEnum.AUTO) common_2rse(devs, cfg) devs.samcam_auto(AutoEnum.ONCE) if devs.tell.is_position("pPark") and devs.tell.get_mounted_sample() is None: - logger.info("TELL already in pPark; skipping dry/park preparation for dewar transfer") + logger.info( + "TELL already in pPark; skipping dry/park preparation for dewar transfer" + ) return if devs.tell.get_mounted_sample() is not None: @@ -234,17 +276,21 @@ def common2dh(devs: BeamlineDevices, cfg: BeamlineConfig): print(f"Error for unmounting: {e}") devs.tell.dry(wait_cold=-1, wait=False) + def dh2sa(devs: BeamlineDevices, cfg: BeamlineConfig): devs.bec_worker.move_to(BeamlineState.SAMPLE_ALIGNMENT) + def common2maintenance(devs: BeamlineDevices, cfg: BeamlineConfig): devs.bec_worker.move_to(BeamlineState.MAINTENANCE) + if __name__ == "__main__": - from aare.common.beamline import mx_beamline + from aarecommon.config.beamline import mx_beamline + beamline = mx_beamline() config = BeamlineConfig(beamline) devices = BeamlineDevices(beamline) common_2rse(devices, config) time.sleep(10.0) - rse2sa(devices, config) \ No newline at end of file + rse2sa(devices, config) diff --git a/src/aare/devices/aerotech.py b/src/aare/devices/aerotech.py index 0d69a767..cbd073ff 100644 --- a/src/aare/devices/aerotech.py +++ b/src/aare/devices/aerotech.py @@ -1,29 +1,40 @@ from typing import Optional, Union +from aarecommon.config.beamline import cfg_get, mx_beamline +from aarecommon.errors.exception_handler import AerotechCommunicationError +from aarecommon.math.coordinate import AerotechCoordinate, Coordinate +from aarecommon.models.beamline import MXBeamline +from aarescan_client import ( + ApiClient, + AxisStatus, + Configuration, + DefaultApi, + RotationRequest, + Status, + Target, +) from aarescan_client.models.grid_request import GridRequest from aarescan_client.models.screen_request import ScreenRequest -from aare.common.beamline import MXBeamline, mx_beamline, cfg_get -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) -from aare.common.logger_config import setup_logger +from aarecommon.config.logger import setup_logger logger = setup_logger("aareDAQ") -class AerotechController(object): +class AerotechController(object): def __init__(self, bl: MXBeamline): if bl == MXBeamline.X06DA: self.__simulated = False - self.__base = cfg_get("daq.hardware.aerotech_url", "http://mx-x06da-queue-01.psi.ch:5234") + self.__base = cfg_get( + "daq.hardware.aerotech_url", "http://mx-x06da-queue-01.psi.ch:5234" + ) elif bl == MXBeamline.X10SA: self.__simulated = False - self.__base = cfg_get("daq.hardware.aerotech_url", "http://mx-x10sa-queue-01.psi.ch:5234") + self.__base = cfg_get( + "daq.hardware.aerotech_url", "http://mx-x10sa-queue-01.psi.ch:5234" + ) elif bl == MXBeamline.X06SA: raise NotImplementedError("Not implemented aerotech url for X06SA") elif bl == MXBeamline.SIMULATED: @@ -38,10 +49,10 @@ class AerotechController(object): self.__api = DefaultApi(self.__client) def __make_aerotech_target( - self, - coord: AerotechCoordinate, - wait: bool = False, - incremental: bool = False, + self, + coord: AerotechCoordinate, + wait: bool = False, + incremental: bool = False, ) -> Target: at_mm = coord.at_mm @@ -55,7 +66,9 @@ class AerotechController(object): ) 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) + return AerotechCoordinate( + at_mm=Coordinate(x=target.x, y=target.y, z=target.z), omega_deg=target.u + ) def cancel(self): try: @@ -71,7 +84,7 @@ class AerotechController(object): def is_idle(self) -> bool: try: status = self.__api.status_get() - return status.state == 'Idle' + return status.state == "Idle" except Exception as e: raise AerotechCommunicationError( "Aerotech status check failed", @@ -93,16 +106,41 @@ class AerotechController(object): 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 Status( + state=Status.State.IDLE, + x=AxisStatus( + pos=self.__pos.x, + vel=self.__vel, + enabled=False, + homed=False, + moving=False, + fault=False, + ), + y=AxisStatus( + pos=self.__pos.y, + vel=self.__vel, + enabled=False, + homed=False, + moving=False, + fault=False, + ), + z=AxisStatus( + pos=self.__pos.z, + vel=self.__vel, + enabled=False, + homed=False, + moving=False, + fault=False, + ), + u=AxisStatus( + pos=self.__pos.u, + vel=self.__vel, + enabled=False, + homed=False, + moving=False, + fault=False, + ), + ) try: return self.__api.status_get() except Exception as e: @@ -113,7 +151,7 @@ class AerotechController(object): operation="GET", ) from e - def move_home(self, wait:bool=True, incremental:bool=False): + def move_home(self, wait: bool = True, incremental: bool = False): if self.__simulated: self.__pos = AEROTECH_HOME return self.__pos @@ -142,16 +180,18 @@ class AerotechController(object): ) from e def position( - self, - target: AerotechCoordinate, - /, - wait: bool = True, - incremental: bool = False, + self, + target: AerotechCoordinate, + /, + wait: bool = True, + incremental: bool = False, ): if self.__simulated: return self.__pos - payload = self.__make_aerotech_target(target, wait=wait, incremental=incremental) + payload = self.__make_aerotech_target( + target, wait=wait, incremental=incremental + ) try: return self.__api.position_post(payload) except Exception as e: @@ -162,15 +202,18 @@ class AerotechController(object): operation="POST", ) from e - def rotation_scan(self, rotation_deg: float | int, - time_sec: float | int, - start_pos_deg: float | int, - run_async: bool = False): + def rotation_scan( + self, + rotation_deg: float | int, + time_sec: float | int, + start_pos_deg: float | int, + run_async: bool = False, + ): payload = RotationRequest( rotation_deg=rotation_deg, time_sec=time_sec, start_pos_deg=start_pos_deg, - run_async=run_async + run_async=run_async, ) if self.__simulated: return payload @@ -185,16 +228,15 @@ class AerotechController(object): operation="POST", ) from e - - def grid_scan(self, - grid_elem_count_y: int, - grid_elem_size_y_um: int | float, - time_sec: int|float, - grid_elem_size_x_um: Optional[Union[float, int]] = None, - grid_elem_count_x: Optional[int] = None, - run_async: Optional[bool] = False - - ): + def grid_scan( + self, + grid_elem_count_y: int, + grid_elem_size_y_um: int | float, + time_sec: int | float, + grid_elem_size_x_um: Optional[Union[float, int]] = None, + grid_elem_count_x: Optional[int] = None, + run_async: Optional[bool] = False, + ): payload = GridRequest( grid_elem_count_x=grid_elem_count_x, grid_elem_count_y=grid_elem_count_y, @@ -215,24 +257,25 @@ class AerotechController(object): operation="POST", ) from e - def screening_scan(self, - rotation_deg: float | int, - wedge_deg: float | int, - time_sec: float | int, - steps: int, - run_async: bool = False - ): + def screening_scan( + self, + rotation_deg: float | int, + wedge_deg: float | int, + time_sec: float | int, + steps: int, + run_async: bool = False, + ): payload = ScreenRequest( rotation_deg=rotation_deg, wedge_deg=wedge_deg, time_sec=time_sec, steps=steps, - run_async=run_async + run_async=run_async, ) if self.__simulated: return payload try: - logger.info('sending screening scan request to aerotech') + logger.info("sending screening scan request to aerotech") return self.__api.screening_post(payload) except Exception as e: raise AerotechCommunicationError( @@ -245,8 +288,8 @@ class AerotechController(object): if __name__ == "__main__": beamline = mx_beamline() - #print(beamline) + # print(beamline) controller = AerotechController(beamline) - #controller.print_status(colored=True, compact=True) + # controller.print_status(colored=True, compact=True) print(controller.get_position()) controller.cancel() diff --git a/src/aare/devices/bec_worker.py b/src/aare/devices/bec_worker.py index c4fccdd3..c30f36bf 100644 --- a/src/aare/devices/bec_worker.py +++ b/src/aare/devices/bec_worker.py @@ -1,44 +1,48 @@ +import time from enum import Enum -from typing import List, Optional, Any +from typing import Any, List, Optional +from aarecommon.config.beamline import cfg_get, mx_beamline +from aarecommon.config.logger import setup_logger +from aarecommon.config.logger_events import log_timing +from aarecommon.errors.exception_handler import BECCommunicationError +from aarecommon.models.beamline import MXBeamline from bec_ipython_client import BECIPythonClient from bec_ipython_client.signals import OperationMode +from bec_lib.procedures.helper import BackendProcedureHelper, FrontendProcedureHelper from bec_lib.service_config import ServiceConfig -from bec_lib.procedures.helper import FrontendProcedureHelper, BackendProcedureHelper - -from aare.common.beamline import MXBeamline, mx_beamline, cfg_get -from aare.common.logger_config import setup_logger -from aare.common.logger_events import log_timing -from aare.common.exception_handler import BECCommunicationError - -import time logger = setup_logger("aareDAQ") -#specify up to 10 queue to runs in parallel, request more if needed! +# specify up to 10 queue to runs in parallel, request more if needed! # st = client.proc.request_new("sleep", ((), {"time_s":5}), queue="test") -#to see all deevices -#devs.show_all +# to see all deevices +# devs.show_all -#helper fucntions +# helper fucntions # helper.get.active_and_pending_queue_names() # helper.get.running_procedures() # helper.request.abort_queue() + class DetectorCoverEnum(str, Enum): """Enum for the detector cover position position devices can take string or number to move Currently dictated string in DAQ for consitency""" - OPEN = 'open' #2 - CLOSED = 'closed' #1 + + OPEN = "open" # 2 + CLOSED = "closed" # 1 + class BrightnessEnum(str, Enum): """Enum for the backlight brightness position devices can take string or number to move """ - ON = 'on' - OFF = 'off' + + ON = "on" + OFF = "off" + class BeamlineState(str, Enum): ROBOT_SAMPLE_EXCHANGE = "robot_sample_exchange" @@ -54,7 +58,7 @@ class BeamlineState(str, Enum): class BECClientWorker: - def __init__(self, beamline:MXBeamline, name:str = "default"): + def __init__(self, beamline: MXBeamline, name: str = "default"): BEAMLINE = beamline.value.lower() self.beamline = beamline if self.beamline is MXBeamline.X06DA: @@ -76,12 +80,14 @@ class BECClientWorker: logger.debug(f"Initializing BECClientWorker for {BEAMLINE} beamline") host = cfg_get("daq.hardware.bec_url", f"{BEAMLINE}-bec-001.psi.ch") service_config = ServiceConfig(redis={"host": host, "port": 6379}) - service_config.config["log_writer"]["base_path"]='/tmp/logs' - #service_config.config["user_macros"]["base_path"]=f'/sls/{BEAMLINE}/config/bec/production/pxiii_bec/pxiii_bec' - #print(service_config.config) - self.client = BECIPythonClient(config=service_config,mode=OperationMode.Procedure) + service_config.config["log_writer"]["base_path"] = "/tmp/logs" + # service_config.config["user_macros"]["base_path"]=f'/sls/{BEAMLINE}/config/bec/production/pxiii_bec/pxiii_bec' + # print(service_config.config) + self.client = BECIPythonClient( + config=service_config, mode=OperationMode.Procedure + ) self.client.start() - #self.client.config.update_session_with_file("/sls/x10sa/config/bec/production/bec/bec_lib/bec_lib/config_helper.py") + # self.client.config.update_session_with_file("/sls/x10sa/config/bec/production/bec/bec_lib/bec_lib/config_helper.py") self.dev = self.client.device_manager.devices print(self.dev.keys()) self.scans = self.client.scans @@ -94,7 +100,7 @@ class BECClientWorker: self.__init_beamline_environment() except Exception as e: logger.error(f"Error initialising BEC devices: {e}") - #raise self._raise_bec_error(e, operation="create planner") + # raise self._raise_bec_error(e, operation="create planner") self.simulated = True logger.debug(f"simulated is {self.simulated}") @@ -102,8 +108,8 @@ class BECClientWorker: def __init_beamline_environment(self): try: self.position_devices, self.planner = init_beamline_environment() - self.__backlight_brightness = self.position_devices['bl_bright'] - self.__frontlight_brightness = self.position_devices['fl_bright'] + self.__backlight_brightness = self.position_devices["bl_bright"] + self.__frontlight_brightness = self.position_devices["fl_bright"] self.__zoom = self.dev.scam_zoom self._ring_current = self.dev.sls_current except Exception as e: @@ -121,9 +127,11 @@ class BECClientWorker: self.ring_current = None raise Exception(f"Error initialising BEC devices: {e}") - def _raise_bec_error(self, exc: Exception, *, operation: str, tags:Optional[List[str]] = None) -> None: + def _raise_bec_error( + self, exc: Exception, *, operation: str, tags: Optional[List[str]] = None + ) -> None: message = f"BEC operation '{operation}' failed: {type(exc).__name__}: {exc}" - #logger.exception(message) + # logger.exception(message) if tags is None: tags = ["error", "bec"] else: @@ -144,7 +152,7 @@ class BECClientWorker: message=message, error=True, error_message=f"Error during '{operation}': {exc}", - tags=tags + tags=tags, ) except Exception as e: logger.error(f"Error sending scilog message: {e}") @@ -162,7 +170,7 @@ class BECClientWorker: exception=exc, ) from exc - def __set_scilog_tags(self, tags:Optional[List[str]]=None): + def __set_scilog_tags(self, tags: Optional[List[str]] = None): try: if tags: self.client.messaging.scilog.set_default_tags(tags) @@ -173,19 +181,19 @@ class BECClientWorker: self._raise_bec_error(e, operation="set scilog tags") def scilog_msg( - self, - message:str, - error:bool=False, - warning:bool=False, - error_message:Optional[str]=None, - attachments=None, - bold:bool=False, - italic:bool=False, - color:Optional[str]=None, - additonal_text:Optional[List[str]]=None, - tags:Optional[List[str]]=None + self, + message: str, + error: bool = False, + warning: bool = False, + error_message: Optional[str] = None, + attachments=None, + bold: bool = False, + italic: bool = False, + color: Optional[str] = None, + additonal_text: Optional[List[str]] = None, + tags: Optional[List[str]] = None, ): - if color and color not in ["red", "green", "yellow","blue","pink"]: + if color and color not in ["red", "green", "yellow", "blue", "pink"]: logger.warning("specified color not in allowed list,using default") color = None try: @@ -226,7 +234,7 @@ class BECClientWorker: except Exception as e: logger.error(f"Error sending scilog message: {e}") - def run_macro(self, macro_name:str, *args, queue:str = "default", **kwargs): + def run_macro(self, macro_name: str, *args, queue: str = "default", **kwargs): if self.simulated: logger.debug(f"Simulating macro {macro_name}") return None @@ -235,7 +243,9 @@ class BECClientWorker: except Exception as e: self._raise_bec_error(e, operation=f"run_macro:{macro_name}") - def run_macro_blocked(self, macro_name:str, *args, queue:str = "default", **kwargs): + def run_macro_blocked( + self, macro_name: str, *args, queue: str = "default", **kwargs + ): if self.simulated: logger.debug(f"Simulating macro {macro_name}") return None @@ -249,7 +259,7 @@ class BECClientWorker: self._raise_bec_error(e, operation=f"run_macro_blocked:{macro_name}") @log_timing(logger, "BEC move_to") - def move_to(self, state:BeamlineState): + def move_to(self, state: BeamlineState): start = time.perf_counter() logger.debug(f"simulated is {self.simulated}") logger.info(f"BEC move_to requested: {state.value}") @@ -259,10 +269,14 @@ class BECClientWorker: try: self.planner.move_to(state) if self.planner.is_state(state): - logger.info(f"BEC move_to completed: {state.value} in {time.perf_counter() - start:.2f}s") + logger.info( + f"BEC move_to completed: {state.value} in {time.perf_counter() - start:.2f}s" + ) return True else: - logger.warning(f"Failed to move to {state.value} in {time.perf_counter() - start:.2f}s") + logger.warning( + f"Failed to move to {state.value} in {time.perf_counter() - start:.2f}s" + ) return False except Exception as e: self._raise_bec_error(e, operation=f"planner.move_to:{state.value}") @@ -294,7 +308,9 @@ class BECClientWorker: return [] try: self.__list_all_macros() - raw_macros = [name for name, _ in self.client.macros._update_handler.macros.items()] + raw_macros = [ + name for name, _ in self.client.macros._update_handler.macros.items() + ] if raw_macros is None: logger.warning("BEC returned no user macros; treating as empty list") return [] @@ -337,14 +353,18 @@ class BECClientWorker: List of position device names after reinitialisation. """ if self.simulated: - logger.debug(f"Simulating reinitialise_planner_and_position_devices(method={method})") + logger.debug( + f"Simulating reinitialise_planner_and_position_devices(method={method})" + ) return [] try: self.client.config.update_session_with_file( - f"/sls/{self.beamline}/config/bec/production/{self._beamline_name}_bec/{self._beamline_name}_bec/device_configs/{self._beamline_name}-devices.yaml" + f"/sls/{self.beamline}/config/bec/production/{self._beamline_name}_bec/{self._beamline_name}_bec/device_configs/{self._beamline_name}-devices.yaml" ) self.__init_beamline_environment() - logger.info(f"Reinitialised BEC planner and position devices using method={method}") + logger.info( + f"Reinitialised BEC planner and position devices using method={method}" + ) return self.list_position_devices() except Exception as e: self._raise_bec_error( @@ -359,17 +379,23 @@ class BECClientWorker: try: mono_pitch_scan(plot) except Exception as e: - self._raise_bec_error(e, operation="mono_pitch_scan", tags=["mono_pitch_scan"]) + self._raise_bec_error( + e, operation="mono_pitch_scan", tags=["mono_pitch_scan"] + ) if self.beamline is MXBeamline.X06DA: - addtional_text = [f"New dcm_pitch position: {self.dev.dcm_pitch.position:5f}"] + addtional_text = [ + f"New dcm_pitch position: {self.dev.dcm_pitch.position:5f}" + ] else: - addtional_text = [f"New dcm_theta2 position: {self.dev.dccm_theta2.position:5f}"] + addtional_text = [ + f"New dcm_theta2 position: {self.dev.dccm_theta2.position:5f}" + ] self.scilog_msg( message="Mono pitch scan completed", bold=True, color="green", tags=["mono_pitch_scan"], - additonal_text=addtional_text + additonal_text=addtional_text, ) def check_current_energy(self): @@ -381,47 +407,54 @@ class BECClientWorker: def change_energy(self, value: float | int, plot: bool = False): current_energy = self.check_current_energy() logger.info(f"Current energy: {current_energy:.1f} eV") - logger.info(f"Change energy requested: from {current_energy:.1f} to {value:.1f} eV") + logger.info( + f"Change energy requested: from {current_energy:.1f} to {value:.1f} eV" + ) try: bl_energy(value, move_gap=False, mono_scan=True, plot=plot) except Exception as e: self._raise_bec_error( e, operation=f"Requested energy change from:" - f"{current_energy:.1f} to {value} eV", - tags=["energy_change"] + f"{current_energy:.1f} to {value} eV", + tags=["energy_change"], ) if abs(value - self.check_current_energy()) > 1: - logger.warning(f"Energy change may have failed, current energy: {self.check_current_energy()} eV") + logger.warning( + f"Energy change may have failed, current energy: {self.check_current_energy()} eV" + ) if beamline is MXBeamline.X10SA: - additonal_text = [f"New dcm_bragg position: {self.dev.dcm_bragg.position:4g} mrad", - f"New dcm_pitch position: {self.dev.dcm_pitch.position:4g} ", - f"Previous energy: {current_energy:.1f} eV ", - f"Requested energy: {value:.1f} eV ", - f"New current energy: {self.check_current_energy():.1f} eV"] + additonal_text = [ + f"New dcm_bragg position: {self.dev.dcm_bragg.position:4g} mrad", + f"New dcm_pitch position: {self.dev.dcm_pitch.position:4g} ", + f"Previous energy: {current_energy:.1f} eV ", + f"Requested energy: {value:.1f} eV ", + f"New current energy: {self.check_current_energy():.1f} eV", + ] else: - additonal_text = [f"New dccm_theta1 position: {self.dev.dccm_theta1.position:4g} mrad", - f"New dccm_theta2 position: {self.dev.dccm_theta2.position:4g} mrad", - f"Previous energy: {current_energy:.1f} eV ", - f"Requested energy: {value:.1f} eV ", - f"New current energy: {self.check_current_energy():.1f} eV"] + additonal_text = [ + f"New dccm_theta1 position: {self.dev.dccm_theta1.position:4g} mrad", + f"New dccm_theta2 position: {self.dev.dccm_theta2.position:4g} mrad", + f"Previous energy: {current_energy:.1f} eV ", + f"Requested energy: {value:.1f} eV ", + f"New current energy: {self.check_current_energy():.1f} eV", + ] self.scilog_msg( message=f"Moved from {current_energy:.1f} eV to {value:.1f} eV", bold=True, color="green", tags=["energy_change"], - additonal_text=additonal_text + additonal_text=additonal_text, ) - def get_det_z(self): try: return self.dev.det_z.position except Exception as e: self._raise_bec_error(e, operation=f"get_det_z", tags=["det_z"]) - def det_z(self, value:float, timeout:int | None = None ): + def det_z(self, value: float, timeout: int | None = None): """timeout is None or integer in s""" try: status = self.scans.mv(self.dev.det_z, value, relative=False) @@ -429,7 +462,9 @@ class BECClientWorker: status.wait(timeout=timeout) return status except Exception as e: - self._raise_bec_error(e, operation=f"scans.mv:det_z:{value}", tags=["det_z"]) + self._raise_bec_error( + e, operation=f"scans.mv:det_z:{value}", tags=["det_z"] + ) def get_det_y(self): try: @@ -437,7 +472,7 @@ class BECClientWorker: except Exception as e: self._raise_bec_error(e, operation=f"get_det_z", tags=["det_z"]) - def det_y(self, value:float, timeout:int | None = None ): + def det_y(self, value: float, timeout: int | None = None): """timeout is None or integer in s""" try: status = self.scans.mv(self.dev.det_y, value, relative=False) @@ -445,7 +480,9 @@ class BECClientWorker: status.wait(timeout=timeout) return status except Exception as e: - self._raise_bec_error(e, operation=f"scans.mv:det_y:{value}", tags=["det_z"]) + self._raise_bec_error( + e, operation=f"scans.mv:det_y:{value}", tags=["det_z"] + ) @property def backlight_brightness(self) -> BrightnessEnum: @@ -457,19 +494,21 @@ class BECClientWorker: self._raise_bec_error( e, operation=f"backlight brightness, could not get backlight brightness", - tags=["backlight"] + tags=["backlight"], ) raise @backlight_brightness.setter - def backlight_brightness(self, value:int | str): + def backlight_brightness(self, value: int | str): """Set the backlight brightness to the specified value""" if self.simulated: return try: self.__backlight_brightness.move(value) except Exception as e: - self._raise_bec_error(e, operation=f"backlight_brightness:{value}", tags=["backlight"]) + self._raise_bec_error( + e, operation=f"backlight_brightness:{value}", tags=["backlight"] + ) raise def get_backlight_pos(self) -> BrightnessEnum: @@ -491,19 +530,19 @@ class BECClientWorker: self._raise_bec_error( e, operation=f"backlight toggle, could not change backlight on/off ", - tags=["backlight"] + tags=["backlight"], ) def save_current_bs_pos(self): - save_current_position(self.dev.bs_z, 'safe') + save_current_position(self.dev.bs_z, "safe") def save_current_collimator_pos(self): - save_current_position(self.dev.coll_y, 'work') + save_current_position(self.dev.coll_y, "work") def save_current_aerotech_position(self): - save_current_position(self.dev.aerotech, 'work', axis = 'x') - save_current_position(self.dev.aerotech, 'work', axis='y') - save_current_position(self.dev.aerotech, 'work', axis='z') + save_current_position(self.dev.aerotech, "work", axis="x") + save_current_position(self.dev.aerotech, "work", axis="y") + save_current_position(self.dev.aerotech, "work", axis="z") self.save_config_and_reload_devices() def save_config_and_reload_devices(self): @@ -514,19 +553,20 @@ class BECClientWorker: return self.__zoom.position @zoom.setter - def zoom(self, value:float): + def zoom(self, value: float): self.scans.umv(self.__zoom, value, relative=False) @property def ring_current(self) -> float: return self._ring_current.get() + if __name__ == "__main__": import time - print(time.ctime(), ' starting BEC Client') + + print(time.ctime(), " starting BEC Client") beamline = mx_beamline() try: - client = BECClientWorker(beamline) except Exception as e: print(f"Error: {e}") @@ -534,12 +574,13 @@ if __name__ == "__main__": client.shutdown_client() except Exception as e: import sys + sys.exit(1) - #print(client.get_det_cov(actual=True)) - #print(client.is_state(BeamlineState.ROBOT_SAMPLE_EXCHANGE)) + # print(client.get_det_cov(actual=True)) + # print(client.is_state(BeamlineState.ROBOT_SAMPLE_EXCHANGE)) try: - print('startting backlight brightness test') - print('initial value') + print("startting backlight brightness test") + print("initial value") # print(client.client.show_last_alarm()) # print(client._raise_bec_error(exc=Exception("test"), operation="test", tags=["test"])) print(client.ring_current) @@ -568,13 +609,13 @@ if __name__ == "__main__": # print(client.backlight_brightness) # print('setting to OFF') # client.backlight_brightness = BrightnessEnum.OFF - #client.backlight_toggle() + # client.backlight_toggle() # print(client.get_backlight_pos()) # print(client.backlight_brightness) - #client.scilog_msg("Testing scilog messages with color = yellow and italic", italic=True, + # client.scilog_msg("Testing scilog messages with color = yellow and italic", italic=True, # color="green", warning=False) except Exception as e: - #client._raise_bec_error(e, operation="send message") + # client._raise_bec_error(e, operation="send message") client.shutdown_client() print(f"Error: {e}") @@ -605,23 +646,22 @@ if __name__ == "__main__": # time.sleep(2.0) # client.planner.current_state() - # client.load_user_macros() - #client.macros.mono_pitch_scan(False) + # client.macros.mono_pitch_scan(False) # status = client.run_macro("planner.current_state", queue="default") - #client.planner.move_to('manual_sample_exchange') + # client.planner.move_to('manual_sample_exchange') # print(status) # status.wait() # print(status) - #try: - #status=client.proc.run_macro("mono_pitch_scan", queue="test") - #planner.move_to('manual_sample_exchange') + # try: + # status=client.proc.run_macro("mono_pitch_scan", queue="test") + # planner.move_to('manual_sample_exchange') # print(status) # status.wait() # status.cancel() # print(status) - #client.client.macros.mono_pictch_scan(False) + # client.client.macros.mono_pictch_scan(False) # try: # # a=client.a2e_runner(160, "iln") # # print(a) @@ -672,5 +712,4 @@ if __name__ == "__main__": # client.shutdown() -#backend wont work unless bec server will work, frontend anywehre with user access - +# backend wont work unless bec server will work, frontend anywehre with user access diff --git a/src/aare/devices/experimental_hutch_shutter.py b/src/aare/devices/experimental_hutch_shutter.py index 83526763..c72541dd 100644 --- a/src/aare/devices/experimental_hutch_shutter.py +++ b/src/aare/devices/experimental_hutch_shutter.py @@ -1,14 +1,14 @@ -from aare.common.beamline import MXBeamline +from aarecommon.models.beamline import MXBeamline from epics import PV + class ExperimentalHutchShutter: def __init__(self, beamline: MXBeamline): BEAMLINE = beamline.value.upper() - self.__close= PV(f"{BEAMLINE}-EH1-PSYS:SH-A-CLOSE-SET") # 1, 0 + self.__close = PV(f"{BEAMLINE}-EH1-PSYS:SH-A-CLOSE-SET") # 1, 0 self.__open = PV(f"{BEAMLINE}-EH1-PSYS:SH-A-OPEN-SET") self.__state = PV(f"{BEAMLINE}-OP-PSH1-EMLS-0010:OPEN") - def state(self): state = self.__state.get() if state == "Open" or state == 1: @@ -22,9 +22,9 @@ class ExperimentalHutchShutter: def open(self): self.__open.put(0) self.__open.put(1) - #print(self.__open.get()) + # print(self.__open.get()) self.__open.put(0) def close(self): self.__close.put(1) - self.__close.put(0) \ No newline at end of file + self.__close.put(0) diff --git a/src/aare/devices/filter_transmission.py b/src/aare/devices/filter_transmission.py index 46dee47f..77db425f 100644 --- a/src/aare/devices/filter_transmission.py +++ b/src/aare/devices/filter_transmission.py @@ -1,9 +1,8 @@ import time +from aarecommon.models.beamline import MXBeamline from epics import PV, poll -from aare.common.beamline import MXBeamline - class FilterTransmission: def __init__(self, bl: MXBeamline, timeout=60.0, fail_on_timeout=True): @@ -51,4 +50,6 @@ class FilterTransmission: timeisup = time.time() > timeout poll(0.1) if timeisup and self._fail_on_timeout: - raise RuntimeError("timeout waiting for filters to achieve requested transmission.") + raise RuntimeError( + "timeout waiting for filters to achieve requested transmission." + ) diff --git a/src/aare/devices/fluorimeter.py b/src/aare/devices/fluorimeter.py index ca5b298c..9e9e8190 100644 --- a/src/aare/devices/fluorimeter.py +++ b/src/aare/devices/fluorimeter.py @@ -1,31 +1,33 @@ import time +from aarecommon.config.logger import setup_logger +from aarecommon.models.beamline import MXBeamline from epics import PV, poll -from aare.common.beamline import MXBeamline - -from aare.common.logger_config import setup_logger - logger = setup_logger("aareaDAQ") -class Fluorimeter(object): +class Fluorimeter(object): def __init__(self, beamline: MXBeamline, **kwargs): BEAMLINE = beamline.value.upper() - self.__start = PV(f"{BEAMLINE}-ES-SiD:mca1Start") #0 done ,1 start - self.__stop = PV(f"{BEAMLINE}-ES-SiD:mca1Stop") #0 done ,1 stop - self.__erase_and_start = PV(f"{BEAMLINE}-ES-SiD:mca1EraseStart") #0 done,1 start - self.__erase = PV(f"{BEAMLINE}-ES-SiD:mca1Erase") #0 done,1 erase + self.__start = PV(f"{BEAMLINE}-ES-SiD:mca1Start") # 0 done ,1 start + self.__stop = PV(f"{BEAMLINE}-ES-SiD:mca1Stop") # 0 done ,1 stop + self.__erase_and_start = PV( + f"{BEAMLINE}-ES-SiD:mca1EraseStart" + ) # 0 done,1 start + self.__erase = PV(f"{BEAMLINE}-ES-SiD:mca1Erase") # 0 done,1 erase - self.__preset_mode = PV(f"{BEAMLINE}-ES-SiD:dxp1:PresetMode") # set mode - self.__status = PV(f"{BEAMLINE}-ES-SiD:mca1.ACQG") # 0 done, 1 acquire + self.__preset_mode = PV(f"{BEAMLINE}-ES-SiD:dxp1:PresetMode") # set mode + self.__status = PV(f"{BEAMLINE}-ES-SiD:mca1.ACQG") # 0 done, 1 acquire - self.__real_time = PV(f"{BEAMLINE}-ES-SiD:mca1.PRTM") #float - self.__live_time = PV(f"{BEAMLINE}-ES-SiD:mca1.PLTM") #float - self.__elapsed_real_time = PV(f"{BEAMLINE}-ES-SiD:mca1.ERTM") #float - self.__elapsed_live_time = PV(f"{BEAMLINE}-ES-SiD:mca1.ELTM") #float - self.__elapsed_trigger_live_time = PV(f"{BEAMLINE}-ES-SiD:dxp1:ElapsedTriggerLiveTime") + self.__real_time = PV(f"{BEAMLINE}-ES-SiD:mca1.PRTM") # float + self.__live_time = PV(f"{BEAMLINE}-ES-SiD:mca1.PLTM") # float + self.__elapsed_real_time = PV(f"{BEAMLINE}-ES-SiD:mca1.ERTM") # float + self.__elapsed_live_time = PV(f"{BEAMLINE}-ES-SiD:mca1.ELTM") # float + self.__elapsed_trigger_live_time = PV( + f"{BEAMLINE}-ES-SiD:dxp1:ElapsedTriggerLiveTime" + ) self.__instant_dead_time = PV(f"{BEAMLINE}-ES-SiD:mca1.IDTIM") self.__average_dead_time = PV(f"{BEAMLINE}-ES-SiD:mca1.DTIM") @@ -39,14 +41,14 @@ class Fluorimeter(object): self.__saveFile = PV(f"{BEAMLINE}-ES-SiD:SaveSystemFile") self.__saveFile_name = PV(f"{BEAMLINE}-ES-SiD:SaveSystem") - self.__saveFile_rbv = PV(f"{BEAMLINE}-ES-SiD:SaveSystem_RBV") #1 ssave, 0 done - self.__roi_1 = self.get_roi(BEAMLINE, 1) #currently SiEsc - self.__roi_2 = self.get_roi(BEAMLINE, 2) #currently MnKa + self.__saveFile_rbv = PV(f"{BEAMLINE}-ES-SiD:SaveSystem_RBV") # 1 ssave, 0 done + self.__roi_1 = self.get_roi(BEAMLINE, 1) # currently SiEsc + self.__roi_2 = self.get_roi(BEAMLINE, 2) # currently MnKa self.__calibration_offset = PV(f"{BEAMLINE}-ES-SiD:mca1.CALO") self.__calibration_slope = PV(f"{BEAMLINE}-ES-SiD:mca1.CALS") - def get_roi(self, bl, roi_num:int): + def get_roi(self, bl, roi_num: int): if roi_num not in [1, 2]: raise ValueError(f"Invalid ROI number: {roi_num}") roi_name = f"{bl}-ES-SiD:mca1.R{roi_num}NM" @@ -105,7 +107,7 @@ class Fluorimeter(object): poll(0.1) def set_preset_mode(self, mode: int): - """ 0 No preset, 1 Real time, 2 Live time, 3 Events, 4 Triggers""" + """0 No preset, 1 Real time, 2 Live time, 3 Events, 4 Triggers""" if mode not in [0, 1, 2, 3, 4]: raise ValueError(f"Invalid preset mode: {mode}") try: @@ -113,7 +115,7 @@ class Fluorimeter(object): except Exception as e: logger.error(f"Error setting preset mode to {mode}: {e}") - def save_file(self, filename:str|None = None, timeout_s: float = 30.0): + def save_file(self, filename: str | None = None, timeout_s: float = 30.0): if not filename: filename = f"KETEK_{time.strftime('%Y%m%d_%H%M%S')}.txt" self.__saveFile_name.put(filename) @@ -173,4 +175,3 @@ class Fluorimeter(object): @poll_time.setter def poll_time(self, value): self.__poll_time.put(value) - diff --git a/src/aare/devices/jfjoch.py b/src/aare/devices/jfjoch.py index fa5b1270..683bbfed 100644 --- a/src/aare/devices/jfjoch.py +++ b/src/aare/devices/jfjoch.py @@ -3,12 +3,13 @@ import time from enum import Enum import jfjoch_client +from aarecommon.config.beamline import get_jfjoch_url +from aarecommon.errors.exception_handler import JFJochCommunicationError +from aarecommon.models.beamline import MXBeamline +from aarecommon.models.models import DAQStatusModel, FluorescenceSpectrumOutputModel +from aarecommon.models.raster_grid import RasterGridRequest +from aarecommon.models.rotation_scan import RotationScanRequest -from aare.common.beamline import MXBeamline, get_jfjoch_url -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 class ScanTypeEnum(Enum): RASTER = "Raster" @@ -18,9 +19,10 @@ class ScanTypeEnum(Enum): HELICAL = "Helical" UNKNOWN = "Unknown" + class JFJochWrapper: def __init__(self, bl: MXBeamline): - self.__simulated = (bl == MXBeamline.SIMULATED) + self.__simulated = bl == MXBeamline.SIMULATED self.__url = get_jfjoch_url(bl) if self.__url == "simulated": @@ -28,7 +30,9 @@ class JFJochWrapper: self.__api = None return - self.__client = jfjoch_client.ApiClient(jfjoch_client.Configuration(host=self.__url)) + self.__client = jfjoch_client.ApiClient( + jfjoch_client.Configuration(host=self.__url) + ) self.__api = jfjoch_client.DefaultApi(self.__client) self.cancel() # if not self.is_idle(): @@ -48,12 +52,12 @@ class JFJochWrapper: return None def _raise_jfjoch_error( - self, - message: str, - *, - error: Exception, - operation: str, - endpoint: str, + self, + message: str, + *, + error: Exception, + operation: str, + endpoint: str, ) -> None: raise JFJochCommunicationError( message, @@ -87,16 +91,16 @@ class JFJochWrapper: def is_idle(self) -> bool: status = self.__api.status_get() - return status.state == 'Idle' + return status.state == "Idle" def __format_dataset_settings( - self, - r: RasterGridRequest | RotationScanRequest, - s: DAQStatusModel, - f: FluorescenceSpectrumOutputModel | None = None, - async_start: bool = True + self, + r: RasterGridRequest | RotationScanRequest, + s: DAQStatusModel, + f: FluorescenceSpectrumOutputModel | None = None, + async_start: bool = True, ) -> jfjoch_client.DatasetSettings: - #common sample parameter intiialisation + # common sample parameter intiialisation if s.sample is None: pgroup = "p16371" sample = "unknown_sample" @@ -104,7 +108,7 @@ class JFJochWrapper: pgroup = s.sample.user sample = s.sample.sample_name - #raster grid, rotation or screening specific parameter initialisation + # raster grid, rotation or screening specific parameter initialisation wedge = None if isinstance(r, RasterGridRequest): data_folder = f"{pgroup}/raw/raster" @@ -130,11 +134,11 @@ class JFJochWrapper: "Scan request has no detector distance (dtz); setup_datacollection " "must validate and set request.dtz before JFJoch is configured" ) - #build common dataset settings + # build common dataset settings dataset_settings = jfjoch_client.DatasetSettings( beam_x_pxl=s.diffraction.beam_center_pxl[0], beam_y_pxl=s.diffraction.beam_center_pxl[1], - ntrigger= trigger, + ntrigger=trigger, images_per_trigger=images, detector_distance_mm=r.dtz, file_prefix=f"{data_folder}/{r.file_prefix}", @@ -150,17 +154,17 @@ class JFJochWrapper: poni_rot2_rad=s.diffraction.poni_rot2_rad, max_spot_count=1000, detect_ice_rings=True, - async_start=async_start + async_start=async_start, ) - #settup grid or goniometer settings adn add to dataset settings + # settup grid or goniometer settings adn add to dataset settings if isinstance(r, RasterGridRequest): grid_settings = jfjoch_client.GridScan( n_fast=r.n_x, snake=True, step_x_um=r.grid_size_mm.x * 1000.0, step_y_um=r.grid_size_mm.y * 1000.0, - vertical=False + vertical=False, ) dataset_settings.grid_scan = grid_settings @@ -170,15 +174,14 @@ class JFJochWrapper: start=r.start_omega_deg, name="omega", vector=[-1, 0, 0], - screening_wedge_deg=wedge + screening_wedge_deg=wedge, ) dataset_settings.goniometer = goniometer_settings - #if measuring fluoresence, add to dataset settings + # if measuring fluoresence, add to dataset settings if f is not None: xrf = jfjoch_client.DatasetSettingsXrayFluorescenceSpectrum( - energy_eV=f.energy_eV, - data=f.spectrum + energy_eV=f.energy_eV, data=f.spectrum ) dataset_settings.xray_fluorescence_spectrum = xrf @@ -193,7 +196,7 @@ class JFJochWrapper: c=unit_cell_floats[2], alpha=unit_cell_floats[3], beta=unit_cell_floats[4], - gamma=unit_cell_floats[5] + gamma=unit_cell_floats[5], ) dataset_settings.unit_cell = unit_cell if s.sample.aaredb_params.spacegroupnumber: @@ -203,14 +206,16 @@ class JFJochWrapper: return dataset_settings def __start_scan( - self, - scan_type: ScanTypeEnum, - r: RasterGridRequest | RotationScanRequest, - s: DAQStatusModel, - f: FluorescenceSpectrumOutputModel | None = None, - async_start: bool = True + self, + scan_type: ScanTypeEnum, + r: RasterGridRequest | RotationScanRequest, + s: DAQStatusModel, + f: FluorescenceSpectrumOutputModel | None = None, + async_start: bool = True, ): - dataset_settings = self.__format_dataset_settings(r, s, f, async_start=async_start) + dataset_settings = self.__format_dataset_settings( + r, s, f, async_start=async_start + ) try: self.__api.start_post(dataset_settings=dataset_settings) except Exception as e: @@ -221,27 +226,30 @@ class JFJochWrapper: endpoint="start_post", ) - def measure_rotation(self, - r: RotationScanRequest, - s: DAQStatusModel, - f: FluorescenceSpectrumOutputModel | None = None, - async_start:bool = True) -> None: + def measure_rotation( + self, + r: RotationScanRequest, + s: DAQStatusModel, + f: FluorescenceSpectrumOutputModel | None = None, + async_start: bool = True, + ) -> None: if r.screening: self.__start_scan(ScanTypeEnum.SCREENING, r, s, f, async_start=async_start) else: self.__start_scan(ScanTypeEnum.ROTATION, r, s, f, async_start=async_start) - def measure_raster(self, - r: RasterGridRequest, - s: DAQStatusModel, - async_start: bool = True): + def measure_raster( + self, r: RasterGridRequest, s: DAQStatusModel, async_start: bool = True + ): self.__start_scan(ScanTypeEnum.RASTER, r, s, async_start=async_start) def wait_till_running(self, timeout: int | float = 60): if self.__simulated: return None try: - self.__api.wait_until_running_post_with_http_info(timeout=math.ceil(timeout)) + self.__api.wait_until_running_post_with_http_info( + timeout=math.ceil(timeout) + ) return True except Exception as e: self._raise_jfjoch_error( @@ -251,7 +259,9 @@ class JFJochWrapper: endpoint="wait_until_running_post", ) - def wait_till_done(self, timeout : int | float) -> jfjoch_client.models.ScanResult | None: + def wait_till_done( + self, timeout: int | float + ) -> jfjoch_client.models.ScanResult | None: if self.__simulated: return None try: @@ -296,7 +306,9 @@ class JFJochWrapper: ) def take_pedestal(self): - raise NotImplementedError("take_pedestal is not implemented in DAQ through the JFJoch API yet") + raise NotImplementedError( + "take_pedestal is not implemented in DAQ through the JFJoch API yet" + ) def get_diffraction_image( self, @@ -321,17 +333,25 @@ class JFJochWrapper: time.sleep(wait_between_retries_s) raise last_error + if __name__ == "__main__": import logging - from aare.common.coordinate import SmargonCoordinate, Coordinate - from aare.common.models import SessionStatus, BeamlineStateEnum, SampleCameraSettings, BeamlineStatus - from aare.common.sample_geometry import SampleGeometryModel - from aare.common.diffraction_geometry import DiffractionGeometry + + from aarecommon.math.coordinate import Coordinate, SmargonCoordinate + from aarecommon.math.diffraction_geometry import DiffractionGeometry + from aarecommon.math.sample_geometry import SampleGeometryModel + from aarecommon.models.models import ( + BeamlineStateEnum, + BeamlineStatus, + SampleCameraSettings, + SessionStatus, + ) + logging.basicConfig(level=logging.DEBUG) bl = MXBeamline.X10SA wrapper = JFJochWrapper(bl) - #wrapper.initialize() - #wrapper.wait_till_done(timeout=360) + # wrapper.initialize() + # wrapper.wait_till_done(timeout=360) status = DAQStatusModel( bl=BeamlineStatus( name="X10SA", @@ -347,7 +367,7 @@ if __name__ == "__main__": zoom=1.0, commissioning_mode=False, dtz_min=120.0, - dtz_max=1600.0 + dtz_max=1600.0, ), diffraction=DiffractionGeometry( detector_description="EIGER", @@ -358,23 +378,38 @@ if __name__ == "__main__": beam_center_pxl=(1000.0, 1000.0), detector_size_pxl=(4000, 4000), poni_rot1_rad=0.0, - poni_rot2_rad=0.0 + poni_rot2_rad=0.0, ), geom=SampleGeometryModel( beam_location_pxl=Coordinate(x=1000, y=1000), pixel_in_mm=0.001, aerotech=Coordinate(x=0, y=0), aerotech_meas=Coordinate(x=0, y=0), - smargon=SmargonCoordinate(sh_mm=Coordinate(x=0,y=0,z=0), phi_deg=0, chi_deg=0), + smargon=SmargonCoordinate( + sh_mm=Coordinate(x=0, y=0, z=0), phi_deg=0, chi_deg=0 + ), omega_deg=10.5, - beam_size_mm=Coordinate(x=0.02, y=0.01) + beam_size_mm=Coordinate(x=0.02, y=0.01), ), session=SessionStatus(), state=BeamlineStateEnum.Maintenance, - busy=False + busy=False, + ) + wrapper.measure_raster( + RasterGridRequest( + n_x=1, + n_y=1, + grid_size_mm=Coordinate(x=0.05, y=0.05), + file_prefix="test", + exp_time_s=0.1, + dtz=150, + transmission=1, + smargon_top_left=SmargonCoordinate(), + omega_deg=0.0, + ), + status, ) - wrapper.measure_raster(RasterGridRequest(n_x=1, n_y=1, grid_size_mm=Coordinate(x=0.05,y=0.05), file_prefix="test", exp_time_s=0.1,dtz=150,transmission=1,smargon_top_left=SmargonCoordinate(),omega_deg=0.0), status) print("waiting for detector to start") wrapper.wait_till_running() print("success we can measure") - wrapper.cancel() \ No newline at end of file + wrapper.cancel() diff --git a/src/aare/devices/pss_state.py b/src/aare/devices/pss_state.py index 38948452..907f78ea 100644 --- a/src/aare/devices/pss_state.py +++ b/src/aare/devices/pss_state.py @@ -12,8 +12,8 @@ This mirrors the simple PV-wrapper pattern used by (prohibited, no alarm) so the simulated mount path runs without hardware. """ -from aare.common.beamline import MXBeamline -from aare.common.logger_config import setup_logger +from aarecommon.config.logger import setup_logger +from aarecommon.models.beamline import MXBeamline from epics import PV logger = setup_logger("aareDAQ") diff --git a/src/aare/devices/smargon.py b/src/aare/devices/smargon.py index 131ab0a9..3465c113 100644 --- a/src/aare/devices/smargon.py +++ b/src/aare/devices/smargon.py @@ -1,11 +1,11 @@ from enum import Enum -from time import sleep, time, perf_counter +from time import perf_counter, sleep, time import requests - -from aare.common.beamline import MXBeamline, cfg_get -from aare.common.coordinate import SmargonCoordinate, Coordinate, AerotechCoordinate -from aare.common.exception_handler import SmargonCommunicationError +from aarecommon.config.beamline import cfg_get +from aarecommon.errors.exception_handler import SmargonCommunicationError +from aarecommon.math.coordinate import AerotechCoordinate, Coordinate, SmargonCoordinate +from aarecommon.models.beamline import MXBeamline class SmargonMode(Enum): @@ -24,10 +24,14 @@ class Smargon(object): def __init__(self, bl: MXBeamline): if bl == MXBeamline.X06DA: self.__simulated = False - self.__base = cfg_get("daq.hardware.smargon_url", "http://x06da-smargopolo.psi.ch:3000") + self.__base = cfg_get( + "daq.hardware.smargon_url", "http://x06da-smargopolo.psi.ch:3000" + ) elif bl == MXBeamline.X10SA: self.__simulated = False - self.__base = cfg_get("daq.hardware.smargon_url", "http://x10sa-smargopolo.psi.ch:3000") + self.__base = cfg_get( + "daq.hardware.smargon_url", "http://x10sa-smargopolo.psi.ch:3000" + ) elif bl == MXBeamline.X06SA: raise NotImplementedError("Not implemented smargon url for X06SA") elif bl == MXBeamline.SIMULATED: @@ -153,15 +157,17 @@ class Smargon(object): return self.__pos_aero acs = self.gonget("readbackAEROTECH") - return AerotechCoordinate(at_mm = Coordinate(x=acs["GMX"], y = acs["GMY"], z = acs["GMZ"]), - omega_deg = acs["GMU"]) + return AerotechCoordinate( + at_mm=Coordinate(x=acs["GMX"], y=acs["GMY"], z=acs["GMZ"]), + omega_deg=acs["GMU"], + ) @property def target(self) -> SmargonCoordinate: if self.__simulated: return self.__pos - scs = self.gonget("targetSCS") #targetAEROTECH, #targetOMEGA + scs = self.gonget("targetSCS") # targetAEROTECH, #targetOMEGA return SmargonCoordinate( sh_mm=Coordinate(x=scs["SHX"], y=scs["SHY"], z=scs["SHZ"]), phi_deg=scs["PHI"], @@ -191,9 +197,11 @@ class Smargon(object): if self.__simulated: return self.__pos_aero - acs = self.gonget("targetAEROTECH") #targetAEROTECH, #targetOMEGA - return AerotechCoordinate(at_mm = Coordinate(x=acs["GMX"], y = acs["GMY"], z = acs["GMZ"]), - omega_deg = acs["GMU"]) + acs = self.gonget("targetAEROTECH") # targetAEROTECH, #targetOMEGA + return AerotechCoordinate( + at_mm=Coordinate(x=acs["GMX"], y=acs["GMY"], z=acs["GMZ"]), + omega_deg=acs["GMU"], + ) @target_aerotech.setter def target_aerotech(self, coord: AerotechCoordinate): @@ -211,7 +219,6 @@ class Smargon(object): if target_string: self.gonput(f"targetAEROTECH?{target_string}") - def wait(self, timeout=60.0, tol=0.01, poll_time=0.01): target = self.target timeout = timeout + time() @@ -236,4 +243,3 @@ if __name__ == "__main__": print(smargon.mode) smargon.initialize(timeout=30.0) print(smargon.mode) - diff --git a/src/aare/devices/tell_backend.py b/src/aare/devices/tell_backend.py index b97fe4f4..4cccfa35 100644 --- a/src/aare/devices/tell_backend.py +++ b/src/aare/devices/tell_backend.py @@ -2,17 +2,15 @@ import json import re import time from typing import Any, Callable, Protocol -from requests.exceptions import HTTPError 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, cfg_get # noqa: F401 +from aarecommon.config.beamline import cfg_get +from aarecommon.config.logger import setup_logger +from aarecommon.errors.exception_handler import TellCommunicationError +from aarecommon.models.beamline import MXBeamline from pshell import PShellClient - +from requests.exceptions import HTTPError logger = setup_logger("aareDAQ") @@ -35,35 +33,27 @@ POSITION_HEATER = "pHeatB" class TellBackend(Protocol): @property - def url(self) -> str | None: - ... + def url(self) -> str | None: ... - def get_state(self) -> str: - ... + def get_state(self) -> str: ... - def get_result(self, command_id: int = -1): - ... + def get_result(self, command_id: int = -1): ... - def wait_state(self, state: str, timeout: float) -> None: - ... + def wait_state(self, state: str, timeout: float) -> None: ... - def wait_state_not(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 wait_events(self, events: dict[str, Any], timeout: float): ... - def eval(self, expr: str): - ... + def eval(self, expr: str): ... - def start_eval(self, expr: str) -> int: - ... + def start_eval(self, expr: str) -> int: ... - def run(self, path: str, pars: list[str] | None = None, background: bool = False) -> None: - ... + def run( + self, path: str, pars: list[str] | None = None, background: bool = False + ) -> None: ... - def abort(self) -> None: - ... + def abort(self) -> None: ... class PShellTellBackend: @@ -98,9 +88,11 @@ class PShellTellBackend: def _resolve_url(bl: MXBeamline) -> str: beamline = bl.value.lower() if bl == MXBeamline.X06DA: - return cfg_get("daq.hardware.tell_url",f"http://{beamline}-tell.psi.ch:22222") + return cfg_get( + "daq.hardware.tell_url", f"http://{beamline}-tell.psi.ch:22222" + ) if bl == MXBeamline.X10SA: - return cfg_get("daq.hardware.tell_url","http://PC17488:22222") + return cfg_get("daq.hardware.tell_url", "http://PC17488:22222") if bl == MXBeamline.X06SA: raise NotImplementedError(f"TellClient not implemented for {beamline}") if bl == MXBeamline.SIMULATED: @@ -150,12 +142,15 @@ class PShellTellBackend: 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: + 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 @@ -264,7 +259,9 @@ class SimTellBackend: return str(self._current_mA) if expr.startswith("smart_magnet.set_current("): - match = re.match(r"smart_magnet\.set_current\(([-+]?\d+(?:\.\d+)?)\)&", expr) + match = re.match( + r"smart_magnet\.set_current\(([-+]?\d+(?:\.\d+)?)\)&", expr + ) if match: self._current_mA = float(match.group(1)) return None @@ -325,7 +322,9 @@ class SimTellBackend: self._set_ready_soon() return cmd_id - def run(self, path: str, pars: list[str] | None = None, background: bool = False) -> None: + def run( + self, path: str, pars: list[str] | None = None, background: bool = False + ) -> None: _ = background if path == "data/set_samples_info" and pars: @@ -350,7 +349,9 @@ class SimTellBackend: class LazyTellBackend: - def __init__(self, factory: Callable[[], TellBackend], *, retry_interval_s: float = 2.0): + 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) @@ -362,7 +363,10 @@ class LazyTellBackend: return self._backend now = time.monotonic() - if now - self._last_attempt_ts < self._retry_interval_s and self._last_error is not None: + 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 @@ -406,8 +410,10 @@ class LazyTellBackend: 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: + 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 + self._get_backend().abort() diff --git a/src/aare/devices/tell_client.py b/src/aare/devices/tell_client.py index 5c06354b..c9086bab 100755 --- a/src/aare/devices/tell_client.py +++ b/src/aare/devices/tell_client.py @@ -4,29 +4,30 @@ import re from enum import Enum from typing import List -from aare.common.beamline import MXBeamline -from aare.common.models import ( - PuckLoadedInfo, +from aarecommon.config.logger import setup_logger +from aarecommon.errors.exception_handler import ( + ManualMountException, + MountingFailed, + SmartMagnetFaultException, + TellCommandWhileBusyException, + TellCommunicationError, + TellConnectionException, +) +from aarecommon.models.beamline import MXBeamline +from aarecommon.models.models import ( DewarAddress, + PuckLoadedInfo, SampleDewarAddress, ) from aareDB import PuckWithTellPosition from aare.devices.tell_backend import ( - TellBackend, - SimTellBackend, + POSITION_COLD, LazyTellBackend, PShellTellBackend, - POSITION_COLD, + SimTellBackend, + TellBackend, ) -from aare.common.exception_handler import ( - SmartMagnetFaultException, - TellConnectionException, - MountingFailed, - TellCommandWhileBusyException, - ManualMountException, TellCommunicationError -) -from aare.common.logger_config import setup_logger logger = setup_logger("aareDAQ") @@ -58,6 +59,7 @@ class TellEventValueEnum(Enum): class TellClient: """High-level Tell robot API using a pluggable backend""" + def __init__(self, bl: MXBeamline, backend: TellBackend | None = None): self.__beamline = bl self.backend = backend or PShellTellBackend(bl) @@ -83,12 +85,12 @@ class TellClient: 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""" + and raises an exception if the robot is not ready""" 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""" + and raises an exception if the robot is busy""" self.backend.wait_state_not("Busy", timeout=timeout) state = self.get_state() if state != "Ready": @@ -100,7 +102,7 @@ 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 """ + :param value""" self.backend.eval("in_mount_position = " + str(value) + "&") def _eval_bool(self, expr: str) -> bool: @@ -171,7 +173,7 @@ class TellClient: def set_samples_info(self, info: List[PuckWithTellPosition]): """sets the samples in the robot dewar based on the given list of PuckWithTellPosition objects - and runs set_sample_info in the background""" + and runs set_sample_info in the background""" j = [] for x in info: @@ -218,7 +220,7 @@ class TellClient: def estimate_mounting_time(self, segment) -> int: """Adds additional time if cooling/drying is expected based on requested segment, current sample segment and gripper position. - :param segment: any - however valid segment ABCDEFX """ + :param segment: any - however valid segment ABCDEFX""" try: current_mounted = self.get_mounted_sample() gripper_in_cold = self.is_in_cold() @@ -271,7 +273,7 @@ class TellClient: event, value = self.backend.wait_events( { "state": None, - "Motion Task": None,#"dry", + "Motion Task": None, # "dry", "Gripper detection": None, "Motion Sync": "Robot Clear after mount", }, @@ -279,51 +281,79 @@ class TellClient: ) logger.info(f"event: {event} occurred with value: {value}") if event is None or event == "state": - state_value = str(value).strip().strip('"\'').lower() + state_value = str(value).strip().strip("\"'").lower() if state_value == "busy": - logger.warning("got busy response from robot, waiting for mount to complete") - logger.info(f"event: {event} occurred with value: {value}, checking command completed okay") + logger.warning( + "got busy response from robot, waiting for mount to complete" + ) + logger.info( + f"event: {event} occurred with value: {value}, checking command completed okay" + ) msg = self.check_command_ok( timeout=wait_timeout, msg=f"Mount {segment}{puck}-{sample}: " ) logger.info(f"Check command okay response: {msg}") if self.get_mounted_sample() is None: - logger.warning("Mount command completed but robot reports no sample mounted") + logger.warning( + "Mount command completed but robot reports no sample mounted" + ) return TellEventValueEnum.NO_PIN_IN_GRIPPER return TellEventValueEnum.SUCCESS elif ( - event == TellEventTypeEnum.GIPPER_DETECTION.value - and value == TellEventValueEnum.NO_PIN_IN_GRIPPER.value + event == TellEventTypeEnum.GIPPER_DETECTION.value + and value == TellEventValueEnum.NO_PIN_IN_GRIPPER.value ): logger.info(f"{TellEventValueEnum.NO_PIN_IN_GRIPPER.value}") return TellEventValueEnum.NO_PIN_IN_GRIPPER - elif event == TellEventTypeEnum.GIPPER_DETECTION.value and value == TellEventValueEnum.PIN_STILL_IN_GRIPPER.value: + elif ( + event == TellEventTypeEnum.GIPPER_DETECTION.value + and value == TellEventValueEnum.PIN_STILL_IN_GRIPPER.value + ): logger.info(f"{TellEventValueEnum.PIN_STILL_IN_GRIPPER.value}") return TellEventValueEnum.PIN_STILL_IN_GRIPPER - elif event == TellEventTypeEnum.GIPPER_DETECTION.value and value == TellEventValueEnum.PIN_IS_LOST_GRIPPER.value: + elif ( + event == TellEventTypeEnum.GIPPER_DETECTION.value + and value == TellEventValueEnum.PIN_IS_LOST_GRIPPER.value + ): logger.info(f"{TellEventValueEnum.PIN_IS_LOST_GRIPPER.value}") return TellEventValueEnum.PIN_IS_LOST_GRIPPER elif event == TellEventTypeEnum.GIPPER_DETECTION.value: return TellEventTypeEnum.GIPPER_DETECTION - elif event == TellEventTypeEnum.MOTION_TASK.value and value == TellEventValueEnum.DRY.value: + elif ( + event == TellEventTypeEnum.MOTION_TASK.value + and value == TellEventValueEnum.DRY.value + ): logger.info(f"{TellEventValueEnum.DRY.value}") return TellEventValueEnum.DRY - elif event == TellEventTypeEnum.MOTION_TASK.value and value == TellEventValueEnum.COLD.value: + elif ( + event == TellEventTypeEnum.MOTION_TASK.value + and value == TellEventValueEnum.COLD.value + ): logger.info(f"{TellEventValueEnum.COLD.value}") - logger.info("As cold, likely robot is cooling from previous mount/dry " - "DAQ will block until command is complete") + logger.info( + "As cold, likely robot is cooling from previous mount/dry " + "DAQ will block until command is complete" + ) self.check_command_ok( timeout=wait_timeout, msg=f"Mount {segment}{puck}-{sample}: " ) return TellEventValueEnum.COLD - elif event == TellEventTypeEnum.MOTION_TASK.value and value == TellEventValueEnum.UNKNOWN.value: + elif ( + event == TellEventTypeEnum.MOTION_TASK.value + and value == TellEventValueEnum.UNKNOWN.value + ): logger.info(f"{TellEventValueEnum.UNKNOWN.value}") return TellEventValueEnum.UNKNOWN - elif event == TellEventTypeEnum.MOTION_SYNC.value and value == TellEventValueEnum.ROBOT_CLEAR_AFTER_MOUNT.value: + elif ( + event == TellEventTypeEnum.MOTION_SYNC.value + and value == TellEventValueEnum.ROBOT_CLEAR_AFTER_MOUNT.value + ): logger.info(f"{TellEventValueEnum.ROBOT_CLEAR_AFTER_MOUNT.value}") return TellEventValueEnum.ROBOT_CLEAR_AFTER_MOUNT else: - logger.info(f"Unexpected event: {event} occurred with value: {value}") + logger.info( + f"Unexpected event: {event} occurred with value: {value}" + ) logger.info("Checking command completed okay anyway") self.check_command_ok( timeout=wait_timeout, msg=f"Mount {segment}{puck}-{sample}: " @@ -340,7 +370,10 @@ class TellClient: raise except Exception as e: logger.error(f"Exception occurred: {e}") - raise TellCommunicationError(message=f"Error during mount {segment}{puck}-{sample}: {e}", critical=True) + raise TellCommunicationError( + message=f"Error during mount {segment}{puck}-{sample}: {e}", + critical=True, + ) def unmount(self, force=False, wait=False, timeout=360.0): if self.is_busy(): @@ -353,7 +386,9 @@ class TellClient: self.check_command_ok(timeout=timeout, msg="Unmount message: ") except MountingFailed as e: result = self.get_result(self._last_cmd_id) - logger.error(f"Unmount failed with status '{result.get('status')}' and payload: {result}") + logger.error( + f"Unmount failed with status '{result.get('status')}' and payload: {result}" + ) raise return self._last_cmd_id @@ -426,7 +461,7 @@ class TellClient: def get_robot_status(self): status = self.backend.eval("robot.take()&") logger.debug(f"robot status: {status}") - #return eval(status) + # return eval(status) return ast.literal_eval(status) def get_detected_pucks(self) -> List[PuckLoadedInfo]: @@ -483,8 +518,10 @@ class TellClient: def is_busy(self): 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): - #Not sure why unused, potentially can remove them + def check_smart_magnet_mounted( + self, timeout: float = 10.0, idle_time: float = 1.0, interval: float = 0.1 + ): + # Not sure why unused, potentially can remove them _ = (timeout, idle_time, interval) initial_state = self.backend.eval("smart_magnet.state&") @@ -506,14 +543,18 @@ class TellClient: 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( + "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") if self.get_mounted_sample(): - logger.error("Check mount: No sample detected, but robot thinks is mounted") + logger.error( + "Check mount: No sample detected, but robot thinks is mounted" + ) raise SmartMagnetFaultException return False elif state == "Paused": @@ -538,32 +579,35 @@ def make_tell_client(bl: MXBeamline) -> TellClient: ) return TellClient(bl, backend=backend) + if __name__ == "__main__": - from aare.common.beamline import mx_beamline from datetime import datetime, timezone + + from aarecommon.config.beamline import mx_beamline + bl = mx_beamline() tell_client = make_tell_client(bl) - #tell_client.toggle_blower() - #tell_client.check_enable_motion() + # tell_client.toggle_blower() + # tell_client.check_enable_motion() print("status ", tell_client.get_robot_status()) - print("dry mount count: ", tell_client.get_setting('dry_mount_counter')) + print("dry mount count: ", tell_client.get_setting("dry_mount_counter")) - ts = float(tell_client.get_setting('dry_timestamp')) + ts = float(tell_client.get_setting("dry_timestamp")) print("dry timestape: ", ts) past = datetime.fromtimestamp(ts, tz=timezone.utc) now = datetime.now(timezone.utc) seconds_ago = int((now - past).total_seconds()) print(seconds_ago) - print("door closer :", tell_client.backend.eval('is_door_closed()&')) - print("manual mode: ", tell_client.backend.eval('is_manual_mode()&')) + print("door closer :", tell_client.backend.eval("is_door_closed()&")) + print("manual mode: ", tell_client.backend.eval("is_manual_mode()&")) print("position :", tell_client.get_robot_status()["pos"]) state = tell_client.get_robot_state() manual_mode = tell_client.is_manual_mode() print("state: ", state) print("is manual mode True: ", manual_mode == True) - print(tell_client.backend.eval('is_manual_mode()&')) + print(tell_client.backend.eval("is_manual_mode()&")) print(tell_client.is_door_closed()) - #print("release safety: ", tell_client.backend.eval('release_safety()&')) + # print("release safety: ", tell_client.backend.eval('release_safety()&')) # time.sleep(5) - #tell_client.blower_off() \ No newline at end of file + # tell_client.blower_off() diff --git a/src/aare/devices/zmq_client.py b/src/aare/devices/zmq_client.py index c4a7bf6b..9494fbaf 100644 --- a/src/aare/devices/zmq_client.py +++ b/src/aare/devices/zmq_client.py @@ -11,9 +11,8 @@ from typing import Optional import cv2 import numpy as np import zmq - -from aare.common.beamline import MXBeamline -from aare.common.logger_config import setup_logger +from aarecommon.config.logger import setup_logger +from aarecommon.models.beamline import MXBeamline logger = setup_logger("aareDAQ") @@ -21,7 +20,7 @@ logger = setup_logger("aareDAQ") class ZMQCameraClient: """ ZMQ camera client that requests the latest image from a ZMQ stream. - + Uses a SUB socket with a short timeout to grab the most recent frame, rather than maintaining a continuous subscription. This is suitable for on-demand image retrieval in the DAQ server. @@ -183,4 +182,4 @@ class ZMQCameraClient: @property def url(self) -> str: - return self.__zmq_url if self.__zmq_url else "simulated" \ No newline at end of file + return self.__zmq_url if self.__zmq_url else "simulated" diff --git a/src/aare/gui/auth.py b/src/aare/gui/auth.py index a9f6514e..7355a9a1 100644 --- a/src/aare/gui/auth.py +++ b/src/aare/gui/auth.py @@ -1,21 +1,20 @@ import json import subprocess + import jwt +from aarecommon.config.logger import setup_logger +from aarecommon.models.auth import get_user +from aarecommon.models.models import TokenData -from aare.common.models import TokenData -from aare.common.auth_models import get_user -from aare.common.logger_config import setup_logger - -logger = setup_logger('aareGUI') +logger = setup_logger("aareGUI") def auth(base_url: str | None, cert_path: str | None) -> str: curr_user = get_user() if base_url is None: - token_data = TokenData(sub=curr_user, - staff=True, - session=15, - pgroups=["p16371", "p22233"]) + token_data = TokenData( + sub=curr_user, staff=True, session=15, pgroups=["p16371", "p22233"] + ) return jwt.encode(token_data.model_dump(), "ABC123") # Single call: Kerberos SPNEGO through Apache, which proxies to the DAQ server. @@ -24,7 +23,7 @@ def auth(base_url: str | None, cert_path: str | None) -> str: cacert = f"{cert_path}" try: token_result = subprocess.run( - ['curl', '-s', '--cacert', cacert, '--negotiate', '-u', ':', url, "-XPOST"], + ["curl", "-s", "--cacert", cacert, "--negotiate", "-u", ":", url, "-XPOST"], capture_output=True, text=True, timeout=120.0, @@ -46,12 +45,12 @@ def auth(base_url: str | None, cert_path: str | None) -> str: ) from e except Exception as e: logger.error(f"Token request curl error: {e}") - raise RuntimeError( - "Cannot reach AareDAQ server (unknown error). " - ) from e + raise RuntimeError("Cannot reach AareDAQ server (unknown error). ") from e if token_result.returncode != 0: - logger.error(f"Token curl exited {token_result.returncode}. stderr: {token_result.stderr[:500]}") + logger.error( + f"Token curl exited {token_result.returncode}. stderr: {token_result.stderr[:500]}" + ) raise RuntimeError( f"Token request failed (curl exit {token_result.returncode}). " "Check Kerberos ticket is valid (kinit) and server is reachable." @@ -68,10 +67,12 @@ def auth(base_url: str | None, cert_path: str | None) -> str: token = response_json.get("access_token") if not token or not isinstance(token, str): - logger.error(f"Missing access_token. Keys: {list(response_json.keys())}. Server response {response_json}") + logger.error( + f"Missing access_token. Keys: {list(response_json.keys())}. Server response {response_json}" + ) raise RuntimeError( "Authentication failed (missing token in server response). " "The server may be starting up." ) - return token \ No newline at end of file + return token diff --git a/src/aare/gui/gui.py b/src/aare/gui/gui.py index bd50867e..ac5189be 100644 --- a/src/aare/gui/gui.py +++ b/src/aare/gui/gui.py @@ -2,35 +2,37 @@ import os import sys import traceback +from aarecommon.config.beamline import cfg_get, mx_beamline +from aarecommon.config.logger import setup_logger +from aarecommon.models.beamline import MXBeamline from PySide6 import QtGui -from PySide6.QtCore import QCommandLineParser, QCommandLineOption -from PySide6.QtWidgets import QApplication, QMessageBox, QSplashScreen, QProgressBar +from PySide6.QtCore import QCommandLineOption, QCommandLineParser +from PySide6.QtWidgets import QApplication, QMessageBox, QProgressBar, QSplashScreen -from aare.common.logger_config import setup_logger +from aare.gui.auth import auth from aare.gui.main_window import MainWindow from aare.gui.widgets.splash_screen import LoadingSplashScreen -from aare.common.beamline import MXBeamline, mx_beamline, cfg_get -from aare.gui.auth import auth logger = setup_logger("aareGUI") + def main(): """Wrapped gui as main function to make tests easier""" splash = None try: - #define application + # define application app = QApplication(sys.argv) app.setApplicationName("AareGUI") app.setApplicationVersion("0.3.1") app.setOrganizationName("PSI") app.setOrganizationDomain("psi.ch") - #set icon + # set icon basedir = os.path.dirname(__file__) icon_path = os.path.join(basedir, "graphics/aaregui_logo.svg") app.setWindowIcon(QtGui.QIcon(icon_path)) - #show splash screen + # show splash screen banner_path = os.path.join(basedir, "graphics/aare_banner.png") splash_pix = QtGui.QPixmap(banner_path) splash = LoadingSplashScreen(splash_pix) @@ -42,30 +44,60 @@ def main(): parser.setApplicationDescription("PSI AareGUI") parser.addHelpOption() # Adds --help option parser.addVersionOption() # Adds --version option - #TODO if zmq and pred stream come from same source, do not need images from both streams, can combine + # TODO if zmq and pred stream come from same source, do not need images from both streams, can combine match mx_beamline(): case MXBeamline.X06DA: - default_url = cfg_get("gui.daq.daq_url", "https://mx-x06da-queue-01.psi.ch") - default_cert_path = cfg_get("gui.daq.cert_path", "/sls/x06da/misc/.cert/6d.crt") - default_zmq_addr = cfg_get("gui.cameras.sample_camera_zmq_url", "tcp://x06da-pserv-01:9089") - default_pred_zmq_addr = cfg_get("gui.cameras.prediction_zmq_url", "tcp://mx-ml:9091") - default_beamline_cam_addr = cfg_get("gui.cameras.beamline_camera_url", "x06da-axis-1.psi.ch") - default_gonio_cam_addr = cfg_get("gui.cameras.gonio_camera_url", "axis-accc8ed2972e.psi.ch") + default_url = cfg_get( + "gui.daq.daq_url", "https://mx-x06da-queue-01.psi.ch" + ) + default_cert_path = cfg_get( + "gui.daq.cert_path", "/sls/x06da/misc/.cert/6d.crt" + ) + default_zmq_addr = cfg_get( + "gui.cameras.sample_camera_zmq_url", "tcp://x06da-pserv-01:9089" + ) + default_pred_zmq_addr = cfg_get( + "gui.cameras.prediction_zmq_url", "tcp://mx-ml:9091" + ) + default_beamline_cam_addr = cfg_get( + "gui.cameras.beamline_camera_url", "x06da-axis-1.psi.ch" + ) + default_gonio_cam_addr = cfg_get( + "gui.cameras.gonio_camera_url", "axis-accc8ed2972e.psi.ch" + ) default_gonio_camera_id = int(cfg_get("gui.cameras.gonio_camera_id", 3)) case MXBeamline.X10SA: - default_url = cfg_get("gui.daq.daq_url", "https://mx-x10sa-queue-01.psi.ch") - default_cert_path = cfg_get("gui.daq.cert_path", "/sls/x10sa/misc/.cert/10s.crt") - default_zmq_addr = cfg_get("gui.cameras.sample_camera_zmq_url", "tcp://x10sa-spark-01:9091") - default_pred_zmq_addr = cfg_get("gui.cameras.prediction_zmq_url", "tcp://x10sa-spark-01:9091") - default_beamline_cam_addr = cfg_get("gui.cameras.beamline_camera_url", "axis-accc8eb02488.psi.ch") - default_gonio_cam_addr = cfg_get("gui.cameras.gonio_camera_url", "axis-accc8ea5e463.psi.ch") + default_url = cfg_get( + "gui.daq.daq_url", "https://mx-x10sa-queue-01.psi.ch" + ) + default_cert_path = cfg_get( + "gui.daq.cert_path", "/sls/x10sa/misc/.cert/10s.crt" + ) + default_zmq_addr = cfg_get( + "gui.cameras.sample_camera_zmq_url", "tcp://x10sa-spark-01:9091" + ) + default_pred_zmq_addr = cfg_get( + "gui.cameras.prediction_zmq_url", "tcp://x10sa-spark-01:9091" + ) + default_beamline_cam_addr = cfg_get( + "gui.cameras.beamline_camera_url", "axis-accc8eb02488.psi.ch" + ) + default_gonio_cam_addr = cfg_get( + "gui.cameras.gonio_camera_url", "axis-accc8ea5e463.psi.ch" + ) default_gonio_camera_id = int(cfg_get("gui.cameras.gonio_camera_id", 1)) case MXBeamline.X06SA: - default_url = cfg_get("gui.daq.daq_url", "https://mx-x06sa-queue-01.psi.ch") - default_cert_path = cfg_get("gui.daq.cert_path", "/sls/x06sa/misc/.cert/6s.crt") + default_url = cfg_get( + "gui.daq.daq_url", "https://mx-x06sa-queue-01.psi.ch" + ) + default_cert_path = cfg_get( + "gui.daq.cert_path", "/sls/x06sa/misc/.cert/6s.crt" + ) default_zmq_addr = cfg_get("gui.cameras.sample_camera_zmq_url", "") default_pred_zmq_addr = cfg_get("gui.cameras.prediction_zmq_url", "") - default_beamline_cam_addr = cfg_get("gui.cameras.beamline_camera_url", "") + default_beamline_cam_addr = cfg_get( + "gui.cameras.beamline_camera_url", "" + ) default_gonio_cam_addr = cfg_get("gui.cameras.gonio_camera_url", "") default_gonio_camera_id = int(cfg_get("gui.cameras.gonio_camera_id", 1)) case _: @@ -73,37 +105,46 @@ def main(): default_cert_path = cfg_get("gui.daq.cert_path", "") default_zmq_addr = cfg_get("gui.cameras.sample_camera_zmq_url", "") default_pred_zmq_addr = cfg_get("gui.cameras.prediction_zmq_url", "") - default_beamline_cam_addr = cfg_get("gui.cameras.beamline_camera_url", "") + default_beamline_cam_addr = cfg_get( + "gui.cameras.beamline_camera_url", "" + ) default_gonio_cam_addr = cfg_get("gui.cameras.gonio_camera_url", "") default_gonio_camera_id = int(cfg_get("gui.cameras.gonio_camera_id", 1)) # Add custom options as needed - urlOption = QCommandLineOption(["u", "aaredaq-url"], - "Base AareDAQ URL", - "url", default_url) + urlOption = QCommandLineOption( + ["u", "aaredaq-url"], "Base AareDAQ URL", "url", default_url + ) parser.addOption(urlOption) - certPath = QCommandLineOption(["c", "aaredaq-cert-path"], - "Server Certificate Path (for self-signed certificates)", - "cert", default_cert_path) + certPath = QCommandLineOption( + ["c", "aaredaq-cert-path"], + "Server Certificate Path (for self-signed certificates)", + "cert", + default_cert_path, + ) parser.addOption(certPath) - defaultImage = QCommandLineOption(["i", "image"], - "Default image to display in absence of the ZMQ stream", - "image") + defaultImage = QCommandLineOption( + ["i", "image"], + "Default image to display in absence of the ZMQ stream", + "image", + ) parser.addOption(defaultImage) - cameraZeroMQ = QCommandLineOption(["s", "sample-camera-zmq"], - "Sample camera ZeroMQ URL", - "sample-camera-zmq", - default_zmq_addr) + cameraZeroMQ = QCommandLineOption( + ["s", "sample-camera-zmq"], + "Sample camera ZeroMQ URL", + "sample-camera-zmq", + default_zmq_addr, + ) parser.addOption(cameraZeroMQ) predZmqOption = QCommandLineOption( ["p", "pred-zmq"], "Prediction ZeroMQ URL (PUB) to subscribe to (e.g. tcp://mx-ml:9091)", "pred-zmq", - default_pred_zmq_addr + default_pred_zmq_addr, ) parser.addOption(predZmqOption) @@ -151,19 +192,21 @@ def main(): QMessageBox.critical( None, "Authentication Error", - f"Cannot connect to AareDAQ server:\n{str(e)}\n\n Please check the server is running and your network connection." + f"Cannot connect to AareDAQ server:\n{str(e)}\n\n Please check the server is running and your network connection.", ) sys.exit(1) splash.set_progress(90, "Loading Main Window...") - win = MainWindow(base_url=base_url, - token=token, - default_image=default_image, - zmq_addr=zmq_addr, - pred_zmq_addr=pred_zmq_addr, - beamline_cam_addr = default_beamline_cam_addr, - gonio_cam_addr = default_gonio_cam_addr, - gonio_cam_id = default_gonio_camera_id) + win = MainWindow( + base_url=base_url, + token=token, + default_image=default_image, + zmq_addr=zmq_addr, + pred_zmq_addr=pred_zmq_addr, + beamline_cam_addr=default_beamline_cam_addr, + gonio_cam_addr=default_gonio_cam_addr, + gonio_cam_id=default_gonio_camera_id, + ) splash.set_progress(100, "Ready") splash.finish(win) @@ -181,12 +224,13 @@ def main(): "Fatal Error", f"An error occurred during startup. See console for details." f"\nPlease check the server is running and your network connection." - f"\n\n{str(e)}\n\n" + f"\n\n{str(e)}\n\n", ) except: pass sys.exit(1) + if __name__ == "__main__": main() diff --git a/src/aare/gui/main_window.py b/src/aare/gui/main_window.py index 80604da9..fb37068a 100644 --- a/src/aare/gui/main_window.py +++ b/src/aare/gui/main_window.py @@ -1,99 +1,106 @@ import time import jwt -from PySide6.QtCore import Qt, Slot, Signal, QTimer, QSettings, QEvent -from PySide6.QtGui import QAction, QPixmap, QKeySequence, QGuiApplication, QActionGroup +from aarecommon.config.logger import setup_logger +from aarecommon.math.coordinate import Coordinate, SmargonCoordinate +from aarecommon.math.diffraction_geometry import DiffractionGeometry +from aarecommon.math.sample_geometry import SampleGeometryModel + +# Common imports +from aarecommon.models.auth import BatonStatus +from aarecommon.models.models import ( + BeamlineStateEnum, + DAQStatusModel, + SampleShortInfoList, + TokenData, +) +from PySide6.QtCore import QEvent, QSettings, Qt, QTimer, Signal, Slot +from PySide6.QtGui import QAction, QActionGroup, QGuiApplication, QKeySequence, QPixmap from PySide6.QtWidgets import ( - QMainWindow, - QWidget, - QHBoxLayout, - QVBoxLayout, - QMessageBox, QDockWidget, - QTabWidget, + QHBoxLayout, + QMainWindow, + QMessageBox, QStackedWidget, + QTabWidget, + QVBoxLayout, + QWidget, ) -#Common imports -from aare.common.auth_models import BatonStatus -from aare.common.coordinate import Coordinate, SmargonCoordinate -from aare.common.diffraction_geometry import DiffractionGeometry -from aare.common.logger_config import setup_logger -from aare.common.models import SampleShortInfoList, TokenData, DAQStatusModel, BeamlineStateEnum -from aare.common.sample_geometry import SampleGeometryModel - -#Gui Models +# Gui Models from aare.gui.models.gui_state_manager import UIStateManager -from aare.gui.styles import build_app_stylesheet, THEME_ORIGINAL, THEME_PORTRAIT - -#panels -from aare.gui.panels.LogPanel import LogDock +from aare.gui.panels.automation_panel import AutomationProgressWidget +from aare.gui.panels.axis_video_panel import AxisVideoPanel from aare.gui.panels.beamline_controls import BeamlineControls +from aare.gui.panels.beamline_recovery_panel import BeamlineRecoveryDialog from aare.gui.panels.beamline_state_panel import BeamlineStatePanel +from aare.gui.panels.compact_automation_panel import CompactAutomationPanel from aare.gui.panels.data_collection_settings import DataCollectionSettings from aare.gui.panels.developer_help_dialog import DeveloperHelpDialog -from aare.gui.panels.beamline_recovery_panel import BeamlineRecoveryDialog +from aare.gui.panels.face_detection_panel import FaceDetectionPanel +from aare.gui.panels.fluorescence_panel import FluorescencePanel from aare.gui.panels.local_contact_panel import LocalContactDialog + +# panels +from aare.gui.panels.LogPanel import LogDock from aare.gui.panels.loop_centering_panel import LoopCenteringPanel from aare.gui.panels.manual_sample_panel import ManualSamplePanel +from aare.gui.panels.portrait_mode import PortraitModePanel from aare.gui.panels.prediction_metrics_panel import PredictionMetricsPanel from aare.gui.panels.reference_tools_panel import ReferenceToolsPanel from aare.gui.panels.sample_queue_panel import SampleQueuePanel +from aare.gui.panels.smargon_trace_panel import SmargonTracePanel from aare.gui.panels.target_stability_panel import TargetStabilityPanel from aare.gui.panels.tell_sample_panel import TellSamplePanel -from aare.gui.panels.face_detection_panel import FaceDetectionPanel -from aare.gui.panels.smargon_trace_panel import SmargonTracePanel -from aare.gui.panels.axis_video_panel import AxisVideoPanel -from aare.gui.panels.fluorescence_panel import FluorescencePanel -from aare.gui.panels.automation_panel import AutomationProgressWidget -from aare.gui.panels.compact_automation_panel import CompactAutomationPanel -from aare.gui.panels.portrait_mode import PortraitModePanel -#Scan Logic +# Scan Logic from aare.gui.scan_logic.raster_grid_manager import RasterGridManager from aare.gui.scan_logic.rotation_scan_manager import RotationScanManager from aare.gui.scan_logic.sample_mount_logic import SampleMountLogic +from aare.gui.styles import THEME_ORIGINAL, THEME_PORTRAIT, build_app_stylesheet -#Threads +# Threads from aare.gui.threads.axis_video_thread import VideoThread -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.threads.prediction_subscriber import PredictionSubscriber +from aare.gui.tutorials.controls_help_dialog import ControlsHelpDialog -#Tutorials +# Tutorials from aare.gui.tutorials.tutorial_actions import TutorialActionExecutor from aare.gui.tutorials.tutorial_manager import TutorialManager +from aare.gui.tutorials.tutorial_registration import register_tutorials from aare.gui.tutorials.tutorial_runtime import DictionaryTextResolver, TutorialEventBus from aare.gui.tutorials.tutorial_targets import MainWindowTutorialTargetResolver -from aare.gui.tutorials.tutorial_registration import register_tutorials from aare.gui.tutorials.tutroial_texts import MANUAL_MOUNT_TUTORIAL -from aare.gui.tutorials.controls_help_dialog import ControlsHelpDialog -from aare.gui.tutorials.tutorial_registration import register_tutorials -#Widgets +# Widgets from aare.gui.widgets.alert_banner import AlertBanner -from aare.gui.widgets.baton_request_dialog import BatonRequestDialog, BatonPendingDialog +from aare.gui.widgets.baton_request_dialog import BatonPendingDialog, BatonRequestDialog +from aare.gui.widgets.busy_overlay import build_busy_overlay_style from aare.gui.widgets.camera_image import SampleCameraImageLabel +from aare.gui.widgets.message_box import precondition_check from aare.gui.widgets.no_wheel_scroll_area import NoWheelScrollArea from aare.gui.widgets.status_bar import StatusBar from aare.gui.widgets.video_image import VideoGraphicsView -from aare.gui.widgets.busy_overlay import build_busy_overlay_style -from aare.gui.widgets.message_box import precondition_check logger = setup_logger("aareGUI") + class MainWindow(QMainWindow): sample_geometry = Signal(SampleGeometryModel) - def __init__(self, base_url: str | None, - token: str, - default_image: str | None, - zmq_addr: str | None, - pred_zmq_addr: str | None, - beamline_cam_addr: str | None, - gonio_cam_addr: str | None, - gonio_cam_id: int | None - ): + def __init__( + self, + base_url: str | None, + token: str, + default_image: str | None, + zmq_addr: str | None, + pred_zmq_addr: str | None, + beamline_cam_addr: str | None, + gonio_cam_addr: str | None, + gonio_cam_id: int | None, + ): super().__init__() self._theme_mode = THEME_ORIGINAL @@ -166,7 +173,7 @@ class MainWindow(QMainWindow): "Authentication Error", "Could not start the GUI because authentication data was invalid.\n\n" "Most commonly the server is not running yet (or is still initialising).\n" - "Please start/restart the server and try again." + "Please start/restart the server and try again.", ) raise @@ -201,17 +208,18 @@ class MainWindow(QMainWindow): detector_description="PILATUS 4", detector_serial_number="1", poni_rot1_rad=-0.001396263, - poni_rot2_rad=-0.003839724 + poni_rot2_rad=-0.003839724, ) - geom = SampleGeometryModel(beam_location_pxl=Coordinate(x=1000, y=1000), - pixel_in_mm=0.001, - aerotech=Coordinate(), - smargon=SmargonCoordinate(sh_mm=Coordinate(), phi_deg=0, chi_deg=0), - omega_deg=0, - beam_size_mm=Coordinate(x=0.01, y=0.01), - aerotech_meas=Coordinate() - ) + geom = SampleGeometryModel( + beam_location_pxl=Coordinate(x=1000, y=1000), + pixel_in_mm=0.001, + aerotech=Coordinate(), + smargon=SmargonCoordinate(sh_mm=Coordinate(), phi_deg=0, chi_deg=0), + omega_deg=0, + beam_size_mm=Coordinate(x=0.01, y=0.01), + aerotech_meas=Coordinate(), + ) self.raster = RasterGridManager(geom=geom) self.rotation = RotationScanManager() @@ -223,10 +231,12 @@ class MainWindow(QMainWindow): self.left_column_layout.setContentsMargins(0, 0, 0, 0) self.left_column_layout.setSpacing(8) - self.data_collection = DataCollectionSettings(s=geom, - parent=self.left_column, - raster_mgr=self.raster, - diffraction=diffraction) + self.data_collection = DataCollectionSettings( + s=geom, + parent=self.left_column, + raster_mgr=self.raster, + diffraction=diffraction, + ) self.loop_centering = LoopCenteringPanel(parent=self.left_column) @@ -254,20 +264,29 @@ class MainWindow(QMainWindow): self.data_collection.set_width, self.loop_centering.sizeHint().width(), self.beamline_state_panel.set_width, - ) + 10 + ) + + 10 ) self.video_tab = QTabWidget(parent=top_widget) - self.sample_camera = SampleCameraImageLabel(geom=geom, raster=self.raster, parent=top_widget, - default_image=default_image) + self.sample_camera = SampleCameraImageLabel( + geom=geom, + raster=self.raster, + parent=top_widget, + default_image=default_image, + ) self.beamline_view = VideoGraphicsView() - self.beamline_view_panel = AxisVideoPanel("Beamline view", self.beamline_view, parent=top_widget) + self.beamline_view_panel = AxisVideoPanel( + "Beamline view", self.beamline_view, parent=top_widget + ) self.beamline_view_panel.refresh_requested.connect(self.refresh_axis_cameras) self.gonio_view = VideoGraphicsView() - self.gonio_view_panel = AxisVideoPanel("Gonio camera", self.gonio_view, parent=top_widget) + self.gonio_view_panel = AxisVideoPanel( + "Gonio camera", self.gonio_view, parent=top_widget + ) self.gonio_view_panel.refresh_requested.connect(self.refresh_axis_cameras) self.beamline_view_container = QWidget(parent=top_widget) @@ -285,7 +304,9 @@ class MainWindow(QMainWindow): self.beamline_view_container, parent=top_widget, ) - self.beamline_combined_panel.refresh_requested.connect(self.refresh_axis_cameras) + self.beamline_combined_panel.refresh_requested.connect( + self.refresh_axis_cameras + ) self.video_tab.addTab(self.sample_camera, "Sample camera") self.video_tab.addTab(self.gonio_view_panel, "Gonio camera") @@ -305,7 +326,9 @@ class MainWindow(QMainWindow): parent=root_widget, default_image=default_image, ) - self.compact_automation_panel = CompactAutomationPanel(self.compact_sample_camera, parent=root_widget) + self.compact_automation_panel = CompactAutomationPanel( + self.compact_sample_camera, parent=root_widget + ) self.compact_automation_page = QWidget(parent=root_widget) self.compact_automation_page.setObjectName("compactAutomationPage") @@ -331,7 +354,9 @@ class MainWindow(QMainWindow): portrait_page_layout = QHBoxLayout(self.portrait_mode_page) portrait_page_layout.setContentsMargins(0, 0, 0, 0) portrait_page_layout.setSpacing(0) - self.portrait_mode_page.setFixedWidth(self.portrait_mode_panel.PORTRAIT_WIDTH + 24) + self.portrait_mode_page.setFixedWidth( + self.portrait_mode_panel.PORTRAIT_WIDTH + 24 + ) portrait_page_layout.addWidget( self.portrait_mode_panel, alignment=Qt.AlignmentFlag.AlignHCenter, @@ -342,7 +367,9 @@ class MainWindow(QMainWindow): self.beamline_controls_scroll = NoWheelScrollArea(top_widget) - self.beamline = BeamlineControls(self.beamline_controls_scroll, staff=self.__decoded_token.staff) + self.beamline = BeamlineControls( + self.beamline_controls_scroll, staff=self.__decoded_token.staff + ) top_widget_layout.addWidget(self.beamline_controls_scroll) self.beamline_controls_scroll.setWidget(self.beamline) self.beamline_controls_scroll.setHorizontalScrollBarPolicy( @@ -354,17 +381,29 @@ class MainWindow(QMainWindow): self.ref_tools_panel = ReferenceToolsPanel(samples=SampleShortInfoList(s=[])) self.job_list_panel = SampleQueuePanel(show_user=self.__decoded_token.staff) - self.compact_automation_panel.play_pause_clicked.connect(self.job_list_panel.run) - self.compact_automation_panel.skip_clicked.connect(self.job_list_panel.skip_current_sample) - self.compact_automation_panel.step_through_toggled.connect(self.job_list_panel.set_step_through) - self.compact_automation_panel.show_full_view_requested.connect(self._return_from_compact_automation_view) - self.compact_automation_panel.annotation_selected.connect(self._handle_compact_annotation) + self.compact_automation_panel.play_pause_clicked.connect( + self.job_list_panel.run + ) + self.compact_automation_panel.skip_clicked.connect( + self.job_list_panel.skip_current_sample + ) + self.compact_automation_panel.step_through_toggled.connect( + self.job_list_panel.set_step_through + ) + self.compact_automation_panel.show_full_view_requested.connect( + self._return_from_compact_automation_view + ) + self.compact_automation_panel.annotation_selected.connect( + self._handle_compact_annotation + ) self.tell_samples_dock = QDockWidget("Sample List", self) self.tell_samples_dock.setObjectName("tell_samples_dock") self.tell_samples_dock.setWidget(self.tell_samples) self.tell_samples_dock.setAllowedAreas(Qt.DockWidgetArea.BottomDockWidgetArea) - self.addDockWidget(Qt.DockWidgetArea.BottomDockWidgetArea, self.tell_samples_dock) + self.addDockWidget( + Qt.DockWidgetArea.BottomDockWidgetArea, self.tell_samples_dock + ) self.ref_tools_dock = QDockWidget("Reference Tools", self) self.ref_tools_dock.setObjectName("ref_tools_dock") @@ -391,21 +430,28 @@ class MainWindow(QMainWindow): self.manual_sample_dock.setObjectName("manual_sample_dock") self.manual_sample_dock.setWidget(self.manual_sample_panel) self.manual_sample_dock.setAllowedAreas(Qt.DockWidgetArea.BottomDockWidgetArea) - self.addDockWidget(Qt.DockWidgetArea.BottomDockWidgetArea, self.manual_sample_dock) + self.addDockWidget( + Qt.DockWidgetArea.BottomDockWidgetArea, self.manual_sample_dock + ) self.automation_progress_panel = AutomationProgressWidget() self.automation_progress_dock = QDockWidget("Automation progress", self) self.automation_progress_dock.setObjectName("automation_progress_dock") self.automation_progress_dock.setWidget(self.automation_progress_panel) - self.automation_progress_dock.setAllowedAreas(Qt.DockWidgetArea.BottomDockWidgetArea) - self.addDockWidget(Qt.DockWidgetArea.BottomDockWidgetArea, self.automation_progress_dock) + self.automation_progress_dock.setAllowedAreas( + Qt.DockWidgetArea.BottomDockWidgetArea + ) + self.addDockWidget( + Qt.DockWidgetArea.BottomDockWidgetArea, self.automation_progress_dock + ) self.face_panel = FaceDetectionPanel() self.face_panel_dock = QDockWidget("Face detection", self) self.face_panel_dock.setObjectName("face_panel_dock") self.face_panel_dock.setWidget(self.face_panel) self.face_panel_dock.setAllowedAreas( - Qt.DockWidgetArea.RightDockWidgetArea | Qt.DockWidgetArea.LeftDockWidgetArea) + Qt.DockWidgetArea.RightDockWidgetArea | Qt.DockWidgetArea.LeftDockWidgetArea + ) self.addDockWidget(Qt.DockWidgetArea.RightDockWidgetArea, self.face_panel_dock) self.face_panel_dock.hide() @@ -414,8 +460,13 @@ class MainWindow(QMainWindow): self.fluor_panel_dock.setObjectName("fluor_panel_dock") self.fluor_panel_dock.setWidget(self.fluor_panel) self.fluor_panel_dock.setAllowedAreas( - Qt.DockWidgetArea.TopDockWidgetArea | Qt.DockWidgetArea.BottomDockWidgetArea | Qt.DockWidgetArea.RightDockWidgetArea) - self.addDockWidget(Qt.DockWidgetArea.BottomDockWidgetArea, self.fluor_panel_dock) + Qt.DockWidgetArea.TopDockWidgetArea + | Qt.DockWidgetArea.BottomDockWidgetArea + | Qt.DockWidgetArea.RightDockWidgetArea + ) + self.addDockWidget( + Qt.DockWidgetArea.BottomDockWidgetArea, self.fluor_panel_dock + ) self.fluor_panel_dock.hide() self.log_dock = LogDock("Console Log", self) @@ -467,7 +518,9 @@ class MainWindow(QMainWindow): ) self.automation_progress_panel.set_running(self.job_list_panel.is_running()) self.compact_automation_panel.set_running(self.job_list_panel.is_running()) - self.compact_automation_panel.set_step_through(self.job_list_panel.is_step_through()) + self.compact_automation_panel.set_step_through( + self.job_list_panel.is_step_through() + ) self.compact_automation_panel.set_samples_in_queue( len(self.job_list_panel.table_model.samples) ) @@ -483,10 +536,12 @@ class MainWindow(QMainWindow): | Qt.DockWidgetArea.LeftDockWidgetArea | Qt.DockWidgetArea.BottomDockWidgetArea ) - self.addDockWidget(Qt.DockWidgetArea.RightDockWidgetArea, self.smargon_trace_dock) + self.addDockWidget( + Qt.DockWidgetArea.RightDockWidgetArea, self.smargon_trace_dock + ) self.smargon_trace_dock.hide() - #Target stability panel + # Target stability panel self.target_stability_panel = TargetStabilityPanel() self.target_stability_dock = QDockWidget("Target stability", self) self.target_stability_dock.setObjectName("target_stability_dock") @@ -496,7 +551,9 @@ class MainWindow(QMainWindow): | Qt.DockWidgetArea.LeftDockWidgetArea | Qt.DockWidgetArea.BottomDockWidgetArea ) - self.addDockWidget(Qt.DockWidgetArea.RightDockWidgetArea, self.target_stability_dock) + self.addDockWidget( + Qt.DockWidgetArea.RightDockWidgetArea, self.target_stability_dock + ) self.target_stability_dock.hide() # Prediction Metrics Panel @@ -509,7 +566,9 @@ class MainWindow(QMainWindow): | Qt.DockWidgetArea.LeftDockWidgetArea | Qt.DockWidgetArea.BottomDockWidgetArea ) - self.addDockWidget(Qt.DockWidgetArea.RightDockWidgetArea, self.prediction_metrics_dock) + self.addDockWidget( + Qt.DockWidgetArea.RightDockWidgetArea, self.prediction_metrics_dock + ) self.prediction_metrics_dock.hide() self.content_stack.addWidget(top_widget) @@ -558,7 +617,9 @@ class MainWindow(QMainWindow): self._remote_close_timer.timeout.connect(self._check_remote_close_deadline) self._axis_camera_refresh_timer = QTimer(self) - self._axis_camera_refresh_timer.setInterval(self._axis_camera_refresh_interval_ms) + self._axis_camera_refresh_timer.setInterval( + self._axis_camera_refresh_interval_ms + ) self._axis_camera_refresh_timer.timeout.connect(self.refresh_axis_cameras) self._axis_camera_refresh_timer.start() @@ -566,8 +627,12 @@ class MainWindow(QMainWindow): job_list_panel=self.job_list_panel, tell_samples=self.tell_samples, ) - self.portrait_mode_panel._back_btn.clicked.connect(self._return_from_portrait_mode) - self.portrait_mode_panel.grab_session_requested.connect(self.status_bar.request_baton) + self.portrait_mode_panel._back_btn.clicked.connect( + self._return_from_portrait_mode + ) + self.portrait_mode_panel.grab_session_requested.connect( + self.status_bar.request_baton + ) # Route alert banner signals through portrait-aware interceptors self.daq.polled_devices_status.connect(self._portrait_alert_primary) @@ -578,7 +643,9 @@ class MainWindow(QMainWindow): 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.automation_progress.connect(self.automation_progress_panel.set_progress) + self.daq.automation_progress.connect( + self.automation_progress_panel.set_progress + ) self.daq.automation_progress.connect(self.compact_automation_panel.set_progress) self.daq.automation_progress.connect(self.portrait_mode_panel.set_progress) @@ -592,7 +659,9 @@ 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.samcam.save_beam_location_setting.connect(self.daq.save_beam_location_camera_setting) + self.beamline.samcam.save_beam_location_setting.connect( + self.daq.save_beam_location_camera_setting + ) self.loop_centering.find_tip.clicked.connect(self.daq.center_loop) self.loop_centering.bounding_box.clicked.connect(self.daq.ml_bounding_box) self.daq.raster_generated_by_ml.connect(self.raster.update_active_grid_request) @@ -612,8 +681,12 @@ class MainWindow(QMainWindow): self.beamline.illumination_panel.back_light.connect(self.daq.back_light) if self.__decoded_token.staff: - self.beamline.monochromator_panel.mono_pitch_scan.connect(self.daq.mono_pitch_scan) - self.beamline.monochromator_panel.change_energy.connect(self.daq.change_energy) + self.beamline.monochromator_panel.mono_pitch_scan.connect( + self.daq.mono_pitch_scan + ) + self.beamline.monochromator_panel.change_energy.connect( + self.daq.change_energy + ) self.beamline.abr_tweak.abr_tweak.connect(self.daq.abr_tweak) self.beamline.abr_tweak.abr_save.connect(self.daq.abr_save) self.beamline.abr_tweak.abr_goto_meas.connect(self.daq.abr_goto_meas) @@ -624,38 +697,88 @@ class MainWindow(QMainWindow): self.beamline.beam_size.beam_size.connect(self.daq.beam_size_mm) self.sample_camera.load_image.connect(self.raster.load_image) - self.sample_camera.switch_raster_grid.connect(self.data_collection.switch_to_raster) - self.beamline.samcam.show_detections_changed.connect(self.sample_camera.set_show_detections) - self.beamline.samcam.show_detection_polygons_changed.connect(self.sample_camera.set_show_detection_polygons) - self.beamline.samcam.show_target_point_changed.connect(self.sample_camera.set_show_target_point) - self.beamline.samcam.show_target_coordinates_changed.connect(self.sample_camera.set_show_target_coordinates) - self.beamline.samcam.show_overlay_legend_changed.connect(self.sample_camera.set_show_overlay_legend) - self.beamline.samcam.compact_overlay_legend_changed.connect(self.sample_camera.set_compact_overlay_legend) - self.beamline.samcam.target_color_changed.connect(self.sample_camera.set_target_color) + self.sample_camera.switch_raster_grid.connect( + self.data_collection.switch_to_raster + ) + self.beamline.samcam.show_detections_changed.connect( + self.sample_camera.set_show_detections + ) + self.beamline.samcam.show_detection_polygons_changed.connect( + self.sample_camera.set_show_detection_polygons + ) + self.beamline.samcam.show_target_point_changed.connect( + self.sample_camera.set_show_target_point + ) + self.beamline.samcam.show_target_coordinates_changed.connect( + self.sample_camera.set_show_target_coordinates + ) + self.beamline.samcam.show_overlay_legend_changed.connect( + self.sample_camera.set_show_overlay_legend + ) + self.beamline.samcam.compact_overlay_legend_changed.connect( + self.sample_camera.set_compact_overlay_legend + ) + self.beamline.samcam.target_color_changed.connect( + self.sample_camera.set_target_color + ) self._restore_samcam_overlay_settings() sample_feed_addr = pred_zmq_addr or zmq_addr if sample_feed_addr is not None: logger.debug(f"Starting prediction subscriber thread {sample_feed_addr}") - self.prediction_thread = PredictionSubscriber(pred_zmq_url=sample_feed_addr, topic=b"") + self.prediction_thread = PredictionSubscriber( + pred_zmq_url=sample_feed_addr, topic=b"" + ) self.prediction_thread.image.connect(self.sample_camera.update_pixmap) - self.prediction_thread.image.connect(self.compact_sample_camera.update_pixmap) - self.prediction_thread.image.connect(self.portrait_sample_camera.update_pixmap) - self.prediction_thread.prediction.connect(self.sample_camera.update_detections) - self.prediction_thread.prediction.connect(self.compact_sample_camera.update_detections) - self.prediction_thread.prediction.connect(self.portrait_sample_camera.update_detections) - self.prediction_thread.prediction.connect(self.prediction_metrics_panel.update_from_prediction) - self.prediction_thread.target_point.connect(self.sample_camera.update_target_point) - self.prediction_thread.target_point.connect(self.compact_sample_camera.update_target_point) - self.prediction_thread.target_point.connect(self.portrait_sample_camera.update_target_point) - self.prediction_thread.target_point.connect(self.target_stability_panel.update_target_point) - self.prediction_thread.focus_measure.connect(self.status_bar.update_sharpness) - self.prediction_thread.fps_measure.connect(self.status_bar.update_samcam_fps) - self.prediction_thread.camera_availability_changed.connect(self.sample_camera.set_camera_available) - self.prediction_thread.camera_availability_changed.connect(self.compact_sample_camera.set_camera_available) - self.prediction_thread.camera_availability_changed.connect(self.portrait_sample_camera.set_camera_available) - self.prediction_thread.camera_availability_changed.connect(self._on_sample_camera_availability_changed) + self.prediction_thread.image.connect( + self.compact_sample_camera.update_pixmap + ) + self.prediction_thread.image.connect( + self.portrait_sample_camera.update_pixmap + ) + self.prediction_thread.prediction.connect( + self.sample_camera.update_detections + ) + self.prediction_thread.prediction.connect( + self.compact_sample_camera.update_detections + ) + self.prediction_thread.prediction.connect( + self.portrait_sample_camera.update_detections + ) + self.prediction_thread.prediction.connect( + self.prediction_metrics_panel.update_from_prediction + ) + self.prediction_thread.target_point.connect( + self.sample_camera.update_target_point + ) + self.prediction_thread.target_point.connect( + self.compact_sample_camera.update_target_point + ) + self.prediction_thread.target_point.connect( + self.portrait_sample_camera.update_target_point + ) + self.prediction_thread.target_point.connect( + self.target_stability_panel.update_target_point + ) + self.prediction_thread.focus_measure.connect( + self.status_bar.update_sharpness + ) + self.prediction_thread.fps_measure.connect( + self.status_bar.update_samcam_fps + ) + self.prediction_thread.camera_availability_changed.connect( + self.sample_camera.set_camera_available + ) + self.prediction_thread.camera_availability_changed.connect( + self.compact_sample_camera.set_camera_available + ) + self.prediction_thread.camera_availability_changed.connect( + self.portrait_sample_camera.set_camera_available + ) + self.prediction_thread.camera_availability_changed.connect( + self._on_sample_camera_availability_changed + ) self.prediction_thread.camera_error.connect(self._on_sample_camera_error) self.prediction_thread.start() else: @@ -663,7 +786,9 @@ class MainWindow(QMainWindow): self.sample_camera.set_camera_available(False) self.compact_sample_camera.set_camera_available(False) self.portrait_sample_camera.set_camera_available(False) - self._show_samcam_feed_banner("Sample camera feed unavailable: no stream configured") + self._show_samcam_feed_banner( + "Sample camera feed unavailable: no stream configured" + ) # # self.data_collection.helical.helical_scan.connect(self.worker.helical_scan) @@ -676,11 +801,17 @@ class MainWindow(QMainWindow): self.sample_camera.evaluate_grid.connect(self.raster.run_grid_scan) self.data_collection.raster.evaluate_grid.connect(self.raster.run_grid_scan) - self.data_collection.raster.evaluate_grid_auto.connect(self.raster.run_grid_scan_auto) + self.data_collection.raster.evaluate_grid_auto.connect( + self.raster.run_grid_scan_auto + ) - self.sample_camera.clear_evaluated_grids.connect(self.raster.clear_completed_grids) + self.sample_camera.clear_evaluated_grids.connect( + self.raster.clear_completed_grids + ) self.sample_camera.clear_grid.connect(self.raster.clear_active_grid) - self.daq.run_number_incremented.connect(self.data_collection.file_path_panel.increment_run_number) + self.daq.run_number_incremented.connect( + self.data_collection.file_path_panel.increment_run_number + ) self.job_list_panel.auto_scan.connect(self.daq.automated_scan) @@ -694,21 +825,33 @@ class MainWindow(QMainWindow): self.ref_tools_panel.mount.connect(self._on_manual_mount_requested) self.ref_tools_panel.unmount.connect(self._on_manual_unmount_requested) - self.data_collection.raster.grid_size_updated.connect(self.raster.update_grid_size) - self.data_collection.raster.exp_time_updated.connect(self.raster.update_exposure_time) - self.data_collection.raster.transmission_updated.connect(self.raster.update_transmission) + self.data_collection.raster.grid_size_updated.connect( + self.raster.update_grid_size + ) + self.data_collection.raster.exp_time_updated.connect( + self.raster.update_exposure_time + ) + self.data_collection.raster.transmission_updated.connect( + self.raster.update_transmission + ) self.data_collection.raster.dtz_updated.connect(self.raster.update_dtz) self.data_collection.raster.grid_metric_updated.connect(self.raster.metric) - self.data_collection.raster.raster_alpha_changed.connect(self.sample_camera.raster_alpha) + self.data_collection.raster.raster_alpha_changed.connect( + self.sample_camera.raster_alpha + ) self.data_collection.cancel.connect(self.daq.cancel) self.raster.grid_scan.connect(self.daq.raster_scan) self.raster.grid_scan_auto.connect(self.daq.raster_scan_auto) self.data_collection.screening.rotation_scan.connect(self.daq.standard_scan) - self.data_collection.simple.rotation_scan.connect(self._on_simple_rotation_requested) + self.data_collection.simple.rotation_scan.connect( + self._on_simple_rotation_requested + ) 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.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) @@ -736,20 +879,28 @@ class MainWindow(QMainWindow): self.beamline_state_panel.sample_exchange.connect(self.daq.sample_exchange) self.beamline_state_panel.sample_alignment.connect(self.daq.sample_alignment) self.beamline_state_panel.beam_location.connect(self.daq.beam_location) - self.beamline_state_panel.beamstop_alignment.connect(self.daq.beamstop_alignment) + self.beamline_state_panel.beamstop_alignment.connect( + self.daq.beamstop_alignment + ) self.beamline_state_panel.flux_measurement.connect(self.daq.flux_measurement) self.beamline_state_panel.data_collection.connect(self.daq.data_collection) self.beamline_state_panel.xtal_snapshot.connect(self.daq.xtal_snapshot) self.beamline_state_panel.xray_fluorescence.connect(self.daq.xray_fluorescence) - self.beamline_state_panel.robot_sample_exchange.connect(self.daq.robot_sample_exchange) + self.beamline_state_panel.robot_sample_exchange.connect( + self.daq.robot_sample_exchange + ) self.rotation.file_ready.connect(self.viewer.load_image) self.raster.image_selected.connect(self.viewer.load_image) self.raster.viewer_track_online.connect(self.viewer.load_online) - self.data_collection.screening.viewer_track_online.connect(self.viewer.load_online) + self.data_collection.screening.viewer_track_online.connect( + self.viewer.load_online + ) self.job_list_panel.viewer_track_online.connect(self.viewer.load_online) - self.sample_logic.sample_changed.connect(self.data_collection.file_path_panel.update_sample) + self.sample_logic.sample_changed.connect( + self.data_collection.file_path_panel.update_sample + ) self.daq.update.connect(self.beamline.omega_panel.update_daq_status) self.daq.update.connect(self.beamline.smargon_panel.update_daq_status) @@ -794,9 +945,15 @@ class MainWindow(QMainWindow): self.daq.raster_scan_completed.connect(self.raster.grid_scan_completed) self.daq.automated_scan_done.connect(self.job_list_panel.automated_scan_done) - self.daq.automation_critical_failure.connect(self._on_automation_critical_failure) - self.daq.manual_collection_critical_failure.connect(self._on_manual_collection_critical_failure) - self.daq.recovery_action_completed.connect(lambda _msg: self._clear_automation_critical_banner()) + self.daq.automation_critical_failure.connect( + self._on_automation_critical_failure + ) + self.daq.manual_collection_critical_failure.connect( + self._on_manual_collection_critical_failure + ) + self.daq.recovery_action_completed.connect( + lambda _msg: self._clear_automation_critical_banner() + ) self.manual_sample_panel.sample_manual.connect(self.daq.sample_manual) self.face_panel.face_detection.connect(self.daq.face_detection) @@ -807,7 +964,9 @@ class MainWindow(QMainWindow): self.data_collection.fluo.fluo_scan.connect(self._on_fluo_scan_requested) self.daq.fluorimeter_spectrum_update.connect(self.fluor_panel.update_plot) - self.daq.fluorimeter_spectrum_update.connect(lambda: self.fluor_panel_dock.setVisible(True)) + self.daq.fluorimeter_spectrum_update.connect( + lambda: self.fluor_panel_dock.setVisible(True) + ) # === Alert/Status Message Routing === # Status bar: General status messages (not device connection status) @@ -819,17 +978,26 @@ class MainWindow(QMainWindow): self._shortcut_manual_sample = QAction("Raise Manual Sample Dock", self) self._shortcut_manual_sample.setShortcut(QKeySequence("Ctrl+M")) self._shortcut_manual_sample.triggered.connect( - lambda: (self.manual_sample_dock.setVisible(True), self.manual_sample_dock.raise_())) + lambda: ( + self.manual_sample_dock.setVisible(True), + self.manual_sample_dock.raise_(), + ) + ) self.addAction(self._shortcut_manual_sample) self._shortcut_raise_sample_list = QAction("Raise sample list", self) self._shortcut_raise_sample_list.setShortcut(QKeySequence("Ctrl+L")) self._shortcut_raise_sample_list.triggered.connect( - lambda: (self.tell_samples_dock.setVisible(True), self.tell_samples_dock.raise_()) + lambda: ( + self.tell_samples_dock.setVisible(True), + self.tell_samples_dock.raise_(), + ) ) self.addAction(self._shortcut_raise_sample_list) - self._shortcut_raise_reference_tools_list = QAction("Raise reference tools", self) + self._shortcut_raise_reference_tools_list = QAction( + "Raise reference tools", self + ) self._shortcut_raise_reference_tools_list.setShortcut(QKeySequence("Ctrl+R")) self._shortcut_raise_reference_tools_list.triggered.connect( lambda: (self.ref_tools_dock.setVisible(True), self.ref_tools_dock.raise_()) @@ -843,24 +1011,38 @@ class MainWindow(QMainWindow): ) self.addAction(self._shortcut_raise_job_list) - self._shortcut_toggle_target_stability = QAction("Toggle target stability panel", self) + self._shortcut_toggle_target_stability = QAction( + "Toggle target stability panel", self + ) self._shortcut_toggle_target_stability.setShortcut(QKeySequence("Ctrl+Shift+T")) self._shortcut_toggle_target_stability.triggered.connect( - lambda: self.target_stability_dock.setVisible(not self.target_stability_dock.isVisible()) + lambda: self.target_stability_dock.setVisible( + not self.target_stability_dock.isVisible() + ) ) self.addAction(self._shortcut_toggle_target_stability) - self._shortcut_toggle_prediction_metrics = QAction("Toggle prediction metrics panel", self) - self._shortcut_toggle_prediction_metrics.setShortcut(QKeySequence("Ctrl+Shift+P")) + self._shortcut_toggle_prediction_metrics = QAction( + "Toggle prediction metrics panel", self + ) + self._shortcut_toggle_prediction_metrics.setShortcut( + QKeySequence("Ctrl+Shift+P") + ) self._shortcut_toggle_prediction_metrics.triggered.connect( - lambda: self.prediction_metrics_dock.setVisible(not self.prediction_metrics_dock.isVisible()) + lambda: self.prediction_metrics_dock.setVisible( + not self.prediction_metrics_dock.isVisible() + ) ) self.addAction(self._shortcut_toggle_prediction_metrics) - self._shortcut_toggle_smargon_trace = QAction("Toggle smargon trace panel", self) + self._shortcut_toggle_smargon_trace = QAction( + "Toggle smargon trace panel", self + ) self._shortcut_toggle_smargon_trace.setShortcut(QKeySequence("Ctrl+Shift+S")) self._shortcut_toggle_smargon_trace.triggered.connect( - lambda: self.smargon_trace_dock.setVisible(not self.smargon_trace_dock.isVisible()) + lambda: self.smargon_trace_dock.setVisible( + not self.smargon_trace_dock.isVisible() + ) ) self.addAction(self._shortcut_toggle_smargon_trace) @@ -878,7 +1060,10 @@ class MainWindow(QMainWindow): current_widget = self.content_stack.currentWidget() - if hasattr(self, "portrait_mode_page") and current_widget is self.portrait_mode_page: + if ( + hasattr(self, "portrait_mode_page") + and current_widget is self.portrait_mode_page + ): self._return_from_portrait_mode() elif bool(getattr(self, "_in_compact_automation_view", False)): self._return_from_compact_automation_view() @@ -891,11 +1076,19 @@ class MainWindow(QMainWindow): def _restore_samcam_overlay_settings(self) -> None: settings = QSettings("PSI", "AareGUI") show_detections = settings.value("samcam/show_detections", True, type=bool) - show_detection_polygons = settings.value("samcam/show_detection_polygons", True, type=bool) + show_detection_polygons = settings.value( + "samcam/show_detection_polygons", True, type=bool + ) show_target_point = settings.value("samcam/show_target_point", True, type=bool) - show_target_coordinates = settings.value("samcam/show_target_coordinates", True, type=bool) - show_overlay_legend = settings.value("samcam/show_overlay_legend", True, type=bool) - compact_overlay_legend = settings.value("samcam/compact_overlay_legend", False, type=bool) + show_target_coordinates = settings.value( + "samcam/show_target_coordinates", True, type=bool + ) + show_overlay_legend = settings.value( + "samcam/show_overlay_legend", True, type=bool + ) + compact_overlay_legend = settings.value( + "samcam/compact_overlay_legend", False, type=bool + ) target_color = settings.value("samcam/target_color", "Cyan", type=str) self.beamline.samcam.apply_overlay_settings( @@ -919,11 +1112,17 @@ class MainWindow(QMainWindow): settings = QSettings("PSI", "AareGUI") overlay = self.sample_camera.target_overlay_settings() settings.setValue("samcam/show_detections", overlay["show_detections"]) - settings.setValue("samcam/show_detection_polygons", overlay["show_detection_polygons"]) # NEW + settings.setValue( + "samcam/show_detection_polygons", overlay["show_detection_polygons"] + ) # NEW settings.setValue("samcam/show_target_point", overlay["show_target_point"]) - settings.setValue("samcam/show_target_coordinates", overlay["show_target_coordinates"]) + settings.setValue( + "samcam/show_target_coordinates", overlay["show_target_coordinates"] + ) settings.setValue("samcam/show_overlay_legend", overlay["show_overlay_legend"]) - settings.setValue("samcam/compact_overlay_legend", overlay["compact_overlay_legend"]) + settings.setValue( + "samcam/compact_overlay_legend", overlay["compact_overlay_legend"] + ) settings.setValue("samcam/target_color", overlay["target_color"]) @Slot(bool) @@ -990,14 +1189,22 @@ class MainWindow(QMainWindow): if self._beamline_cam_addr: self.beamline_camera_thread = VideoThread(ip=self._beamline_cam_addr) - self.beamline_camera_thread.frame_ready.connect(self.beamline_view.update_frame) - self.beamline_camera_thread.frame_ready.connect(self.beamline_view_2_combined.update_frame) + self.beamline_camera_thread.frame_ready.connect( + self.beamline_view.update_frame + ) + self.beamline_camera_thread.frame_ready.connect( + self.beamline_view_2_combined.update_frame + ) self.beamline_camera_thread.start() if self._gonio_cam_addr and self._gonio_cam_id: - self.gonio_camera_thread = VideoThread(ip=self._gonio_cam_addr, camera=self._gonio_cam_id) + self.gonio_camera_thread = VideoThread( + ip=self._gonio_cam_addr, camera=self._gonio_cam_id + ) self.gonio_camera_thread.frame_ready.connect(self.gonio_view.update_frame) - self.gonio_camera_thread.frame_ready.connect(self.beamline_view_1_combined.update_frame) + self.gonio_camera_thread.frame_ready.connect( + self.beamline_view_1_combined.update_frame + ) self.gonio_camera_thread.start() def _all_tell_samples_in_default_order(self) -> list: @@ -1028,8 +1235,12 @@ class MainWindow(QMainWindow): if not self._in_compact_automation_view: self._pre_automation_window_state = self.saveState() self._pre_automation_ref_tools_visible = self.ref_tools_dock.isVisible() - self._pre_automation_left_column_visible = self.collection_controls_scroll.isVisible() - self._pre_automation_right_column_visible = self.beamline_controls_scroll.isVisible() + self._pre_automation_left_column_visible = ( + self.collection_controls_scroll.isVisible() + ) + self._pre_automation_right_column_visible = ( + self.beamline_controls_scroll.isVisible() + ) self.tell_samples_dock.setVisible(False) self.job_list_dock.setVisible(False) @@ -1059,8 +1270,12 @@ class MainWindow(QMainWindow): if self._pre_automation_window_state is not None: self.restoreState(self._pre_automation_window_state) - self.collection_controls_scroll.setVisible(self._pre_automation_left_column_visible) - self.beamline_controls_scroll.setVisible(self._pre_automation_right_column_visible) + self.collection_controls_scroll.setVisible( + self._pre_automation_left_column_visible + ) + self.beamline_controls_scroll.setVisible( + self._pre_automation_right_column_visible + ) if self.__decoded_token.staff: self.ref_tools_dock.setVisible(self._pre_automation_ref_tools_visible) @@ -1074,7 +1289,9 @@ class MainWindow(QMainWindow): @Slot() def _refresh_compact_queue_preview(self) -> None: - current_sample, next_sample, next_next_sample = self.job_list_panel.queue_preview() + current_sample, next_sample, next_next_sample = ( + self.job_list_panel.queue_preview() + ) self.compact_automation_panel.set_samples( current_sample, next_sample, @@ -1115,17 +1332,17 @@ class MainWindow(QMainWindow): # Hide all dock widgets for dock_attr in ( - "tell_samples_dock", - "job_list_dock", - "manual_sample_dock", - "automation_progress_dock", - "face_panel_dock", - "fluor_panel_dock", - "smargon_trace_dock", - "target_stability_dock", - "prediction_metrics_dock", - "log_dock", - "ref_tools_dock", + "tell_samples_dock", + "job_list_dock", + "manual_sample_dock", + "automation_progress_dock", + "face_panel_dock", + "fluor_panel_dock", + "smargon_trace_dock", + "target_stability_dock", + "prediction_metrics_dock", + "log_dock", + "ref_tools_dock", ): dock = getattr(self, dock_attr, None) if dock is not None: @@ -1238,7 +1455,9 @@ class MainWindow(QMainWindow): return mapping.get(annotation, str(annotation).strip()) @staticmethod - def _append_annotation_to_comment(existing_comment: str | None, annotation: str) -> str: + def _append_annotation_to_comment( + existing_comment: str | None, annotation: str + ) -> str: token = MainWindow._annotation_token(annotation) current = str(existing_comment or "").strip() @@ -1255,7 +1474,9 @@ class MainWindow(QMainWindow): def _handle_compact_annotation(self, annotation: str) -> None: current_sample, _, _ = self.job_list_panel.queue_preview() if current_sample is None: - self.status_bar.show_connection_message("No sample selected for annotation.", True) + self.status_bar.show_connection_message( + "No sample selected for annotation.", True + ) return updated_comment = self._append_annotation_to_comment( @@ -1263,7 +1484,9 @@ class MainWindow(QMainWindow): annotation, ) - self.job_list_panel.annotate_sample_comment(current_sample.db_id, updated_comment) + self.job_list_panel.annotate_sample_comment( + current_sample.db_id, updated_comment + ) self.tell_samples.annotate_sample_comment(current_sample.db_id, updated_comment) self._refresh_compact_queue_preview() @@ -1317,12 +1540,16 @@ class MainWindow(QMainWindow): self._enter_automation_view_action = QAction("Automation View", self) self._enter_automation_view_action.setShortcut(QKeySequence("Ctrl+5")) - self._enter_automation_view_action.triggered.connect(self.enter_compact_automation_view) + self._enter_automation_view_action.triggered.connect( + self.enter_compact_automation_view + ) menu_bar.addAction(self._enter_automation_view_action) self._return_main_view_action = QAction("Return to Main View", self) self._return_main_view_action.setShortcut(QKeySequence("Ctrl+Shift+5")) - self._return_main_view_action.triggered.connect(self._return_from_compact_automation_view) + self._return_main_view_action.triggered.connect( + self._return_from_compact_automation_view + ) menu_bar.addAction(self._return_main_view_action) self._portrait_mode_action = QAction("Portrait Mode", self) @@ -1352,9 +1579,13 @@ class MainWindow(QMainWindow): view_menu.addSeparator() if self._beamline_state_panel_enabled: - self._show_beamline_state_action = QAction("Show Beamline State Panel", self) + self._show_beamline_state_action = QAction( + "Show Beamline State Panel", self + ) self._show_beamline_state_action.setCheckable(True) - self._show_beamline_state_action.setChecked(self.beamline_state_panel.isVisible()) + self._show_beamline_state_action.setChecked( + self.beamline_state_panel.isVisible() + ) self._show_beamline_state_action.triggered.connect( lambda checked: self.beamline_state_panel.setVisible(checked) ) @@ -1363,7 +1594,9 @@ class MainWindow(QMainWindow): show_samples_action = QAction("Show Sample List", self) show_samples_action.setCheckable(True) show_samples_action.setChecked(True) - show_samples_action.triggered.connect(lambda checked: self.tell_samples_dock.setVisible(checked)) + 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) @@ -1374,44 +1607,66 @@ class MainWindow(QMainWindow): 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) + 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(lambda checked: self.job_list_dock.setVisible(checked)) + 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(lambda checked: self.manual_sample_dock.setVisible(checked)) - self.manual_sample_dock.visibilityChanged.connect(show_manual_sample_action.setChecked) + 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(lambda checked: self.face_panel_dock.setVisible(checked)) - self.face_panel_dock.visibilityChanged.connect(show_face_panel_action.setChecked) + 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) show_fluor_panel_action = QAction("Show fluorescence", self) show_fluor_panel_action.setCheckable(True) show_fluor_panel_action.setChecked(False) - show_fluor_panel_action.triggered.connect(lambda checked: self.fluor_panel_dock.setVisible(checked)) - self.fluor_panel_dock.visibilityChanged.connect(show_fluor_panel_action.setChecked) + show_fluor_panel_action.triggered.connect( + lambda checked: self.fluor_panel_dock.setVisible(checked) + ) + self.fluor_panel_dock.visibilityChanged.connect( + show_fluor_panel_action.setChecked + ) view_menu.addAction(show_fluor_panel_action) show_smargon_trace_action = QAction("Show Smargon trace", self) 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) + show_smargon_trace_action.triggered.connect( + lambda checked: self.smargon_trace_dock.setVisible(checked) + ) self.smargon_trace_dock.visibilityChanged.connect( - lambda visible: self.smargon_trace_panel.refresh_plot(force=True) if visible else None + show_smargon_trace_action.setChecked + ) + self.smargon_trace_dock.visibilityChanged.connect( + lambda visible: ( + self.smargon_trace_panel.refresh_plot(force=True) if visible else None + ) ) view_menu.addAction(show_smargon_trace_action) @@ -1421,7 +1676,9 @@ class MainWindow(QMainWindow): show_target_stability_action.triggered.connect( lambda checked: self.target_stability_dock.setVisible(checked) ) - self.target_stability_dock.visibilityChanged.connect(show_target_stability_action.setChecked) + self.target_stability_dock.visibilityChanged.connect( + show_target_stability_action.setChecked + ) view_menu.addAction(show_target_stability_action) show_prediction_metrics_action = QAction("Show Prediction Metrics", self) @@ -1430,13 +1687,17 @@ class MainWindow(QMainWindow): show_prediction_metrics_action.triggered.connect( lambda checked: self.prediction_metrics_dock.setVisible(checked) ) - self.prediction_metrics_dock.visibilityChanged.connect(show_prediction_metrics_action.setChecked) + self.prediction_metrics_dock.visibilityChanged.connect( + show_prediction_metrics_action.setChecked + ) view_menu.addAction(show_prediction_metrics_action) show_log_action = QAction("Show Log", self) show_log_action.setCheckable(True) show_log_action.setChecked(False) - show_log_action.triggered.connect(lambda checked: self.log_dock.setVisible(checked)) + show_log_action.triggered.connect( + lambda checked: self.log_dock.setVisible(checked) + ) self.log_dock.visibilityChanged.connect(show_log_action.setChecked) view_menu.addAction(show_log_action) @@ -1444,22 +1705,30 @@ class MainWindow(QMainWindow): sample_camera_tab_action = QAction("Sample camera tab", self) sample_camera_tab_action.setShortcut(QKeySequence("Ctrl+1")) - sample_camera_tab_action.triggered.connect(lambda: self.video_tab.setCurrentIndex(0)) + sample_camera_tab_action.triggered.connect( + lambda: self.video_tab.setCurrentIndex(0) + ) view_menu.addAction(sample_camera_tab_action) gonio_camera_tab_action = QAction("Gonio camera tab", self) gonio_camera_tab_action.setShortcut(QKeySequence("Ctrl+2")) - gonio_camera_tab_action.triggered.connect(lambda: self.video_tab.setCurrentIndex(1)) + gonio_camera_tab_action.triggered.connect( + lambda: self.video_tab.setCurrentIndex(1) + ) view_menu.addAction(gonio_camera_tab_action) beamline_view_tab_action = QAction("Beamline view tab", self) beamline_view_tab_action.setShortcut(QKeySequence("Ctrl+3")) - beamline_view_tab_action.triggered.connect(lambda: self.video_tab.setCurrentIndex(2)) + beamline_view_tab_action.triggered.connect( + lambda: self.video_tab.setCurrentIndex(2) + ) view_menu.addAction(beamline_view_tab_action) beamline_combined_tab_action = QAction("Beamline combined view tab", self) beamline_combined_tab_action.setShortcut(QKeySequence("Ctrl+4")) - beamline_combined_tab_action.triggered.connect(lambda: self.video_tab.setCurrentIndex(3)) + beamline_combined_tab_action.triggered.connect( + lambda: self.video_tab.setCurrentIndex(3) + ) view_menu.addAction(beamline_combined_tab_action) view_menu.addSeparator() @@ -1500,8 +1769,12 @@ class MainWindow(QMainWindow): start_text_tutorial_action.triggered.connect(self.start_text_tutorial) help_menu.addAction(start_text_tutorial_action) - start_interactive_tutorial_action = QAction("Start Tutorial (Interactive)", self) - start_interactive_tutorial_action.triggered.connect(self.start_interactive_tutorial) + start_interactive_tutorial_action = QAction( + "Start Tutorial (Interactive)", self + ) + start_interactive_tutorial_action.triggered.connect( + self.start_interactive_tutorial + ) help_menu.addAction(start_interactive_tutorial_action) def _capture_default_window_state(self) -> None: @@ -1572,13 +1845,13 @@ class MainWindow(QMainWindow): self._dev_help_dialog.activateWindow() def _show_runtime_notification( - self, - *, - title: str, - message: str, - level: str = "error", - sticky: bool = True, - auto_clear_ms: int | None = None, + self, + *, + title: str, + message: str, + level: str = "error", + sticky: bool = True, + auto_clear_ms: int | None = None, ) -> None: self.log_dock.show_notification( title=title, @@ -1620,9 +1893,9 @@ class MainWindow(QMainWindow): def _is_detector_state_error(message: str) -> bool: text = (message or "").lower() return ( - "daq state error" in text - or "must be idle to start measurement" in text - or "must be idle" in text + "daq state error" in text + or "must be idle to start measurement" in text + or "must be idle" in text ) def _detector_error_dialog_title(self, message: str) -> str: @@ -1714,7 +1987,9 @@ class MainWindow(QMainWindow): "complete the safety search before mounting." ) if getattr(bl, "pss_alarm", False): - return "The hutch safety alarm is active. Mounting is blocked until it clears." + return ( + "The hutch safety alarm is active. Mounting is blocked until it clears." + ) return None def _on_manual_mount_requested(self, sample, reference: bool = False) -> None: @@ -1816,27 +2091,42 @@ class MainWindow(QMainWindow): if self.daq is not None: self.daq.send_status_request() except Exception as e: - logger.error(f"Failed to pause automation queue after critical failure: {e}") + logger.error( + f"Failed to pause automation queue after critical failure: {e}" + ) # 2. Mark the automation progress widget as finished-with-error so # _is_automation_active() returns False and idle/close timers behave. try: - from aare.common.automation_models import ( + from aarecommon.models.automation import ( AutomationProgress, StepState, StepStatus, WorkflowStateKind, ) + progress = getattr(self.automation_progress_panel, "_progress", None) if progress is None: progress = AutomationProgress( current_step=None, steps=[ - StepState(step=WorkflowStateKind.MOUNT, status=StepStatus.PENDING), - StepState(step=WorkflowStateKind.LOOP_CENTRE, status=StepStatus.PENDING), - StepState(step=WorkflowStateKind.RASTER, status=StepStatus.PENDING), - StepState(step=WorkflowStateKind.DATA_COLLECTION, status=StepStatus.PENDING), - StepState(step=WorkflowStateKind.FINAL, status=StepStatus.PENDING), + StepState( + step=WorkflowStateKind.MOUNT, status=StepStatus.PENDING + ), + StepState( + step=WorkflowStateKind.LOOP_CENTRE, + status=StepStatus.PENDING, + ), + StepState( + step=WorkflowStateKind.RASTER, status=StepStatus.PENDING + ), + StepState( + step=WorkflowStateKind.DATA_COLLECTION, + status=StepStatus.PENDING, + ), + StepState( + step=WorkflowStateKind.FINAL, status=StepStatus.PENDING + ), ], finished=False, success=None, @@ -1852,7 +2142,9 @@ class MainWindow(QMainWindow): progress.success = False self.automation_progress_panel.set_progress(progress) except Exception as e: - logger.error(f"Failed to update automation progress after critical failure: {e}") + logger.error( + f"Failed to update automation progress after critical failure: {e}" + ) self._show_runtime_notification( title="Automation paused", @@ -1870,7 +2162,7 @@ class MainWindow(QMainWindow): self._detector_error_dialog_title(message), ( "Automation has been stopped because there is an error with the detector.\n\n" - f"Details:\n{message}" + f"Details:\n{message}" ), ) else: @@ -1940,7 +2232,7 @@ class MainWindow(QMainWindow): return QMessageBox.warning(self, "No Sample", f"{msg}") -#TODO tidy up mount and sampel view fucntions + # TODO tidy up mount and sampel view fucntions @Slot() def mount_view(self): self.video_tab.setCurrentWidget(self.beamline_combined_panel) @@ -1956,10 +2248,16 @@ class MainWindow(QMainWindow): if self._is_automation_active(): self._refresh_idle_activity(report_backend=False) - if hasattr(self, "beamline_camera_thread") and self.beamline_camera_thread is not None: + if ( + hasattr(self, "beamline_camera_thread") + and self.beamline_camera_thread is not None + ): self.beamline_camera_thread.set_busy(s.busy) - if hasattr(self, "gonio_camera_thread") and self.gonio_camera_thread is not None: + if ( + hasattr(self, "gonio_camera_thread") + and self.gonio_camera_thread is not None + ): self.gonio_camera_thread.set_busy(s.busy) busy_style = build_busy_overlay_style( @@ -1968,11 +2266,17 @@ class MainWindow(QMainWindow): session_state=getattr(getattr(s, "session", None), "session", None), ) - if hasattr(self, "beamline_view_panel") and self.beamline_view_panel is not None: + if ( + hasattr(self, "beamline_view_panel") + and self.beamline_view_panel is not None + ): self.beamline_view_panel.set_busy_style(busy_style) if hasattr(self, "gonio_view_panel") and self.gonio_view_panel is not None: self.gonio_view_panel.set_busy_style(busy_style) - if hasattr(self, "beamline_combined_panel") and self.beamline_combined_panel is not None: + if ( + hasattr(self, "beamline_combined_panel") + and self.beamline_combined_panel is not None + ): self.beamline_combined_panel.set_busy_style(busy_style) self.target_stability_panel.set_beam_center( @@ -2005,7 +2309,6 @@ class MainWindow(QMainWindow): self.__mounting = False self.video_tab.setCurrentWidget(self.sample_camera) - # ========== BATON DIALOG HANDLING ========== @Slot(dict) @@ -2028,8 +2331,12 @@ class MainWindow(QMainWindow): timeout_seconds=timeout, parent=self, ) - self._baton_request_dialog.accepted_signal.connect(self.status_bar._on_baton_dialog_accepted) - self._baton_request_dialog.refused_signal.connect(self.status_bar._on_baton_dialog_refused) + self._baton_request_dialog.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() @@ -2054,13 +2361,19 @@ class MainWindow(QMainWindow): self._close_baton_pending_dialog() if status.you_are_holder: - self.alert_banner.show_message("Baton acquired!", False, auto_clear_ms=10000) + 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) + self.alert_banner.show_message( + "Request declined or cancelled", False, auto_clear_ms=10000 + ) # Manage incoming request dialog (when someone requests from us) if not status.incoming_request and self._baton_request_dialog is not None: - logger.info("Incoming baton request no longer active, closing request dialog") + logger.info( + "Incoming baton request no longer active, closing request dialog" + ) self._close_baton_dialog() @Slot(dict) @@ -2073,7 +2386,9 @@ class MainWindow(QMainWindow): """ if result.get("granted"): self._waiting_for_baton_response = False - self.alert_banner.show_message("Baton acquired!", False, auto_clear_ms=10000) + self.alert_banner.show_message( + "Baton acquired!", False, auto_clear_ms=10000 + ) logger.info("Baton acquired") # Close pending dialog; StatusBar will trigger p-group selection via SSE self._close_baton_pending_dialog() @@ -2086,10 +2401,15 @@ class MainWindow(QMainWindow): if getattr(self, "_baton_pending_dialog", None) is None: target_user = holder.replace("Request sent to ", "").replace( - " (Note: beamline is currently busy, transfer will be queued if accepted)", "") - 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) + " (Note: beamline is currently busy, transfer will be queued if accepted)", + "", + ) + 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) @@ -2099,13 +2419,18 @@ class MainWindow(QMainWindow): elif result.get("queued"): self._waiting_for_baton_response = True - self.alert_banner.show_waiting("Control transfer queued - waiting for beamline") + 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 = 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() @@ -2116,8 +2441,9 @@ class MainWindow(QMainWindow): elif result.get("error"): self._waiting_for_baton_response = False - self.alert_banner.show_message(result.get("message", "Request failed"), True, - auto_clear_ms=15000) + self.alert_banner.show_message( + result.get("message", "Request failed"), True, auto_clear_ms=15000 + ) logger.warning(f"Baton request failed: {result.get('message')}") self._close_baton_pending_dialog() @@ -2130,12 +2456,16 @@ class MainWindow(QMainWindow): logger.debug(f"Baton response result: {result}") if result.get("accepted"): self._waiting_for_baton_response = False - self.alert_banner.show_message("Control transferred", False, auto_clear_ms=10000) + self.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) elif result.get("refused"): self._waiting_for_baton_response = False - self.alert_banner.show_message("Request declined", False, auto_clear_ms=10000) + 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) else: @@ -2154,25 +2484,34 @@ class MainWindow(QMainWindow): elif result.get("granted"): self._waiting_for_baton_response = False - self.alert_banner.show_message("Baton acquired!", False, auto_clear_ms=10000) + self.alert_banner.show_message( + "Baton acquired!", False, auto_clear_ms=10000 + ) self._close_baton_pending_dialog() elif result.get("queued"): self._waiting_for_baton_response = True - self.alert_banner.show_waiting("Control transfer queued - waiting for beamline") + self.alert_banner.show_waiting( + "Control transfer queued - waiting for beamline" + ) if getattr(self, "_baton_pending_dialog", None) is not None: self._baton_pending_dialog.set_queued_state() else: - self._baton_pending_dialog = BatonPendingDialog(target_user="Current Holder", timeout_seconds=0, - parent=self) - self._baton_pending_dialog.cancelled_signal.connect(self.daq.cancel_baton_request) + self._baton_pending_dialog = 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.alert_banner.show_message( + "Request declined", False, auto_clear_ms=10000 + ) self._close_baton_pending_dialog() else: @@ -2191,7 +2530,7 @@ class MainWindow(QMainWindow): 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'): + if hasattr(self._baton_pending_dialog, "_timer"): self._baton_pending_dialog._timer.stop() self._baton_pending_dialog.close() finally: @@ -2203,7 +2542,9 @@ class MainWindow(QMainWindow): settings.beginGroup("panel_visibility") settings.setValue("smargon_trace", self.smargon_trace_dock.isVisible()) settings.setValue("target_stability", self.target_stability_dock.isVisible()) - settings.setValue("prediction_metrics", self.prediction_metrics_dock.isVisible()) + settings.setValue( + "prediction_metrics", self.prediction_metrics_dock.isVisible() + ) settings.setValue("face_detection", self.face_panel_dock.isVisible()) settings.setValue("fluorescence", self.fluor_panel_dock.isVisible()) settings.setValue("log", self.log_dock.isVisible()) @@ -2215,22 +2556,32 @@ class MainWindow(QMainWindow): settings.beginGroup("panel_visibility") if settings.contains("smargon_trace"): - self.smargon_trace_dock.setVisible(settings.value("smargon_trace", False, type=bool)) + self.smargon_trace_dock.setVisible( + settings.value("smargon_trace", False, type=bool) + ) if settings.contains("target_stability"): - self.target_stability_dock.setVisible(settings.value("target_stability", False, type=bool)) + self.target_stability_dock.setVisible( + settings.value("target_stability", False, type=bool) + ) if settings.contains("prediction_metrics"): - self.prediction_metrics_dock.setVisible(settings.value("prediction_metrics", False, type=bool)) + self.prediction_metrics_dock.setVisible( + settings.value("prediction_metrics", False, type=bool) + ) if settings.contains("face_detection"): - self.face_panel_dock.setVisible(settings.value("face_detection", False, type=bool)) + self.face_panel_dock.setVisible( + settings.value("face_detection", False, type=bool) + ) if settings.contains("fluorescence"): - self.fluor_panel_dock.setVisible(settings.value("fluorescence", False, type=bool)) + self.fluor_panel_dock.setVisible( + settings.value("fluorescence", False, type=bool) + ) if settings.contains("log"): self.log_dock.setVisible(settings.value("log", False, type=bool)) settings.endGroup() def _restore_window_state(self) -> None: - #TODO put all setting related handlign into state_manager + # TODO put all setting related handlign into state_manager self.state_manager.restore_window(self) self._restore_panel_visibility_settings() @@ -2275,7 +2626,10 @@ class MainWindow(QMainWindow): self._cleanup_done = True try: - if hasattr(self, "_samcam_source_timer") and self._samcam_source_timer is not None: + if ( + hasattr(self, "_samcam_source_timer") + and self._samcam_source_timer is not None + ): self._samcam_source_timer.stop() except Exception as e: logger.warning(f"Failed to stop _samcam_source_timer: {e}") @@ -2287,13 +2641,19 @@ class MainWindow(QMainWindow): logger.warning(f"Failed to stop _idle_timer: {e}") try: - if hasattr(self, "_remote_close_timer") and self._remote_close_timer is not None: + if ( + hasattr(self, "_remote_close_timer") + and self._remote_close_timer is not None + ): self._remote_close_timer.stop() except Exception as e: logger.warning(f"Failed to stop _remote_close_timer: {e}") try: - if hasattr(self, "_axis_camera_refresh_timer") and self._axis_camera_refresh_timer is not None: + if ( + hasattr(self, "_axis_camera_refresh_timer") + and self._axis_camera_refresh_timer is not None + ): self._axis_camera_refresh_timer.stop() except Exception as e: logger.warning(f"Failed to stop _axis_camera_refresh_timer: {e}") @@ -2306,9 +2666,7 @@ class MainWindow(QMainWindow): self._stop_axis_camera_threads() - for attr_name in ( - "prediction_thread", - ): + for attr_name in ("prediction_thread",): thread = getattr(self, attr_name, None) if thread is None: continue @@ -2345,7 +2703,10 @@ class MainWindow(QMainWindow): return now = time.monotonic() - if now - self._last_interaction_report_ts < self._interaction_report_min_interval_s: + if ( + now - self._last_interaction_report_ts + < self._interaction_report_min_interval_s + ): return self._last_interaction_report_ts = now @@ -2375,7 +2736,9 @@ class MainWindow(QMainWindow): logger.debug(f"GUI interaction event filter error: {e}") return super().eventFilter(obj, event) - def _start_remote_close_countdown(self, requested_by: str | None, grace_seconds: int | None) -> None: + def _start_remote_close_countdown( + self, requested_by: str | None, grace_seconds: int | None + ) -> None: grace = max(1, int(grace_seconds or 60)) self._remote_close_deadline_ts = time.time() + grace self._remote_close_reason = requested_by or "staff" @@ -2435,8 +2798,10 @@ class MainWindow(QMainWindow): return if not self._can_close_for_idle_or_remote(): - logger.info("Idle timeout reached, but GUI remains open because beamline is busy or automation is active.") + logger.info( + "Idle timeout reached, but GUI remains open because beamline is busy or automation is active." + ) return logger.warning("Closing GUI after inactivity timeout.") - self.close() \ No newline at end of file + self.close() diff --git a/src/aare/gui/models/bookmark.py b/src/aare/gui/models/bookmark.py index 39e265aa..8483d8f6 100644 --- a/src/aare/gui/models/bookmark.py +++ b/src/aare/gui/models/bookmark.py @@ -1,9 +1,8 @@ from typing import Literal +from aarecommon.math.coordinate import SmargonCoordinate from PySide6.QtGui import QColor -from aare.common.coordinate import SmargonCoordinate - class SmargonBookmark: coord: SmargonCoordinate diff --git a/src/aare/gui/models/sample_queue_model.py b/src/aare/gui/models/sample_queue_model.py index 77e0569f..c2a20ed4 100644 --- a/src/aare/gui/models/sample_queue_model.py +++ b/src/aare/gui/models/sample_queue_model.py @@ -1,8 +1,7 @@ +from aarecommon.models.models import SampleShortInfo, SampleShortInfoList from PySide6.QtCore import QAbstractTableModel, Qt from PySide6.QtGui import QBrush, QColor -from aare.common.models import SampleShortInfo, SampleShortInfoList - def get_entry(sample: SampleShortInfo, column: int, *, show_user: bool = False): if show_user: @@ -55,7 +54,9 @@ class SampleQueueSpreadsheet(QAbstractTableModel): def data(self, index, role=None): if role == Qt.ItemDataRole.DisplayRole: - return get_entry(self.samples[index.row()], index.column(), show_user=self._show_user) + return get_entry( + self.samples[index.row()], index.column(), show_user=self._show_user + ) elif role == Qt.ItemDataRole.TextAlignmentRole: return Qt.AlignmentFlag.AlignCenter elif role == Qt.ItemDataRole.BackgroundRole: @@ -75,13 +76,16 @@ class SampleQueueSpreadsheet(QAbstractTableModel): return str(section + 1) return None - def updateData(self, samples: list[SampleShortInfo],): + def updateData( + self, + samples: list[SampleShortInfo], + ): self.beginResetModel() self.samples = list(samples) self.endResetModel() def mimeTypes(self): - return ['text/plain'] + return ["text/plain"] def canDropMimeData(self, data, action, row, column, parent): if data.hasText(): @@ -152,8 +156,7 @@ class SampleQueueSpreadsheet(QAbstractTableModel): self.beginResetModel() self.samples = [ - SampleShortInfo.model_validate(d) - for d in state.get("samples", []) + SampleShortInfo.model_validate(d) for d in state.get("samples", []) ] - self.endResetModel() \ No newline at end of file + self.endResetModel() diff --git a/src/aare/gui/models/user_sample_model.py b/src/aare/gui/models/user_sample_model.py index e9163eb2..1637c8f4 100644 --- a/src/aare/gui/models/user_sample_model.py +++ b/src/aare/gui/models/user_sample_model.py @@ -1,10 +1,9 @@ import re -from PySide6.QtCore import QAbstractTableModel, Qt, QMimeData +from aarecommon.models.models import SampleShortInfo, SampleShortInfoList +from PySide6.QtCore import QAbstractTableModel, QMimeData, Qt from PySide6.QtGui import QBrush, QColor -from aare.common.models import SampleShortInfo, SampleShortInfoList - def get_entry(sample: SampleShortInfo, column: int): if column == 0: @@ -31,14 +30,13 @@ def get_entry(sample: SampleShortInfo, column: int): return sample.comment - class UserSampleSpreadsheet(QAbstractTableModel): def __init__( self, parent=None, samples: list[SampleShortInfo] | None = None, current_puck: str | None = None, - current_sample: int | None = None + current_sample: int | None = None, ): super().__init__(parent) if samples is None: @@ -111,10 +109,7 @@ class UserSampleSpreadsheet(QAbstractTableModel): self.current_puck = current_puck self.current_sample = current_sample - def updateData( - self, - samples: list[SampleShortInfo] - ): + def updateData(self, samples: list[SampleShortInfo]): if samples != self.samples: self.beginResetModel() self.samples = samples @@ -128,7 +123,6 @@ class UserSampleSpreadsheet(QAbstractTableModel): self._sort() self.layoutChanged.emit() - def _sort(self): filtered = self._apply_filter(self.samples) if self.__sort_col == 3: @@ -151,7 +145,9 @@ class UserSampleSpreadsheet(QAbstractTableModel): def _apply_filter(self, rows: list[SampleShortInfo]) -> list[SampleShortInfo]: # Default filter by User using current p-group if no explicit filter set - filters: dict[int, str] = {col: v for col, v in (self.__filters or {}).items() if (v or "").strip()} + filters: dict[int, str] = { + col: v for col, v in (self.__filters or {}).items() if (v or "").strip() + } if self.__filter_col is not None and (self.__filter_value or "").strip(): filters[self.__filter_col] = self.__filter_value @@ -179,7 +175,7 @@ class UserSampleSpreadsheet(QAbstractTableModel): return default_flags def mimeTypes(self): - return ['text/plain'] + return ["text/plain"] def mimeData(self, indexes): mime_data = QMimeData() @@ -195,7 +191,6 @@ class UserSampleSpreadsheet(QAbstractTableModel): def get_id(self, row: int) -> SampleShortInfo: return self.__sorted_samples[row] - def set_filter(self, field: str, text: str | None): try: col = self.header.index(field) @@ -253,7 +248,7 @@ class UserSampleSpreadsheet(QAbstractTableModel): # Restore the filter if temp_filter is not None: self.__filters[column] = temp_filter - + seen: set[str] = set() out: list[str] = [] for s in filtered_samples: @@ -268,11 +263,17 @@ class UserSampleSpreadsheet(QAbstractTableModel): out.append(txt) if len(out) >= limit: break - + # Sort appropriately if column == 5: # User/pgroup column try: - out.sort(key=lambda x: int(x[1:]) if x and x[0].lower() == "p" and x[1:].isdigit() else float("inf")) + out.sort( + key=lambda x: ( + int(x[1:]) + if x and x[0].lower() == "p" and x[1:].isdigit() + else float("inf") + ) + ) except Exception: out.sort() else: @@ -286,7 +287,7 @@ class UserSampleSpreadsheet(QAbstractTableModel): filtered_samples = self._apply_filter(self.samples) if temp_filter is not None: self.__filters[0] = temp_filter - + rx = re.compile(r"^([A-Za-z]+)") counts: dict[str, int] = {} for s in filtered_samples: @@ -303,14 +304,16 @@ class UserSampleSpreadsheet(QAbstractTableModel): items = sorted(counts.items(), key=lambda kv: (-kv[1], kv[0])) return [k for k, _ in items[:limit]] - def suggested_prefixes_for_location(self, limit: int = 200) -> tuple[list[str], list[str]]: + def suggested_prefixes_for_location( + self, limit: int = 200 + ) -> tuple[list[str], list[str]]: """Get location prefixes from currently filtered samples (excluding column 3 filter).""" # Get currently filtered samples, excluding the location filter temp_filter = self.__filters.pop(3, None) filtered_samples = self._apply_filter(self.samples) if temp_filter is not None: self.__filters[3] = temp_filter - + seg_seen: set[str] = set() segpos_seen: set[str] = set() for s in filtered_samples: @@ -323,5 +326,7 @@ class UserSampleSpreadsheet(QAbstractTableModel): if seg: seg_seen.add(seg) segs = sorted(seg_seen) - segpos = sorted(segpos_seen, key=lambda x: (x[0], int(x[1:]) if x[1:].isdigit() else 0)) - return (segs[:limit], segpos[:limit]) \ No newline at end of file + segpos = sorted( + segpos_seen, key=lambda x: (x[0], int(x[1:]) if x[1:].isdigit() else 0) + ) + return (segs[:limit], segpos[:limit]) diff --git a/src/aare/gui/panels/LogPanel.py b/src/aare/gui/panels/LogPanel.py index ef4986a8..3ad76b36 100644 --- a/src/aare/gui/panels/LogPanel.py +++ b/src/aare/gui/panels/LogPanel.py @@ -1,3 +1,4 @@ +from aarecommon.config.logger import attach_to_logger, find_existing_formatter from PySide6.QtCore import Qt, QTimer, Signal, Slot from PySide6.QtWidgets import ( QDockWidget, @@ -11,7 +12,6 @@ from PySide6.QtWidgets import ( QWidget, ) -from aare.common.logger_config import attach_to_logger, find_existing_formatter from aare.gui.log import QtLogEmitter, QtLogHandler diff --git a/src/aare/gui/panels/abr_tweak_panel.py b/src/aare/gui/panels/abr_tweak_panel.py index 9eb589cf..497d886f 100644 --- a/src/aare/gui/panels/abr_tweak_panel.py +++ b/src/aare/gui/panels/abr_tweak_panel.py @@ -1,8 +1,8 @@ +from aarecommon.math.coordinate import AerotechCoordinate, Coordinate +from aarecommon.models.models import DAQStatusModel 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, AerotechCoordinate -from aare.common.models import DAQStatusModel +from PySide6.QtWidgets import QGridLayout, QLabel, QPushButton, QWidget from aare.gui.widgets.button_with_payload import ButtonWithPayload from aare.gui.widgets.number_line_edit import NumberLineEdit @@ -10,6 +10,7 @@ from aare.gui.widgets.title_label import TitleLabel DEFAULT_ABR_STEP_UM = 5 + class AbrTweakButtons(QWidget): abr_tweak = Signal(AerotechCoordinate) @@ -66,12 +67,21 @@ class AbrTweakButtons(QWidget): @Slot(dict) def abr_button(self, payload: dict): - 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"]))) + 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(AerotechCoordinate) abr_save = Signal() @@ -87,8 +97,8 @@ class AbrTweakWidget(QWidget): grid_layout.addWidget(TitleLabel("ABR meas. pos.", self), 0, 0, 1, 3) - self.__abr_buttons = AbrTweakButtons(DEFAULT_ABR_STEP_UM/1000, parent=self) - grid_layout.addWidget(self.__abr_buttons, 1, 0, 1 ,4) + self.__abr_buttons = AbrTweakButtons(DEFAULT_ABR_STEP_UM / 1000, parent=self) + grid_layout.addWidget(self.__abr_buttons, 1, 0, 1, 4) self.__abr_buttons.abr_tweak.connect(self.abr_button_pressed) grid_layout.addWidget(QLabel("Step", parent=self), 2, 0) diff --git a/src/aare/gui/panels/automation_panel.py b/src/aare/gui/panels/automation_panel.py index c44e1fc2..6e4186ab 100644 --- a/src/aare/gui/panels/automation_panel.py +++ b/src/aare/gui/panels/automation_panel.py @@ -4,23 +4,22 @@ import copy import time from datetime import datetime -from PySide6.QtCore import Slot, QTimer -from PySide6.QtWidgets import ( - QWidget, - QVBoxLayout, - QLabel, -) - -from aare.common.automation_models import ( +from aarecommon.config.logger import setup_logger +from aarecommon.models.automation import ( AutomationProgress, StepStatus, WorkflowStateKind, ) - -from aare.common.logger_config import setup_logger +from PySide6.QtCore import QTimer, Slot +from PySide6.QtWidgets import ( + QLabel, + QVBoxLayout, + QWidget, +) logger = setup_logger("aareGUI") + class AutomationProgressWidget(QWidget): """Compact fixed-step widget for DAQ automation progress.""" @@ -96,7 +95,8 @@ class AutomationProgressWidget(QWidget): @staticmethod def _make_step(step: WorkflowStateKind, status: StepStatus): - from aare.common.automation_models import StepState + from aarecommon.models.automation import StepState + return StepState(step=step, status=status, message="") @staticmethod @@ -128,16 +128,33 @@ class AutomationProgressWidget(QWidget): ) if status == StepStatus.SUCCESS: - return base + " background-color: #ECFDF3; color: #166534; border-color: #A7F3D0;" + return ( + base + + " background-color: #ECFDF3; color: #166534; border-color: #A7F3D0;" + ) if status == StepStatus.RUNNING: - return base + " background-color: #EFF6FF; color: #1D4ED8; font-weight: 700; border-color: #BFDBFE;" + return ( + base + + " background-color: #EFF6FF; color: #1D4ED8; font-weight: 700; border-color: #BFDBFE;" + ) if status == StepStatus.FAILED: - return base + " background-color: #FEF2F2; color: #B91C1C; font-weight: 700; border-color: #FECACA;" + return ( + base + + " background-color: #FEF2F2; color: #B91C1C; font-weight: 700; border-color: #FECACA;" + ) if status == StepStatus.PAUSED: - return base + " background-color: #FFF7ED; color: #C2410C; font-weight: 700; border-color: #FED7AA;" + return ( + base + + " background-color: #FFF7ED; color: #C2410C; font-weight: 700; border-color: #FED7AA;" + ) if status == StepStatus.SKIPPED: - return base + " background-color: #F8FAFC; color: #475569; border-color: #E2E8F0;" - return base + " background-color: #F8FAFC; color: #64748B; border-color: #E2E8F0;" + return ( + base + + " background-color: #F8FAFC; color: #475569; border-color: #E2E8F0;" + ) + return ( + base + " background-color: #F8FAFC; color: #64748B; border-color: #E2E8F0;" + ) @staticmethod def _format_duration(seconds: float | None) -> str: @@ -227,8 +244,14 @@ class AutomationProgressWidget(QWidget): self._refresh_timer.stop() if self._stats_label: - measured_avg_seconds = progress.avg_time_per_sample if progress.avg_time_per_sample > 0 else None - displayed_avg_seconds = measured_avg_seconds or self.DEFAULT_SAMPLE_ESTIMATE_S + measured_avg_seconds = ( + progress.avg_time_per_sample + if progress.avg_time_per_sample > 0 + else None + ) + displayed_avg_seconds = ( + measured_avg_seconds or self.DEFAULT_SAMPLE_ESTIMATE_S + ) current_sample = progress.current_sample_name or "None" samples_left = ( @@ -237,7 +260,11 @@ class AutomationProgressWidget(QWidget): else max(0, int(progress.samples_in_queue or 0)) ) - if progress.finished or self._is_paused or not self._has_active_sample(progress): + if ( + progress.finished + or self._is_paused + or not self._has_active_sample(progress) + ): queue_remaining = displayed_avg_seconds * samples_left else: queue_remaining = displayed_avg_seconds * max(0, samples_left) @@ -264,14 +291,20 @@ class AutomationProgressWidget(QWidget): duration_str = "" if step_state.started_at is not None: end = step_state.completed_at or time.time() - duration_str = f" ({self._format_duration(end - step_state.started_at)})" + duration_str = ( + f" ({self._format_duration(end - step_state.started_at)})" + ) if step_state.status == StepStatus.FAILED and step_state.message: message = f" — Failed: {step_state.message}" else: message = f" — {step_state.message}" if step_state.message else "" - error_str = f"
Error: {step_state.error_code}" if step_state.error_code else "" + error_str = ( + f"
Error: {step_state.error_code}" + if step_state.error_code + else "" + ) label.setText(f"{icon} {title}{duration_str}{message}{error_str}") label.setStyleSheet(self._style_for_status(step_state.status)) diff --git a/src/aare/gui/panels/beam_center_panel.py b/src/aare/gui/panels/beam_center_panel.py index e59c0c58..c14a39a1 100644 --- a/src/aare/gui/panels/beam_center_panel.py +++ b/src/aare/gui/panels/beam_center_panel.py @@ -1,7 +1,7 @@ +from aarecommon.models.models import DAQStatusModel from PySide6.QtCore import Signal, Slot -from PySide6.QtWidgets import QWidget, QGridLayout, QLabel +from PySide6.QtWidgets import QGridLayout, QLabel, QWidget -from aare.common.models import DAQStatusModel from aare.gui.widgets.number_line_edit import NumberLineEdit from aare.gui.widgets.title_label import TitleLabel diff --git a/src/aare/gui/panels/beam_mark_panel.py b/src/aare/gui/panels/beam_mark_panel.py index 0d064db9..de1b3da0 100644 --- a/src/aare/gui/panels/beam_mark_panel.py +++ b/src/aare/gui/panels/beam_mark_panel.py @@ -1,7 +1,7 @@ +from aarecommon.models.models import DAQStatusModel from PySide6.QtCore import Signal, Slot -from PySide6.QtWidgets import QWidget, QGridLayout, QPushButton, QLabel +from PySide6.QtWidgets import QGridLayout, QLabel, QPushButton, QWidget -from aare.common.models import DAQStatusModel from aare.gui.widgets.title_label import TitleLabel @@ -13,7 +13,7 @@ class BeamMarkWidget(QWidget): grid_layout = QGridLayout(self) - grid_layout.addWidget(TitleLabel("Beam mark (image)", self), 0, 0,1,5) + grid_layout.addWidget(TitleLabel("Beam mark (image)", self), 0, 0, 1, 5) self.x = QLabel("0") self.y = QLabel("0") @@ -25,7 +25,7 @@ class BeamMarkWidget(QWidget): grid_layout.addWidget(QLabel("pxl"), 1, 4) clear_button = QPushButton("Clear marks") - grid_layout.addWidget(clear_button, 2, 0,1,5) + grid_layout.addWidget(clear_button, 2, 0, 1, 5) clear_button.pressed.connect(self.clear_button_pressed) @Slot() diff --git a/src/aare/gui/panels/beam_size_panel.py b/src/aare/gui/panels/beam_size_panel.py index 31c3ed88..3a9219ff 100644 --- a/src/aare/gui/panels/beam_size_panel.py +++ b/src/aare/gui/panels/beam_size_panel.py @@ -1,7 +1,7 @@ +from aarecommon.models.models import DAQStatusModel from PySide6.QtCore import Signal, Slot -from PySide6.QtWidgets import QWidget, QGridLayout, QLabel +from PySide6.QtWidgets import QGridLayout, QLabel, QWidget -from aare.common.models import DAQStatusModel from aare.gui.widgets.number_line_edit import NumberLineEdit from aare.gui.widgets.title_label import TitleLabel diff --git a/src/aare/gui/panels/beamline_recovery_panel.py b/src/aare/gui/panels/beamline_recovery_panel.py index f1c90ed9..de779de3 100644 --- a/src/aare/gui/panels/beamline_recovery_panel.py +++ b/src/aare/gui/panels/beamline_recovery_panel.py @@ -1,19 +1,19 @@ from __future__ import annotations +from aarecommon.models.models import DAQStatusModel from PySide6.QtCore import Slot from PySide6.QtWidgets import ( QDialog, - QVBoxLayout, - QLabel, - QPushButton, - QWidget, QDialogButtonBox, QInputDialog, + QLabel, QLineEdit, QMessageBox, + QPushButton, + QVBoxLayout, + QWidget, ) -from aare.common.models import DAQStatusModel from aare.gui.threads.daq_worker import DAQWorker @@ -163,7 +163,9 @@ class RecoveryPanel(QWidget): def _sample_appears_mounted(self) -> bool: try: - return self._last_status is not None and self._last_status.sample is not None + return ( + self._last_status is not None and self._last_status.sample is not None + ) except Exception: return False @@ -179,16 +181,22 @@ class RecoveryPanel(QWidget): self._recovery_unmount_btn.setEnabled(sample_mounted) self._recovery_unmount_btn.setToolTip( - "" if sample_mounted else "Disabled because no mounted sample is visible in current status." + "" + if sample_mounted + else "Disabled because no mounted sample is visible in current status." ) self._free_beamline_btn.setEnabled(beamline_busy) self._free_beamline_btn.setToolTip( - "" if beamline_busy else "Disabled because beamline does not currently appear busy." + "" + if beamline_busy + else "Disabled because beamline does not currently appear busy." ) self._resync_sample_btn.setEnabled(True) - self._resync_sample_btn.setToolTip("Force a one-shot sample reconciliation against TELL.") + self._resync_sample_btn.setToolTip( + "Force a one-shot sample reconciliation against TELL." + ) def _confirm(self, title: str, msg: str) -> bool: reply = QMessageBox.warning( @@ -296,4 +304,4 @@ class BeamlineRecoveryDialog(QDialog): buttons = QDialogButtonBox(QDialogButtonBox.StandardButton.Close, parent=self) buttons.rejected.connect(self.reject) buttons.accepted.connect(self.accept) - layout.addWidget(buttons) \ No newline at end of file + layout.addWidget(buttons) diff --git a/src/aare/gui/panels/beamline_state_panel.py b/src/aare/gui/panels/beamline_state_panel.py index 9e41a56c..70e4dbdd 100644 --- a/src/aare/gui/panels/beamline_state_panel.py +++ b/src/aare/gui/panels/beamline_state_panel.py @@ -1,12 +1,11 @@ from collections import deque from dataclasses import dataclass +from aarecommon.models.models import BeamlineStateEnum, DAQStatusModel from PySide6.QtCore import QPoint, QRect, Qt, Signal, Slot from PySide6.QtGui import QColor, QPainter, QPen from PySide6.QtWidgets import QFrame, QLabel, QPushButton -from aare.common.models import BeamlineStateEnum, DAQStatusModel - @dataclass(frozen=True) class StationSpec: @@ -118,21 +117,86 @@ class BeamlineStatePanel(QFrame): } self._stations = [ - StationSpec(BeamlineStateEnum.DewarTransfer, "Dewar transfer", 54, 140, True, "Dewar transfer mode"), - StationSpec(BeamlineStateEnum.SampleExchange, "Manual sample exchange", 54, 176, True, - "Manual sample exchange mode"), - StationSpec(BeamlineStateEnum.RobotSampleExchange, "Robot sample exchange", 54, 212, True, - "Robot-assisted sample exchange"), - StationSpec(BeamlineStateEnum.SampleAlignment, "Sample alignment", 54, 248, True, - "Sample centring and alignment mode"), - StationSpec(BeamlineStateEnum.BeamLocation, "Beam location", 54, 284, True, "Beam location mode"), - StationSpec(BeamlineStateEnum.BeamstopAlignment, "Beamstop alignment", 54, 320, True, - "Beamstop alignment mode"), - StationSpec(BeamlineStateEnum.FluxMeasurement, "Flux measurement", 54, 356, True, "Flux measurement mode"), - StationSpec(BeamlineStateEnum.DataCollection, "Data collection", 54, 392, True, - "Measurement / collection mode"), - StationSpec(BeamlineStateEnum.XtalSnapshot, "Crystal snapshot", 54, 428, True, "Crystal snapshot mode"), - StationSpec(BeamlineStateEnum.XrayFluorescence, "XRF", 54, 464, True, "X-ray fluorescence mode"), + StationSpec( + BeamlineStateEnum.DewarTransfer, + "Dewar transfer", + 54, + 140, + True, + "Dewar transfer mode", + ), + StationSpec( + BeamlineStateEnum.SampleExchange, + "Manual sample exchange", + 54, + 176, + True, + "Manual sample exchange mode", + ), + StationSpec( + BeamlineStateEnum.RobotSampleExchange, + "Robot sample exchange", + 54, + 212, + True, + "Robot-assisted sample exchange", + ), + StationSpec( + BeamlineStateEnum.SampleAlignment, + "Sample alignment", + 54, + 248, + True, + "Sample centring and alignment mode", + ), + StationSpec( + BeamlineStateEnum.BeamLocation, + "Beam location", + 54, + 284, + True, + "Beam location mode", + ), + StationSpec( + BeamlineStateEnum.BeamstopAlignment, + "Beamstop alignment", + 54, + 320, + True, + "Beamstop alignment mode", + ), + StationSpec( + BeamlineStateEnum.FluxMeasurement, + "Flux measurement", + 54, + 356, + True, + "Flux measurement mode", + ), + StationSpec( + BeamlineStateEnum.DataCollection, + "Data collection", + 54, + 392, + True, + "Measurement / collection mode", + ), + StationSpec( + BeamlineStateEnum.XtalSnapshot, + "Crystal snapshot", + 54, + 428, + True, + "Crystal snapshot mode", + ), + StationSpec( + BeamlineStateEnum.XrayFluorescence, + "XRF", + 54, + 464, + True, + "X-ray fluorescence mode", + ), ] self._segments = [ @@ -182,7 +246,9 @@ class BeamlineStatePanel(QFrame): self._update_collapsed_state() @staticmethod - def _canon_segment(a: BeamlineStateEnum, b: BeamlineStateEnum) -> tuple[BeamlineStateEnum, BeamlineStateEnum]: + def _canon_segment( + a: BeamlineStateEnum, b: BeamlineStateEnum + ) -> tuple[BeamlineStateEnum, BeamlineStateEnum]: return tuple(sorted((a, b), key=lambda state: state.value)) def _build_graph( @@ -240,11 +306,20 @@ class BeamlineStatePanel(QFrame): def _active_hover_route(self) -> set[tuple[BeamlineStateEnum, BeamlineStateEnum]]: if self._hovered_state is not None: - route_source = self._last_stable_state if self._current_state == BeamlineStateEnum.Moving else self._current_state + route_source = ( + self._last_stable_state + if self._current_state == BeamlineStateEnum.Moving + else self._current_state + ) return self._path_segments_between(route_source, self._hovered_state) - if self._current_state == BeamlineStateEnum.Moving and self._pending_target_state is not None: - return self._path_segments_between(self._last_stable_state, self._pending_target_state) + if ( + self._current_state == BeamlineStateEnum.Moving + and self._pending_target_state is not None + ): + return self._path_segments_between( + self._last_stable_state, self._pending_target_state + ) return set() @@ -254,7 +329,11 @@ class BeamlineStatePanel(QFrame): widget: QLabel | QPushButton = HoverableButton(station.label, self) widget.setFlat(True) widget.setCursor(Qt.CursorShape.PointingHandCursor) - widget.clicked.connect(lambda _checked=False, state=station.state: self._emit_for_state(state)) + widget.clicked.connect( + lambda _checked=False, state=station.state: self._emit_for_state( + state + ) + ) else: widget = HoverableLabel(station.label, self) @@ -356,7 +435,10 @@ class BeamlineStatePanel(QFrame): widget = self._station_widgets[station.state] is_current = station.state == self._current_state is_hovered = station.state == self._hovered_state - is_pending = station.state == self._pending_target_state and self._current_state == BeamlineStateEnum.Moving + is_pending = ( + station.state == self._pending_target_state + and self._current_state == BeamlineStateEnum.Moving + ) label_color = self._group_label_colors.get(station.state, "rgb(55, 67, 87)") label_bg = self._group_label_backgrounds.get(station.state, "transparent") @@ -460,13 +542,18 @@ class BeamlineStatePanel(QFrame): if segment in active_hover_route: return self._line_hover - current_path = self._path_segments_between(BeamlineStateEnum.DewarTransfer, self._last_stable_state or self._current_state) + current_path = self._path_segments_between( + BeamlineStateEnum.DewarTransfer, + self._last_stable_state or self._current_state, + ) if segment in current_path: return self._line_current return self._line_color - def _draw_segment(self, painter: QPainter, start: QPoint, end: QPoint, color: QColor) -> None: + def _draw_segment( + self, painter: QPainter, start: QPoint, end: QPoint, color: QColor + ) -> None: pen = QPen(color, 3) pen.setCapStyle(Qt.PenCapStyle.RoundCap) painter.setPen(pen) @@ -485,7 +572,8 @@ class BeamlineStatePanel(QFrame): ) if station.state == self._hovered_state or ( - self._current_state == BeamlineStateEnum.Moving and station.state == self._pending_target_state + self._current_state == BeamlineStateEnum.Moving + and station.state == self._pending_target_state ): painter.setPen(QPen(self._station_hover, 3)) painter.setBrush(self._station_hover) @@ -537,7 +625,9 @@ class BeamlineStatePanel(QFrame): tell_state = status.tell_state tell_text = f"Tell: {tell_state.activity.display_name()}" - tell_phase = tell_state.phase.display_name() if tell_state.phase is not None else "" + tell_phase = ( + tell_state.phase.display_name() if tell_state.phase is not None else "" + ) tell_message = (tell_state.message or "").strip() if tell_phase: @@ -547,7 +637,12 @@ class BeamlineStatePanel(QFrame): if tell_state.activity.value == "error": tell_color = "red" - elif tell_state.activity.value in {"mounting", "unmounting", "drying", "cooling"}: + elif tell_state.activity.value in { + "mounting", + "unmounting", + "drying", + "cooling", + }: tell_color = "orange" else: tell_color = "green" @@ -567,4 +662,4 @@ class BeamlineStatePanel(QFrame): label = state.display_name() if state is not None else "—" self.current_label.setText(f"Current: {label}") self.current_label.adjustSize() - self._apply_station_highlight() \ No newline at end of file + self._apply_station_highlight() diff --git a/src/aare/gui/panels/compact_automation_panel.py b/src/aare/gui/panels/compact_automation_panel.py index d3545f34..de2ca6fe 100644 --- a/src/aare/gui/panels/compact_automation_panel.py +++ b/src/aare/gui/panels/compact_automation_panel.py @@ -1,17 +1,17 @@ +from aarecommon.models.automation import AutomationProgress +from aarecommon.models.models import SampleShortInfo from PySide6.QtCore import Qt, Signal, Slot from PySide6.QtWidgets import ( QFrame, - QVBoxLayout, QHBoxLayout, QLabel, + QMenu, QPushButton, QToolButton, - QMenu, + QVBoxLayout, QWidget, ) -from aare.common.automation_models import AutomationProgress -from aare.common.models import SampleShortInfo from aare.gui.widgets.automation_progress import CompactAutomationProgressStrip @@ -103,7 +103,9 @@ class CompactAutomationPanel(QFrame): self.annotation_button = QToolButton(self) self.annotation_button.setObjectName("compactSecondaryButton") self.annotation_button.setText("Annotate") - self.annotation_button.setPopupMode(QToolButton.ToolButtonPopupMode.InstantPopup) + self.annotation_button.setPopupMode( + QToolButton.ToolButtonPopupMode.InstantPopup + ) annotation_menu = QMenu(self.annotation_button) for label in ["Heart", "Thumbs Up", "Thumbs Down", "Eyes", "Scan Again"]: @@ -228,4 +230,4 @@ class CompactAutomationPanel(QFrame): def _format_sample(sample: SampleShortInfo | None) -> str: if sample is None: return "—" - return f"{sample.sample_name} ({sample.loc_str()})" \ No newline at end of file + return f"{sample.sample_name} ({sample.loc_str()})" diff --git a/src/aare/gui/panels/data_collection_settings.py b/src/aare/gui/panels/data_collection_settings.py index 5ca96870..6912f5e9 100644 --- a/src/aare/gui/panels/data_collection_settings.py +++ b/src/aare/gui/panels/data_collection_settings.py @@ -1,15 +1,14 @@ +from aarecommon.math.diffraction_geometry import DiffractionGeometry +from aarecommon.math.sample_geometry import SampleGeometryModel +from aarecommon.models.models import DAQStatusModel from PySide6.QtCore import Signal, Slot from PySide6.QtWidgets import ( QFrame, - QVBoxLayout, - QTabWidget, QPushButton, + QTabWidget, + QVBoxLayout, ) -from aare.common.diffraction_geometry import DiffractionGeometry -from aare.common.models import DAQStatusModel -from aare.common.sample_geometry import SampleGeometryModel - from aare.gui.panels.file_path_panel import FilePathPanel from aare.gui.panels.fluorescence_data_collection import FluorescenceDataCollectionPanel from aare.gui.panels.raster_data_collection import RasterDataCollectionPanel @@ -22,11 +21,14 @@ class DataCollectionSettings(QFrame): cancel = Signal() set_width = 400 - def __init__(self, - s: SampleGeometryModel, - raster_mgr: RasterGridManager, - diffraction: DiffractionGeometry, - parent=None): + + def __init__( + self, + s: SampleGeometryModel, + raster_mgr: RasterGridManager, + diffraction: DiffractionGeometry, + parent=None, + ): super().__init__(parent) self.setFixedWidth(self.set_width) self.setFrameShape(QFrame.Shape.StyledPanel) @@ -38,10 +40,14 @@ class DataCollectionSettings(QFrame): self.__tab_widget = QTabWidget() - self.raster = RasterDataCollectionPanel(parent=self, raster_mgr=raster_mgr, diffraction=diffraction) + self.raster = RasterDataCollectionPanel( + parent=self, raster_mgr=raster_mgr, diffraction=diffraction + ) self.__tab_widget.addTab(self.raster, "Raster scan") - self.screening = RotationDataCollectionPanel(parent=self, diffraction=diffraction) + self.screening = RotationDataCollectionPanel( + parent=self, diffraction=diffraction + ) self.__tab_widget.addTab(self.screening, "Rotation") self.simple = SimpleRotationSettingsPanel(parent=self) @@ -70,7 +76,6 @@ class DataCollectionSettings(QFrame): self.__tab_widget.currentChanged.connect(self._on_tab_changed) - @Slot() def switch_to_raster(self): self.__tab_widget.setCurrentIndex(0) @@ -96,4 +101,4 @@ class DataCollectionSettings(QFrame): kind = "raster" elif idx == 1 or idx == 2: kind = "rotation" - self.file_path_panel.set_scan_kind(kind) \ No newline at end of file + self.file_path_panel.set_scan_kind(kind) diff --git a/src/aare/gui/panels/developer_help_dialog.py b/src/aare/gui/panels/developer_help_dialog.py index 1f0c833c..0a92ce38 100644 --- a/src/aare/gui/panels/developer_help_dialog.py +++ b/src/aare/gui/panels/developer_help_dialog.py @@ -5,6 +5,9 @@ import logging import time from typing import Dict +from aarecommon.config.logger import attach_to_logger, find_existing_formatter +from aarecommon.errors.codes import error_code_help +from aarecommon.models.models import OpenGuiSessionInfo from PySide6.QtCore import Qt, QUrl, Slot from PySide6.QtGui import QDesktopServices, QGuiApplication from PySide6.QtWidgets import ( @@ -27,9 +30,6 @@ from PySide6.QtWidgets import ( QWidget, ) -from aare.common.error_codes import error_code_help -from aare.common.logger_config import attach_to_logger, find_existing_formatter -from aare.common.models import OpenGuiSessionInfo from aare.gui.log import QtLogEmitter, QtLogHandler from aare.gui.threads.daq_worker import DAQWorker diff --git a/src/aare/gui/panels/face_detection_panel.py b/src/aare/gui/panels/face_detection_panel.py index 2335d0f6..12c09ebf 100644 --- a/src/aare/gui/panels/face_detection_panel.py +++ b/src/aare/gui/panels/face_detection_panel.py @@ -1,19 +1,17 @@ -from PySide6.QtCore import Signal -from PySide6.QtWidgets import QWidget, QGridLayout, QLabel, QPushButton, QVBoxLayout +import numpy as np +from aarecommon.config.logger import setup_logger from matplotlib.backends.backend_qt5agg import FigureCanvasQTAgg as FigureCanvas from matplotlib.figure import Figure -import numpy as np +from PySide6.QtCore import Signal +from PySide6.QtWidgets import QGridLayout, QLabel, QPushButton, QVBoxLayout, QWidget from aare.gui.widgets.number_line_edit import NumberLineEdit from aare.gui.widgets.title_label import TitleLabel -from aare.common.logger_config import setup_logger - logger = setup_logger("aareGUI") class FaceDetectionPanel(QWidget): - face_detection = Signal(int, int) def __init__(self, parent=None): @@ -89,9 +87,15 @@ class FaceDetectionPanel(QWidget): if running: if self._manual_run_requested: - self.status_lbl.setText(f"Running... angle {angle}" if angle is not None else "Running...") + self.status_lbl.setText( + f"Running... angle {angle}" if angle is not None else "Running..." + ) else: - self.status_lbl.setText(f"Automation running... angle {angle}" if angle is not None else "Automation running...") + self.status_lbl.setText( + f"Automation running... angle {angle}" + if angle is not None + else "Automation running..." + ) else: if self._manual_run_requested: self.status_lbl.setText("Done") @@ -126,7 +130,9 @@ class FaceDetectionPanel(QWidget): height_fit = A + B * np.cos(C * np.deg2rad(ang_grid) - phi) self.ax1.plot(ang_grid, height_fit, color="tab:orange", label="Height fit") if "best_angle_deg" in hf: - self.ax1.axvline(hf["best_angle_deg"], color="tab:orange", ls="--", alpha=0.6) + self.ax1.axvline( + hf["best_angle_deg"], color="tab:orange", ls="--", alpha=0.6 + ) af = data.get("area_fit", {}) or {} if {"A", "B", "phi_rad", "C"} <= af.keys(): @@ -134,7 +140,9 @@ class FaceDetectionPanel(QWidget): area_fit = A2 + B2 * np.cos(C2 * np.deg2rad(ang_grid) - phi2) self.ax2.plot(ang_grid, area_fit, color="tab:red", label="Area fit") if "best_angle_deg" in af: - self.ax2.axvline(af["best_angle_deg"], color="tab:red", ls="--", alpha=0.6) + self.ax2.axvline( + af["best_angle_deg"], color="tab:red", ls="--", alpha=0.6 + ) self.ax1.set_xlabel("Angle (deg)") self.ax1.set_ylabel("Height") @@ -146,4 +154,4 @@ class FaceDetectionPanel(QWidget): self.ax2.legend() self.ax2.grid(True, alpha=0.3) - self.canvas.draw_idle() \ No newline at end of file + self.canvas.draw_idle() diff --git a/src/aare/gui/panels/file_path_panel.py b/src/aare/gui/panels/file_path_panel.py index db141a7f..6cd2c730 100644 --- a/src/aare/gui/panels/file_path_panel.py +++ b/src/aare/gui/panels/file_path_panel.py @@ -2,10 +2,17 @@ import os from datetime import datetime from pathlib import Path -from PySide6.QtCore import Signal, Slot, Qt -from PySide6.QtWidgets import QWidget, QGridLayout, QLabel, QLineEdit, QSpinBox, QMessageBox +from aarecommon.models.models import DAQStatusModel, SampleShortInfo +from PySide6.QtCore import Qt, Signal, Slot +from PySide6.QtWidgets import ( + QGridLayout, + QLabel, + QLineEdit, + QMessageBox, + QSpinBox, + QWidget, +) -from aare.common.models import SampleShortInfo, DAQStatusModel from aare.gui.widgets.title_label import TitleLabel ## Logic for filenames: @@ -15,6 +22,7 @@ from aare.gui.widgets.title_label import TitleLabel ## 4. If sample is registered in the database and has puck information, it is by default placed in /// ## 5. If sample is registered in the database as manual, it is by default placed in /manual/ + class FilePathPanel(QWidget): path_updated = Signal(str) @@ -31,31 +39,35 @@ class FilePathPanel(QWidget): self.__filename = "" self.__scan_kind = "raster" # default: "rotation" | "screening" | "raster" - self.__formatted_date = datetime.now().strftime('%Y%m%d') + self.__formatted_date = datetime.now().strftime("%Y%m%d") grid_layout.addWidget(TitleLabel("Dataset path", self), 0, 0, 1, 2) grid_layout.addWidget(QLabel("Directory", parent=self), 1, 0) self.directory_edit = QLineEdit("{date}/{puck}/{pos}", parent=self) self.directory_edit.setStyleSheet("background-color: rgb(255, 255, 255);") - self.directory_edit.setToolTip("Provide subdirectory for your files. The following macros are allowed:
" - "{date} - date in format yyyymmdd
" - "{puck} - puck name
" - "{pos} - pin position in puck
" - "{sample} - sample name
" - "{sample_id} - sample id
" - "{prefix} - file prefix value") + self.directory_edit.setToolTip( + "Provide subdirectory for your files. The following macros are allowed:
" + "{date} - date in format yyyymmdd
" + "{puck} - puck name
" + "{pos} - pin position in puck
" + "{sample} - sample name
" + "{sample_id} - sample id
" + "{prefix} - file prefix value" + ) grid_layout.addWidget(self.directory_edit, 1, 1) grid_layout.addWidget(QLabel("File prefix", parent=self), 2, 0) self.file_prefix_edit = QLineEdit("{sample}", parent=self) self.file_prefix_edit.setStyleSheet("background-color: rgb(255, 255, 255);") - self.file_prefix_edit.setToolTip("Provide file prefix for your files. The following macros are allowed:
" - "{date} - date in format yyyymmdd
" - "{puck} - puck name
" - "{pos} - pin position in puck
" - "{sample} - sample name
" - "{sample_id} - sample id") + self.file_prefix_edit.setToolTip( + "Provide file prefix for your files. The following macros are allowed:
" + "{date} - date in format yyyymmdd
" + "{puck} - puck name
" + "{pos} - pin position in puck
" + "{sample} - sample name
" + "{sample_id} - sample id" + ) grid_layout.addWidget(self.file_prefix_edit, 2, 1) @@ -79,11 +91,11 @@ class FilePathPanel(QWidget): def _expand_macros(self, base: str, rn: int) -> str: name = f"{base}_{rn:03d}" - name = name.replace('{date}', self.__formatted_date) - name = name.replace('{sample}', self.__sample_name) - name = name.replace('{puck}', self.__puck_name) - name = name.replace('{pos}', f"{self.__puck_pos:02d}") - name = name.replace('{sample_id}', f"{self.__sample_id}") + name = name.replace("{date}", self.__formatted_date) + name = name.replace("{sample}", self.__sample_name) + name = name.replace("{puck}", self.__puck_name) + name = name.replace("{pos}", f"{self.__puck_pos:02d}") + name = name.replace("{sample_id}", f"{self.__sample_id}") return name def _effective_dataset_base(self, base_no_run: str) -> str: @@ -111,7 +123,7 @@ class FilePathPanel(QWidget): file_prefix = self.file_prefix_edit.text() run_number = self.run_number_edit.value() - dir_name = dir_name.replace('{prefix}', file_prefix) + dir_name = dir_name.replace("{prefix}", file_prefix) if dir_name == "": base = "" @@ -144,7 +156,9 @@ class FilePathPanel(QWidget): effective = self._effective_dataset_base(self.__filename) exists = os.path.exists(f"{effective}_master.h5") or os.path.exists(effective) self.file_name_label.setText(effective + "_master.h5") - self.file_name_label.setStyleSheet("color: rgb(200, 0, 0);" if exists else "color: rgb(0, 0, 0);") + self.file_name_label.setStyleSheet( + "color: rgb(200, 0, 0);" if exists else "color: rgb(0, 0, 0);" + ) self.path_updated.emit(self.__filename) @Slot() @@ -189,12 +203,19 @@ class FilePathPanel(QWidget): self.__puck_pos = sample.pin self.run_number_edit.setValue(1) - if sample.aaredb_params is not None and sample.aaredb_params.directory is not None: + if ( + sample.aaredb_params is not None + and sample.aaredb_params.directory is not None + ): self.directory_edit.setText(f"{sample.aaredb_params.directory}") elif sample.location is not None: - self.directory_edit.setText(f"{self.__formatted_date}/{self.__puck_name}/{self.__puck_pos:02d}") + self.directory_edit.setText( + f"{self.__formatted_date}/{self.__puck_name}/{self.__puck_pos:02d}" + ) else: - self.directory_edit.setText(f"{self.__formatted_date}/manual/{self.__sample_name}") + self.directory_edit.setText( + f"{self.__formatted_date}/manual/{self.__sample_name}" + ) self.update_filename() @Slot(DAQStatusModel) @@ -213,9 +234,15 @@ class FilePathPanel(QWidget): def next_free_run_from(self, start_rn: int) -> tuple[int, str]: # Compute next free run number and updated base - dir_name = self.directory_edit.text().replace('{prefix}', self.file_prefix_edit.text()) - base = (dir_name if dir_name.endswith("/") or dir_name == "" else dir_name + "/") - base += ("run" if self.file_prefix_edit.text() == "" else self.file_prefix_edit.text()) + dir_name = self.directory_edit.text().replace( + "{prefix}", self.file_prefix_edit.text() + ) + base = dir_name if dir_name.endswith("/") or dir_name == "" else dir_name + "/" + base += ( + "run" + if self.file_prefix_edit.text() == "" + else self.file_prefix_edit.text() + ) rn = start_rn while rn <= self.run_number_edit.maximum(): @@ -233,10 +260,10 @@ class FilePathPanel(QWidget): reply = QMessageBox.question( self, "File exists", - #f"This file already exists:\n{effective}_master.h5\nDo you wish to overwrite?", + # f"This file already exists:\n{effective}_master.h5\nDo you wish to overwrite?", f"This file already exists:\n{effective}_master.h5\n Updating run number", QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No, - QMessageBox.StandardButton.No + QMessageBox.StandardButton.No, ) if reply == QMessageBox.StandardButton.No: # bump run and update @@ -249,4 +276,4 @@ class FilePathPanel(QWidget): return False return False - return True \ No newline at end of file + return True diff --git a/src/aare/gui/panels/fluorescence_data_collection.py b/src/aare/gui/panels/fluorescence_data_collection.py index df68c90c..e7d57ea8 100644 --- a/src/aare/gui/panels/fluorescence_data_collection.py +++ b/src/aare/gui/panels/fluorescence_data_collection.py @@ -1,7 +1,7 @@ +from aarecommon.models.models import FluorescenceSpectrumParameterModel from PySide6.QtCore import Signal, Slot -from PySide6.QtWidgets import QWidget, QGridLayout, QLabel, QPushButton, QCheckBox +from PySide6.QtWidgets import QCheckBox, QGridLayout, QLabel, QPushButton, QWidget -from aare.common.models import FluorescenceSpectrumParameterModel from aare.gui.widgets.number_line_edit import NumberLineEdit @@ -15,7 +15,9 @@ class FluorescenceDataCollectionPanel(QWidget): # Beam transmission (0..1) lay.addWidget(QLabel("Beam transmission", self), 0, 0) - self.transmission = NumberLineEdit(0.0, 1.0, decimals=4, default=0.1, parent=self) + self.transmission = NumberLineEdit( + 0.0, 1.0, decimals=4, default=0.1, parent=self + ) lay.addWidget(self.transmission, 0, 1) # Exposure time (seconds) @@ -25,7 +27,9 @@ class FluorescenceDataCollectionPanel(QWidget): lay.addWidget(QLabel("s", self), 1, 2) # Accumulate checkbox (reverse of erase) - self.accumulate_cb = QCheckBox("Accumulate", self) # accumulate=True => erase=False + self.accumulate_cb = QCheckBox( + "Accumulate", self + ) # accumulate=True => erase=False self.accumulate_cb.setChecked(False) lay.addWidget(self.accumulate_cb, 2, 0, 1, 3) @@ -44,4 +48,8 @@ class FluorescenceDataCollectionPanel(QWidget): t = float(self.transmission.value) exp = float(self.exposure.value) erase = not self.accumulate_cb.isChecked() - self.fluo_scan.emit(FluorescenceSpectrumParameterModel(acq_time_s=exp, transmission=t, erase=erase)) + self.fluo_scan.emit( + FluorescenceSpectrumParameterModel( + acq_time_s=exp, transmission=t, erase=erase + ) + ) diff --git a/src/aare/gui/panels/fluorescence_panel.py b/src/aare/gui/panels/fluorescence_panel.py index a5849ee5..28135d08 100644 --- a/src/aare/gui/panels/fluorescence_panel.py +++ b/src/aare/gui/panels/fluorescence_panel.py @@ -1,14 +1,14 @@ import numpy as np +from aarecommon.config.logger import setup_logger +from aarecommon.models.models import DAQStatusModel, FluorescenceSpectrumOutputModel from PySide6.QtCharts import QChart, QChartView, QLineSeries, QValueAxis -from PySide6.QtCore import QPointF, Qt, Slot, QEvent -from PySide6.QtGui import QPainter, QColor, QPen -from PySide6.QtWidgets import QWidget, QGridLayout, QGraphicsSimpleTextItem, QLabel - -from aare.common.logger_config import setup_logger -from aare.common.models import FluorescenceSpectrumOutputModel, DAQStatusModel +from PySide6.QtCore import QEvent, QPointF, Qt, Slot +from PySide6.QtGui import QColor, QPainter, QPen +from PySide6.QtWidgets import QGraphicsSimpleTextItem, QGridLayout, QLabel, QWidget logger = setup_logger("aareGUI") + class FluorescencePanel(QWidget): def __init__(self, parent=None): super().__init__(parent) @@ -78,7 +78,10 @@ class FluorescencePanel(QWidget): def eventFilter(self, obj, event): try: - if obj is self.chart_view.viewport() and event.type() == QEvent.Type.MouseMove: + if ( + obj is self.chart_view.viewport() + and event.type() == QEvent.Type.MouseMove + ): pos = event.position() if hasattr(event, "position") else event.pos() p = QPointF(pos.x(), pos.y()) plot = self.chart.plotArea() @@ -88,7 +91,8 @@ class FluorescencePanel(QWidget): return False # Map pixel X -> chart X x_val = self.axis_x.min() + (self.axis_x.max() - self.axis_x.min()) * ( - (p.x() - plot.left()) / plot.width()) + (p.x() - plot.left()) / plot.width() + ) # Snap to the largest Y within +/- 3 indices around nearest index center = self._nearest_index(x_val) @@ -105,7 +109,9 @@ class FluorescencePanel(QWidget): best_i = i pt = self.series.at(best_i) - self.chart_view.setToolTip(f"Energy {pt.x():.3f} keV counts {pt.y():.3f}") + self.chart_view.setToolTip( + f"Energy {pt.x():.3f} keV counts {pt.y():.3f}" + ) self._update_vline(pt.x()) return False except Exception as e: @@ -114,7 +120,8 @@ class FluorescencePanel(QWidget): def _show_context_menu(self, pos): try: - from PySide6.QtWidgets import QMenu, QFileDialog + from PySide6.QtWidgets import QFileDialog, QMenu + menu = QMenu(self.chart_view) act_save_csv = menu.addAction("Save spectrum as CSV...") global_pos = self.chart_view.mapToGlobal(pos) @@ -127,7 +134,11 @@ class FluorescencePanel(QWidget): # Prefer cached arrays; fall back to reading from series x_vals = self._last_x_keV y_vals = self._last_y_counts - if x_vals is None or y_vals is None or len(x_vals) != self.series.count(): + if ( + x_vals is None + or y_vals is None + or len(x_vals) != self.series.count() + ): x_vals = [self.series.at(i).x() for i in range(self.series.count())] y_vals = [self.series.at(i).y() for i in range(self.series.count())] @@ -135,11 +146,14 @@ class FluorescencePanel(QWidget): if self._sample_name is not None: default_name = f"{self._sample_name}_spectrum.csv" - path, _ = QFileDialog.getSaveFileName(self, "Save Spectrum CSV", default_name, "CSV files (*.csv)") + path, _ = QFileDialog.getSaveFileName( + self, "Save Spectrum CSV", default_name, "CSV files (*.csv)" + ) if not path: return try: import csv + with open(path, "w", newline="") as f: writer = csv.writer(f) writer.writerow(["Energy_keV", "Counts"]) @@ -156,7 +170,9 @@ class FluorescencePanel(QWidget): try: ymin = self.axis_y.min() ymax = self.axis_y.max() - self._vline.replace(0, x_val, ymin) if self._vline.count() > 0 else self._vline.append(x_val, ymin) + self._vline.replace( + 0, x_val, ymin + ) if self._vline.count() > 0 else self._vline.append(x_val, ymin) if self._vline.count() == 1: self._vline.append(x_val, ymax) else: @@ -198,7 +214,9 @@ class FluorescencePanel(QWidget): # Update average dead time label (if provided) try: if hasattr(f, "average_dead_time") and f.average_dead_time is not None: - self.avg_dead_label.setText(f"Average dead time: {f.average_dead_time * 100.0:.3f}%") + self.avg_dead_label.setText( + f"Average dead time: {f.average_dead_time * 100.0:.3f}%" + ) else: self.avg_dead_label.setText("Average dead time: -") except Exception: @@ -226,7 +244,9 @@ class FluorescencePanel(QWidget): max_idx = int(np.argmax(y)) peak_keV = float(x[max_idx]) peak_counts = float(y[max_idx]) - self.peak_info_label.setText(f"Peak: {peak_keV:.3f} keV, {peak_counts:.3f} counts") + self.peak_info_label.setText( + f"Peak: {peak_keV:.3f} keV, {peak_counts:.3f} counts" + ) # Axes xmin = float(np.min(x)) if x.size else -1.0 xmax = float(np.max(x)) if x.size else 1.0 diff --git a/src/aare/gui/panels/illumination_panel.py b/src/aare/gui/panels/illumination_panel.py index 376b2c94..7d624ac5 100644 --- a/src/aare/gui/panels/illumination_panel.py +++ b/src/aare/gui/panels/illumination_panel.py @@ -1,6 +1,6 @@ +from aarecommon.models.models import DAQStatusModel from PySide6.QtCore import Qt, Signal, Slot -from PySide6.QtWidgets import QWidget, QGridLayout, QSlider, QLabel -from aare.common.models import DAQStatusModel +from PySide6.QtWidgets import QGridLayout, QLabel, QSlider, QWidget from aare.gui.widgets.title_label import TitleLabel @@ -20,7 +20,9 @@ class IlluminationPanel(QWidget): grid_layout.addWidget(front_label, 1, 0, 1, 2) self.is_sliding = False - self.front_light_slider = QSlider(orientation=Qt.Orientation.Horizontal, parent=self) + self.front_light_slider = QSlider( + orientation=Qt.Orientation.Horizontal, parent=self + ) self.front_light_slider.setRange(0, 100) self.front_light_slider.sliderPressed.connect(self.on_slider_pressed) self.front_light_slider.sliderReleased.connect(self.on_front_slider_released) @@ -30,7 +32,9 @@ class IlluminationPanel(QWidget): back_label.setAlignment(Qt.AlignmentFlag.AlignCenter) grid_layout.addWidget(back_label, 3, 0, 1, 2) - self.back_light_slider = QSlider(orientation=Qt.Orientation.Horizontal, parent=self) + self.back_light_slider = QSlider( + orientation=Qt.Orientation.Horizontal, parent=self + ) self.back_light_slider.setRange(0, 100) self.back_light_slider.sliderPressed.connect(self.on_slider_pressed) self.back_light_slider.sliderReleased.connect(self.on_back_slider_released) diff --git a/src/aare/gui/panels/local_contact_panel.py b/src/aare/gui/panels/local_contact_panel.py index 51921011..9c158c08 100644 --- a/src/aare/gui/panels/local_contact_panel.py +++ b/src/aare/gui/panels/local_contact_panel.py @@ -2,12 +2,15 @@ from __future__ import annotations from collections.abc import Callable -from PySide6.QtCore import Slot, QUrl +from aarecommon.config.logger import setup_logger +from aarecommon.models.models import DAQStatusModel +from PySide6.QtCore import QUrl, Slot from PySide6.QtGui import QDesktopServices from PySide6.QtWidgets import ( QCheckBox, QDialog, QDialogButtonBox, + QDoubleSpinBox, QFrame, QGridLayout, QGroupBox, @@ -19,17 +22,15 @@ from PySide6.QtWidgets import ( QTabWidget, QTextEdit, QVBoxLayout, - QWidget, QDoubleSpinBox, + QWidget, ) -from aare.common.logger_config import setup_logger -from aare.common.models import DAQStatusModel from aare.gui.panels.beamline_recovery_panel import RecoveryPanel from aare.gui.threads.daq_worker import DAQWorker from aare.gui.widgets.local_contact_status_widget import LocalContactStatusWidget +from aare.gui.widgets.number_line_edit import NumberLineEdit from aare.gui.widgets.text_list_dialog import TextListDialog from aare.gui.widgets.title_label import TitleLabel -from aare.gui.widgets.number_line_edit import NumberLineEdit logger = setup_logger("aareGUI") @@ -116,7 +117,9 @@ class LocalContactPanel(QFrame): transfer_error_layout = QVBoxLayout(self._transfer_error_frame) transfer_error_layout.setContentsMargins(10, 10, 10, 10) - self._transfer_error_title = QLabel("Error transferring information from DAQ", self._transfer_error_frame) + self._transfer_error_title = QLabel( + "Error transferring information from DAQ", self._transfer_error_frame + ) self._transfer_error_title.setStyleSheet("font-weight: 700;") transfer_error_layout.addWidget(self._transfer_error_title) @@ -138,12 +141,16 @@ class LocalContactPanel(QFrame): self._tabs.addTab(self._build_detector_tab(), self.TAB_DETECTOR) self._tabs.addTab(self._build_config_tab(), self.TAB_CONFIG) - self._daq.local_contact_simulation_state_loaded.connect(self._apply_simulation_state) + self._daq.local_contact_simulation_state_loaded.connect( + self._apply_simulation_state + ) self._daq.local_contact_device_state_loaded.connect(self._apply_device_state) self._daq.local_contact_links_loaded.connect(self._apply_links) self._daq.local_contact_config_loaded.connect(self._apply_local_contact_config) self._daq.local_contact_config_saved.connect(self._apply_local_contact_config) - self._daq.local_contact_transfer_error.connect(self._show_local_contact_transfer_error) + self._daq.local_contact_transfer_error.connect( + self._show_local_contact_transfer_error + ) self._daq.bec_user_macros_loaded.connect(self._show_bec_user_macros) self._daq.bec_devices_loaded.connect(self._show_bec_devices) self._daq.update.connect(self._update_from_status) @@ -155,12 +162,17 @@ class LocalContactPanel(QFrame): def set_active_tab(self, tab_name: str) -> None: for index in range(self._tabs.count()): - if self._tabs.tabText(index).strip().lower() == str(tab_name).strip().lower(): + if ( + self._tabs.tabText(index).strip().lower() + == str(tab_name).strip().lower() + ): self._tabs.setCurrentIndex(index) return logger.warning(f"Unknown Local Contact tab requested: {tab_name}") - def _register_status_widget(self, widget: LocalContactStatusWidget) -> LocalContactStatusWidget: + def _register_status_widget( + self, widget: LocalContactStatusWidget + ) -> LocalContactStatusWidget: self._status_widgets.append(widget) if self._last_status is not None: widget.set_daq_status(self._last_status) @@ -251,7 +263,13 @@ class LocalContactPanel(QFrame): self._make_status_widget( title="Recovery status", summary="Recovery-relevant DAQ state.", - fields=("beamline_state", "busy", "sample", "tell_connected", "tell_state"), + fields=( + "beamline_state", + "busy", + "sample", + "tell_connected", + "tell_state", + ), ) ) layout.addWidget(RecoveryPanel(daq=self._daq, parent=tab), 1) @@ -267,7 +285,14 @@ class LocalContactPanel(QFrame): self._make_status_widget( title="TELL status", summary="TELL-focused state and connection details.", - fields=("tell", "tell_connected", "tell_state", "tell_error", "busy", "beamline_state"), + fields=( + "tell", + "tell_connected", + "tell_state", + "tell_error", + "busy", + "beamline_state", + ), ) ) @@ -277,16 +302,38 @@ class LocalContactPanel(QFrame): row = 0 row += 1 - grid.addWidget(self._make_button("Unmount", self._daq.unmount, "Requesting sample unmount."), row, 0) - grid.addWidget(self._make_button("Dry", self._daq.tell_dry, "Requesting TELL dry."), row, 1) + grid.addWidget( + self._make_button( + "Unmount", self._daq.unmount, "Requesting sample unmount." + ), + row, + 0, + ) + grid.addWidget( + self._make_button("Dry", self._daq.tell_dry, "Requesting TELL dry."), row, 1 + ) row += 1 - grid.addWidget(self._make_button("Park and dry", self._daq.park_and_dry, "Requesting park and dry."), row, 0) - grid.addWidget(self._make_button("Toggle blower", self._daq.tell_toggle_blower, "Toggling blower."), row, 1) + grid.addWidget( + self._make_button( + "Park and dry", self._daq.park_and_dry, "Requesting park and dry." + ), + row, + 0, + ) + grid.addWidget( + self._make_button( + "Toggle blower", self._daq.tell_toggle_blower, "Toggling blower." + ), + row, + 1, + ) row += 1 grid.addWidget(self._make_button("Anneal", self._anneal_from_dialog), row, 0) - grid.addWidget(self._make_button("TELL access info", self._show_tell_access_info), row, 1) + grid.addWidget( + self._make_button("TELL access info", self._show_tell_access_info), row, 1 + ) wrapper = QGroupBox("TELL actions", tab) wrapper_layout = QVBoxLayout(wrapper) @@ -309,7 +356,14 @@ class LocalContactPanel(QFrame): self._make_status_widget( title="BEC status", summary="BEC-focused state and related backend status.", - fields=("bec", "busy", "beamline_state", "detector", "aerotech", "smargon"), + fields=( + "bec", + "busy", + "beamline_state", + "detector", + "aerotech", + "smargon", + ), ) ) @@ -321,11 +375,26 @@ class LocalContactPanel(QFrame): tools_layout.addWidget(tools) tools_layout.addWidget( - self._make_button("Load BEC user macros", self._daq.bec_load_user_macros, "Loading BEC user macros.")) + self._make_button( + "Load BEC user macros", + self._daq.bec_load_user_macros, + "Loading BEC user macros.", + ) + ) tools_layout.addWidget( - self._make_button("Show BEC user macros", self._daq.bec_list_all_user_macros, "Listing BEC user macros.")) + self._make_button( + "Show BEC user macros", + self._daq.bec_list_all_user_macros, + "Listing BEC user macros.", + ) + ) tools_layout.addWidget( - self._make_button("Show BEC position devices", self._daq.bec_list_all_devices, "Listing BEC devices.")) + self._make_button( + "Show BEC position devices", + self._daq.bec_list_all_devices, + "Listing BEC devices.", + ) + ) tools_layout.addWidget( self._make_button( "Reinitialise BEC planner/devices", @@ -376,7 +445,15 @@ class LocalContactPanel(QFrame): self._make_status_widget( title="Hardware status", summary="Hardware-related backend and beamline status.", - fields=("aerotech", "smargon", "detector", "bec", "tell", "busy", "beamline_state"), + fields=( + "aerotech", + "smargon", + "detector", + "bec", + "tell", + "busy", + "beamline_state", + ), ) ) @@ -425,9 +502,21 @@ class LocalContactPanel(QFrame): self._build_section( "Initialise", [ - self._make_button("Initialise detector", self._daq.initialise_detector, "Initialising detector."), - self._make_button("Initialise Smargon", self._daq.initialise_smargon, "Initialising Smargon."), - self._make_button("Initialise Aerotech", self._daq.initialise_aerotech, "Initialising Aerotech."), + self._make_button( + "Initialise detector", + self._daq.initialise_detector, + "Initialising detector.", + ), + self._make_button( + "Initialise Smargon", + self._daq.initialise_smargon, + "Initialising Smargon.", + ), + self._make_button( + "Initialise Aerotech", + self._daq.initialise_aerotech, + "Initialising Aerotech.", + ), ], ) ) @@ -445,7 +534,15 @@ class LocalContactPanel(QFrame): self._make_status_widget( title="Detector status", summary="Detector-focused state and metadata.", - fields=("detector", "bec", "busy", "beamline_state", "dtz", "detector_description", "detector_serial"), + fields=( + "detector", + "bec", + "busy", + "beamline_state", + "dtz", + "detector_description", + "detector_serial", + ), ) ) @@ -474,11 +571,13 @@ class LocalContactPanel(QFrame): row, 1, ) - grid.addWidget(self._make_button( - "Resync detector/DTZ hardware cache", - self._daq.resync_local_contact_detector_metadata, - "Resyncing detector metadata and DTZ limits cache.", - )) + grid.addWidget( + self._make_button( + "Resync detector/DTZ hardware cache", + self._daq.resync_local_contact_detector_metadata, + "Resyncing detector metadata and DTZ limits cache.", + ) + ) wrapper = QGroupBox("Detector actions", tab) wrapper_layout = QVBoxLayout(wrapper) @@ -523,8 +622,12 @@ class LocalContactPanel(QFrame): form_layout.addWidget(label, 0, 0) form_layout.addWidget(self._mount_to_center_sleep_s, 0, 1) - loop_face_padding_label = QLabel("Line scan loop_face Y padding / side", form_box) - self._line_scan_loop_face_y_padding_fraction_each_side = QDoubleSpinBox(form_box) + loop_face_padding_label = QLabel( + "Line scan loop_face Y padding / side", form_box + ) + self._line_scan_loop_face_y_padding_fraction_each_side = QDoubleSpinBox( + form_box + ) self._line_scan_loop_face_y_padding_fraction_each_side.setRange(0.0, 5.0) self._line_scan_loop_face_y_padding_fraction_each_side.setDecimals(2) self._line_scan_loop_face_y_padding_fraction_each_side.setSingleStep(0.05) @@ -544,9 +647,13 @@ class LocalContactPanel(QFrame): ) form_layout.addWidget(loop_face_padding_label, 1, 0) - form_layout.addWidget(self._line_scan_loop_face_y_padding_fraction_each_side, 1, 1) + form_layout.addWidget( + self._line_scan_loop_face_y_padding_fraction_each_side, 1, 1 + ) form_layout.addWidget(loop_all_padding_label, 2, 0) - form_layout.addWidget(self._line_scan_loop_all_y_padding_fraction_each_side, 2, 1) + form_layout.addWidget( + self._line_scan_loop_all_y_padding_fraction_each_side, 2, 1 + ) button_row = QWidget(tab) button_layout = QHBoxLayout(button_row) @@ -590,7 +697,9 @@ class LocalContactPanel(QFrame): name.setStyleSheet("font-weight: 700;") checkbox = QCheckBox(text, row_widget) - checkbox.toggled.connect(lambda checked, d=device: self._daq.set_local_contact_simulation(d, checked)) + checkbox.toggled.connect( + lambda checked, d=device: self._daq.set_local_contact_simulation(d, checked) + ) self._sim_checkboxes[device] = checkbox row_layout.addWidget(name) @@ -608,10 +717,12 @@ class LocalContactPanel(QFrame): name.setStyleSheet("font-weight: 700;") button = QPushButton(text, row_widget) - button.clicked.connect(lambda: self._run_logged_action( - f"Restarting {device} backend.", - lambda: self._daq.restart_local_contact_device(device), - )) + button.clicked.connect( + lambda: self._run_logged_action( + f"Restarting {device} backend.", + lambda: self._daq.restart_local_contact_device(device), + ) + ) row_layout.addWidget(name) row_layout.addWidget(button, 1) @@ -649,7 +760,9 @@ class LocalContactPanel(QFrame): @Slot(list) def _show_bec_user_macros(self, items: list) -> None: if self._bec_macros_dialog is None: - self._bec_macros_dialog = TextListDialog(title="BEC user macros", parent=self) + self._bec_macros_dialog = TextListDialog( + title="BEC user macros", parent=self + ) self._bec_macros_dialog.set_items( [str(item) for item in items], empty_message="No BEC user macros found.", @@ -661,7 +774,9 @@ class LocalContactPanel(QFrame): @Slot(list) def _show_bec_devices(self, items: list) -> None: if self._bec_devices_dialog is None: - self._bec_devices_dialog = TextListDialog(title="BEC position devices", parent=self) + self._bec_devices_dialog = TextListDialog( + title="BEC position devices", parent=self + ) self._bec_devices_dialog.set_items( [str(item) for item in items], empty_message="No BEC devices found.", @@ -691,7 +806,10 @@ class LocalContactPanel(QFrame): @Slot() def _show_tell_access_info(self) -> None: - message = self._links_payload.get("tell_hint") or "Please check TELL status via Remmina / VNC." + message = ( + self._links_payload.get("tell_hint") + or "Please check TELL status via Remmina / VNC." + ) QMessageBox.information(self, "TELL access", str(message)) def _make_button( @@ -704,7 +822,9 @@ class LocalContactPanel(QFrame): if log_message is None: button.clicked.connect(callback) else: - button.clicked.connect(lambda: self._run_logged_action(log_message, callback)) + button.clicked.connect( + lambda: self._run_logged_action(log_message, callback) + ) return button @Slot() @@ -733,7 +853,9 @@ class LocalContactPanel(QFrame): self._mount_to_center_sleep_s.blockSignals(True) self._mount_to_center_sleep_s.setValue( - float(self._local_contact_config_payload.get("mount_to_center_sleep_s", 0.0)) + float( + self._local_contact_config_payload.get("mount_to_center_sleep_s", 0.0) + ) ) self._mount_to_center_sleep_s.blockSignals(False) @@ -792,4 +914,4 @@ class LocalContactDialog(QDialog): layout.addWidget(buttons) def set_active_tab(self, tab_name: str) -> None: - self._panel.set_active_tab(tab_name) \ No newline at end of file + self._panel.set_active_tab(tab_name) diff --git a/src/aare/gui/panels/manual_sample_panel.py b/src/aare/gui/panels/manual_sample_panel.py index bb0bf81a..64fe051c 100644 --- a/src/aare/gui/panels/manual_sample_panel.py +++ b/src/aare/gui/panels/manual_sample_panel.py @@ -1,8 +1,15 @@ -from PySide6.QtCore import Signal, Slot -from PySide6.QtWidgets import QWidget, QGridLayout, QLabel, QPushButton, QCheckBox, QLineEdit +from aarecommon.models.models import DAQStatusModel, SampleShortInfo from aareDB import DataCollectionParameters +from PySide6.QtCore import Signal, Slot +from PySide6.QtWidgets import ( + QCheckBox, + QGridLayout, + QLabel, + QLineEdit, + QPushButton, + QWidget, +) -from aare.common.models import SampleShortInfo, DAQStatusModel from aare.gui.widgets.number_line_edit import NumberLineEdit from aare.gui.widgets.title_label import TitleLabel @@ -104,7 +111,7 @@ class ManualSamplePanel(QWidget): if self._unit_cell.isChecked(): data_processing = DataCollectionParameters( spacegroupnumber=int(self._text_sg.value), - cellparameters=f"{self._text_a.value:.2f} {self._text_b.value:.2f} {self._text_c.value:.2f} {self._text_alpha.value:.2f} {self._text_beta.value:.2f} {self._text_gamma.value:.2f}" + cellparameters=f"{self._text_a.value:.2f} {self._text_b.value:.2f} {self._text_c.value:.2f} {self._text_alpha.value:.2f} {self._text_beta.value:.2f} {self._text_gamma.value:.2f}", ) s = SampleShortInfo( @@ -112,7 +119,7 @@ class ManualSamplePanel(QWidget): puck_name="", dewar_name="", sample_name=self.__sample_name, - run_number=1, #can't be -1. So have defaulted to 1 for now. + run_number=1, # can't be -1. So have defaulted to 1 for now. pin=1, aaredb_params=data_processing, user=self.__pgroup, diff --git a/src/aare/gui/panels/monochromator_panel.py b/src/aare/gui/panels/monochromator_panel.py index 2164165f..ed637c87 100644 --- a/src/aare/gui/panels/monochromator_panel.py +++ b/src/aare/gui/panels/monochromator_panel.py @@ -1,7 +1,7 @@ +from aarecommon.models.models import DAQStatusModel from PySide6.QtCore import Signal, Slot -from PySide6.QtWidgets import QWidget, QGridLayout, QLabel, QDoubleSpinBox, QPushButton +from PySide6.QtWidgets import QDoubleSpinBox, QGridLayout, QLabel, QPushButton, QWidget -from aare.common.models import DAQStatusModel from aare.gui.widgets.title_label import TitleLabel diff --git a/src/aare/gui/panels/omega_panel.py b/src/aare/gui/panels/omega_panel.py index e578f82c..2c54fe53 100644 --- a/src/aare/gui/panels/omega_panel.py +++ b/src/aare/gui/panels/omega_panel.py @@ -1,6 +1,6 @@ +from aarecommon.models.models import DAQStatusModel from PySide6.QtCore import Signal, Slot -from PySide6.QtWidgets import QWidget, QGridLayout, QLabel, QHBoxLayout -from aare.common.models import DAQStatusModel +from PySide6.QtWidgets import QGridLayout, QHBoxLayout, QLabel, QWidget from aare.gui.widgets.button_with_payload import ButtonWithPayload from aare.gui.widgets.number_line_edit import NumberLineEdit diff --git a/src/aare/gui/panels/portrait_mode.py b/src/aare/gui/panels/portrait_mode.py index d5ae693a..b3f1e145 100644 --- a/src/aare/gui/panels/portrait_mode.py +++ b/src/aare/gui/panels/portrait_mode.py @@ -2,34 +2,50 @@ from __future__ import annotations import math -from PySide6.QtCore import Qt, QPointF, QRectF, Signal, Slot, QTimer +from aarecommon.config.logger import setup_logger +from aarecommon.models.automation import ( + AutomationProgress, + StepStatus, + WorkflowStateKind, +) +from PySide6.QtCore import QPointF, QRectF, Qt, QTimer, Signal, Slot from PySide6.QtGui import ( - QColor, QFont, QFontMetrics, QLinearGradient, QPainter, QPen, + QColor, + QFont, + QFontMetrics, + QLinearGradient, + QPainter, + QPen, ) from PySide6.QtWidgets import ( - QFrame, QHBoxLayout, QLabel, QPushButton, QScrollArea, - QSizePolicy, QStackedWidget, QVBoxLayout, QWidget, + QFrame, + QHBoxLayout, + QLabel, + QPushButton, + QScrollArea, + QSizePolicy, QSpinBox, + QStackedWidget, + QVBoxLayout, + QWidget, ) -from aare.common.automation_models import AutomationProgress, StepStatus, WorkflowStateKind -from aare.common.logger_config import setup_logger - -logger = setup_logger('aareGUI') +logger = setup_logger("aareGUI") # --------------------------------------------------------------------------- # Colour palette (kept identical to gui_designer.py) # --------------------------------------------------------------------------- -BG = "#071018" -CARD_BG = "#0E1A26" -ACCENT = "#62D8C8" +BG = "#071018" +CARD_BG = "#0E1A26" +ACCENT = "#62D8C8" ACCENT_DIM = "#1A3A36" -TEXT = "#F5F7FA" -SUBTEXT = "#8A9BB0" -BUTTON_BG = "#132131" -LED_OFF = "#1C2E3E" +TEXT = "#F5F7FA" +SUBTEXT = "#8A9BB0" +BUTTON_BG = "#132131" +LED_OFF = "#1C2E3E" ACTIVE_STEP = "#FFFFFF" + # --------------------------------------------------------------------------- # LED step indicator # --------------------------------------------------------------------------- @@ -334,7 +350,9 @@ class PortraitModePanel(QWidget): # Camera card — wraps the real compact_sample_camera cam_card = QFrame() - cam_card.setStyleSheet(f"QFrame {{ background: {CARD_BG}; border-radius: 18px; }}") + cam_card.setStyleSheet( + f"QFrame {{ background: {CARD_BG}; border-radius: 18px; }}" + ) cam_card_layout = QVBoxLayout(cam_card) cam_card_layout.setContentsMargins(4, 4, 4, 4) cam_card_layout.setSpacing(0) @@ -479,7 +497,9 @@ class PortraitModePanel(QWidget): layout.addWidget(self._queue_scroll, stretch=1) back_btn = self._accent_button("← BACK TO CAMERA") - back_btn.clicked.connect(lambda: self._stack.setCurrentWidget(self._player_page)) + back_btn.clicked.connect( + lambda: self._stack.setCurrentWidget(self._player_page) + ) layout.addWidget(back_btn) return page @@ -492,9 +512,7 @@ class PortraitModePanel(QWidget): self._job_list_panel = job_list_panel self._tell_samples = tell_samples - self._job_list_panel.loop_restart_requested = ( - self._restart_loop_if_needed - ) + self._job_list_panel.loop_restart_requested = self._restart_loop_if_needed self._populate_queue_from_tell_samples_if_empty() self.refresh_queue_preview() @@ -709,7 +727,9 @@ class PortraitModePanel(QWidget): self._alert_toast.setVisible(False) self._alert_toast_label.clear() - def _flush_portrait_alerts_to_banners(self, primary_banner, secondary_banner) -> None: + def _flush_portrait_alerts_to_banners( + self, primary_banner, secondary_banner + ) -> None: """ Called when returning to main view — replay any error alerts that arrived during portrait mode so the operator doesn't miss them. @@ -727,9 +747,7 @@ class PortraitModePanel(QWidget): self._loop_enabled = enabled if enabled and self._job_list_panel: - self._loop_samples = list( - self._job_list_panel.table_model.samples - ) + self._loop_samples = list(self._job_list_panel.table_model.samples) self._loop_remaining = self._loop_count.value() else: self._loop_remaining = 0 @@ -745,8 +763,8 @@ class PortraitModePanel(QWidget): # Timeout protection if ( - self._loop_restart_deadline is not None - and self._loop_restart_deadline.hasExpired() + self._loop_restart_deadline is not None + and self._loop_restart_deadline.hasExpired() ): self._loop_restart_timer.stop() @@ -761,9 +779,9 @@ class PortraitModePanel(QWidget): # Still busy, wait if getattr( - self._job_list_panel, - "_SampleQueuePanel__busy", - False, + self._job_list_panel, + "_SampleQueuePanel__busy", + False, ): return @@ -795,12 +813,11 @@ class PortraitModePanel(QWidget): self._loop_remaining -= 1 # Start waiting for idle - self._loop_restart_deadline = ( - QTimer().remainingTime() - ) + self._loop_restart_deadline = QTimer().remainingTime() # 30 second safety timeout from PySide6.QtCore import QDeadlineTimer + self._loop_restart_deadline = QDeadlineTimer(30000) self._loop_restart_timer.start() @@ -822,24 +839,15 @@ class PortraitModePanel(QWidget): logger.info("Queue already exists, not populating from tell_samples") return - samples = list( - getattr(self._tell_samples.table_model, "samples", []) - ) + samples = list(getattr(self._tell_samples.table_model, "samples", [])) ordered = sorted( - [ - s for s in samples - if getattr(s, "location", None) is not None - ], - key=lambda s: ( - s.loc_str_sort() - if hasattr(s, "loc_str_sort") - else "" - ), + [s for s in samples if getattr(s, "location", None) is not None], + key=lambda s: s.loc_str_sort() if hasattr(s, "loc_str_sort") else "", ) if ordered: self._job_list_panel.queue_samples( ordered, replace=True, - ) \ No newline at end of file + ) diff --git a/src/aare/gui/panels/prediction_metrics_panel.py b/src/aare/gui/panels/prediction_metrics_panel.py index 895bb07e..50de44e6 100644 --- a/src/aare/gui/panels/prediction_metrics_panel.py +++ b/src/aare/gui/panels/prediction_metrics_panel.py @@ -14,33 +14,32 @@ import time from collections import deque from dataclasses import dataclass, field +from aarecommon.config.logger import setup_logger +from aarecommon.models.models import MLBoxType from PySide6.QtCharts import ( - QChart, - QChartView, + QBarCategoryAxis, QBarSeries, QBarSet, - QBarCategoryAxis, - QValueAxis, + QChart, + QChartView, QLineSeries, + QValueAxis, ) -from PySide6.QtCore import Qt, Slot, QTimer -from PySide6.QtGui import QPainter, QColor, QPen +from PySide6.QtCore import Qt, QTimer, Slot +from PySide6.QtGui import QColor, QPainter, QPen from PySide6.QtWidgets import ( - QWidget, - QVBoxLayout, + QFileDialog, + QGridLayout, + QGroupBox, QHBoxLayout, QLabel, QPushButton, - QGroupBox, - QGridLayout, QSpinBox, - QFileDialog, QSplitter, + QVBoxLayout, + QWidget, ) -from aare.common.logger_config import setup_logger -from aare.common.models import MLBoxType - logger = setup_logger("aareGUI") @@ -52,6 +51,7 @@ logger = setup_logger("aareGUI") @dataclass class PredictionFrame: """Single frame of prediction data.""" + timestamp: float boxes: list[dict] = field(default_factory=list) frame_time_ms: float = 0.0 @@ -64,6 +64,7 @@ class PredictionFrame: @dataclass class GroundTruthComparison: """Comparison result against ground truth.""" + iou_scores: list[float] = field(default_factory=list) false_positives: int = 0 false_negatives: int = 0 @@ -77,12 +78,12 @@ class GroundTruthComparison: # Colors matching the bounding box colors in camera_image.py CLASS_COLORS = { - "Crystal": "#0000ff", # blue - "Loop_face": "#ffff00", # yellow - "Loop_all": "#00ff00", # green - "Pin": "#ff0000", # red - "Ice": "#00ffff", # cyan - "Needle": "#ff00ff", # magenta + "Crystal": "#0000ff", # blue + "Loop_face": "#ffff00", # yellow + "Loop_all": "#00ff00", # green + "Pin": "#ff0000", # red + "Ice": "#00ffff", # cyan + "Needle": "#ff00ff", # magenta } @@ -215,7 +216,9 @@ class ObjectCountWidget(QWidget): # Count label with matching color count_label = QLabel("0") - count_label.setStyleSheet(f"color: {color}; font-size: 14px; font-weight: bold;") + count_label.setStyleSheet( + f"color: {color}; font-size: 14px; font-weight: bold;" + ) count_label.setAlignment(Qt.AlignmentFlag.AlignRight) row = i // 2 @@ -311,7 +314,7 @@ class TimingStatsWidget(QWidget): class ErrorTrackingWidget(QWidget): """ Tracks and displays error metrics. - + Note: FP/FN tracking requires ground truth data to be meaningful. Without ground truth, this shows detection statistics instead. """ @@ -503,7 +506,7 @@ class RollingStatsChart(QWidget): class PredictionMetricsPanel(QWidget): """ Comprehensive panel for real-time ML prediction feedback. - + Connect to PredictionSubscriber.prediction signal: prediction_subscriber.prediction.connect(panel.update_from_prediction) """ @@ -647,7 +650,10 @@ class PredictionMetricsPanel(QWidget): super().showEvent(event) if self._refresh_timer is not None and not self._refresh_timer.isActive(): self._refresh_timer.start() - if self._chart_refresh_timer is not None and not self._chart_refresh_timer.isActive(): + if ( + self._chart_refresh_timer is not None + and not self._chart_refresh_timer.isActive() + ): self._chart_refresh_timer.start() self._dirty = True @@ -786,32 +792,36 @@ class PredictionMetricsPanel(QWidget): try: with open(path, "w", newline="") as f: writer = csv.writer(f) - writer.writerow([ - "timestamp", - "frame_time_ms", - "detection_count", - "mean_confidence", - "max_confidence", - "crystal_count", - "loop_face_count", - "loop_all_count", - "pin_count", - ]) + writer.writerow( + [ + "timestamp", + "frame_time_ms", + "detection_count", + "mean_confidence", + "max_confidence", + "crystal_count", + "loop_face_count", + "loop_all_count", + "pin_count", + ] + ) for frame in self._history: - writer.writerow([ - f"{frame.timestamp:.6f}", - f"{frame.frame_time_ms:.2f}", - len(frame.boxes), - f"{frame.mean_confidence:.4f}", - f"{frame.max_confidence:.4f}", - frame.class_counts.get("Crystal", 0), - frame.class_counts.get("Loop_face", 0), - frame.class_counts.get("Loop_all", 0), - frame.class_counts.get("Pin", 0), - ]) + writer.writerow( + [ + f"{frame.timestamp:.6f}", + f"{frame.frame_time_ms:.2f}", + len(frame.boxes), + f"{frame.mean_confidence:.4f}", + f"{frame.max_confidence:.4f}", + frame.class_counts.get("Crystal", 0), + frame.class_counts.get("Loop_face", 0), + frame.class_counts.get("Loop_all", 0), + frame.class_counts.get("Pin", 0), + ] + ) logger.info(f"Exported {len(self._history)} frames to {path}") except Exception as e: - logger.error(f"Failed to export prediction metrics: {e}") \ No newline at end of file + logger.error(f"Failed to export prediction metrics: {e}") diff --git a/src/aare/gui/panels/raster_data_collection.py b/src/aare/gui/panels/raster_data_collection.py index ce214bdd..1a4199c5 100644 --- a/src/aare/gui/panels/raster_data_collection.py +++ b/src/aare/gui/panels/raster_data_collection.py @@ -1,16 +1,26 @@ -from PySide6.QtCore import Signal, Slot, Qt -from PySide6.QtWidgets import QLabel, QSizePolicy, QSpacerItem, QPushButton, QComboBox, QSlider, QMessageBox +from aarecommon.config.logger import setup_logger +from aarecommon.math.diffraction_geometry import DiffractionGeometry +from aarecommon.models.models import BeamlineStateEnum, DAQStatusModel +from PySide6.QtCore import Qt, Signal, Slot +from PySide6.QtWidgets import ( + QComboBox, + QLabel, + QMessageBox, + QPushButton, + QSizePolicy, + QSlider, + QSpacerItem, +) -from aare.common.diffraction_geometry import DiffractionGeometry -from aare.common.models import DAQStatusModel, BeamlineStateEnum from aare.gui.panels.scan_settings_panel import ScanSettingsPanel from aare.gui.scan_logic.raster_grid_manager import RasterGridManager, RasterGridMetric from aare.gui.widgets.number_line_edit import DbOverrideLineEdit from aare.gui.widgets.raster_grid_table import RasterGridTable -from aare.common.logger_config import setup_logger logger = setup_logger("aareGUI") -#TODO prevent raster if no grid, or at least rpevent smargon from doing danngerous move to 0,0,0!!! + + +# TODO prevent raster if no grid, or at least rpevent smargon from doing danngerous move to 0,0,0!!! class RasterDataCollectionPanel(ScanSettingsPanel): grid_size_updated = Signal(float, float) exp_time_updated = Signal(float) @@ -19,11 +29,16 @@ class RasterDataCollectionPanel(ScanSettingsPanel): grid_metric_updated = Signal(RasterGridMetric) raster_alpha_changed = Signal(int) - def __init__(self, raster_mgr :RasterGridManager, - diffraction: DiffractionGeometry, - default_transmission: float = 1.0, - parent=None): - super().__init__(diffraction, raster_mgr.active_grid.dtz, default_transmission, parent) + def __init__( + self, + raster_mgr: RasterGridManager, + diffraction: DiffractionGeometry, + default_transmission: float = 1.0, + parent=None, + ): + super().__init__( + diffraction, raster_mgr.active_grid.dtz, default_transmission, parent + ) self._previous_sample_was_none_raster = True @@ -31,18 +46,22 @@ class RasterDataCollectionPanel(ScanSettingsPanel): self.__n_y = raster_mgr.active_grid.n_y self.__size_x = raster_mgr.active_grid.grid_size_mm.x * 1000.0 self.__size_y = raster_mgr.active_grid.grid_size_mm.y * 1000.0 - self.__total_time = raster_mgr.active_grid.exp_time_s*self.__n_x*self.__n_y + self.__total_time = raster_mgr.active_grid.exp_time_s * self.__n_x * self.__n_y self._layout.addWidget(QLabel("Grid element size", parent=self), 3, 0) - self.width_enter = DbOverrideLineEdit(5, 100, default=self.__size_x, decimals=0, parent=self) + self.width_enter = DbOverrideLineEdit( + 5, 100, default=self.__size_x, decimals=0, parent=self + ) self.width_enter.valueChanged.connect(self.grid_size) self._register_override_field(self.width_enter) self._layout.addWidget(self.width_enter, 3, 1) self._layout.addWidget(QLabel(" x ", parent=self), 3, 2) - self.height_enter = DbOverrideLineEdit(5, 100, default=self.__size_y, decimals=0, parent=self) + self.height_enter = DbOverrideLineEdit( + 5, 100, default=self.__size_y, decimals=0, parent=self + ) self.height_enter.valueChanged.connect(self.grid_size) self._register_override_field(self.height_enter) @@ -52,7 +71,11 @@ class RasterDataCollectionPanel(ScanSettingsPanel): self._layout.addWidget(QLabel("Image time", parent=self), 4, 0) self.image_time_enter = DbOverrideLineEdit( - 0.0005, 10.0, default=raster_mgr.active_grid.exp_time_s, decimals=4, parent=self + 0.0005, + 10.0, + default=raster_mgr.active_grid.exp_time_s, + decimals=4, + parent=self, ) self._layout.addWidget(self.image_time_enter, 4, 1, 1, 3) self._layout.addWidget(QLabel("s", parent=self), 4, 4) @@ -87,18 +110,26 @@ class RasterDataCollectionPanel(ScanSettingsPanel): self._layout.addWidget(QLabel("Metric", parent=self), 7, 0) self.metric_combo.addItem("Raster score", RasterGridMetric.RASTER_SCORE) - self.metric_combo.addItem("Spot count (low res.)", RasterGridMetric.SPOTS_LOW_RES) + self.metric_combo.addItem( + "Spot count (low res.)", RasterGridMetric.SPOTS_LOW_RES + ) self.metric_combo.addItem("Spot count", RasterGridMetric.SPOTS) - self.metric_combo.addItem("Spot count (indexed)", RasterGridMetric.SPOTS_INDEXED) + self.metric_combo.addItem( + "Spot count (indexed)", RasterGridMetric.SPOTS_INDEXED + ) self.metric_combo.addItem("Spot count (ice)", RasterGridMetric.SPOTS_ICE) - self.metric_combo.addItem("Spot ratio (ice/low res.)", RasterGridMetric.SPOTS_ICE_LOW_RES) + self.metric_combo.addItem( + "Spot ratio (ice/low res.)", RasterGridMetric.SPOTS_ICE_LOW_RES + ) self.metric_combo.addItem("Background estimate", RasterGridMetric.BKG) self.metric_combo.addItem("Indexing result", RasterGridMetric.INDEXING) self.metric_combo.addItem("Profile Radius", RasterGridMetric.PR) self.metric_combo.addItem("B-factor", RasterGridMetric.BFACTOR) self.metric_combo.addItem("Resolution (ML est.)", RasterGridMetric.RES) - self.metric_combo.setCurrentIndex(self.metric_combo.findData(RasterGridMetric.RASTER_SCORE)) + self.metric_combo.setCurrentIndex( + self.metric_combo.findData(RasterGridMetric.RASTER_SCORE) + ) self.metric_combo.currentIndexChanged.connect(self.metric_changed) self._layout.addWidget(self.metric_combo, 7, 1, 1, 3) @@ -107,7 +138,9 @@ class RasterDataCollectionPanel(ScanSettingsPanel): slider = QSlider(orientation=Qt.Orientation.Horizontal, parent=self) slider.setRange(0, 255) slider.setValue(127) - slider.valueChanged.connect(lambda: self.raster_alpha_changed.emit(255 - slider.value())) + slider.valueChanged.connect( + lambda: self.raster_alpha_changed.emit(255 - slider.value()) + ) self._layout.addWidget(slider, 8, 1, 1, 3) self._table = RasterGridTable(raster_mgr) @@ -120,7 +153,9 @@ class RasterDataCollectionPanel(ScanSettingsPanel): self._layout.addWidget(QLabel("Measurement time", parent=self), 11, 0) self.total_time = QLabel(f"{self.__total_time} min 0 s") - self.total_time.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.total_time.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.total_time, 11, 1, 1, 3) self.calculate_total_time() @@ -141,8 +176,7 @@ class RasterDataCollectionPanel(ScanSettingsPanel): def grid_size(self, _: float): self.__size_x = self.width_enter.value self.__size_y = self.height_enter.value - self.grid_size_updated.emit(self.__size_x / 1000.0, - self.__size_y / 1000.0) + self.grid_size_updated.emit(self.__size_x / 1000.0, self.__size_y / 1000.0) self.update_grid_scan_size() @Slot(float) @@ -151,7 +185,9 @@ class RasterDataCollectionPanel(ScanSettingsPanel): self.calculate_total_time() @Slot(int, int, float, float) - def grid_scan_size_change(self, n_x: int, n_y: int, size_x_mm: float, size_y_mm: float): + def grid_scan_size_change( + self, n_x: int, n_y: int, size_x_mm: float, size_y_mm: float + ): self.__size_x = size_x_mm * 1000.0 self.__size_y = size_y_mm * 1000.0 self.__n_x = n_x @@ -193,14 +229,13 @@ class RasterDataCollectionPanel(ScanSettingsPanel): logger.warning(f"Invalid exposure time: {e} reseting to default") self.exp_time_s(0.02) - def metric_changed(self, _: int): self.grid_metric_updated.emit(self.metric_combo.currentData()) def get_parameter_mappings(self): """Return raster-specific parameter mappings.""" return [ - ('exposure', self.image_time_enter, None), + ("exposure", self.image_time_enter, None), # Add other raster-specific parameters here as needed ] @@ -216,25 +251,35 @@ class RasterDataCollectionPanel(ScanSettingsPanel): if self.__n_x <= 0 or self.__n_y <= 0 or self.image_time_enter.value < 0: self.__total_time = 0.0 else: - self.__total_time = self.__n_x * self.__n_y * self.image_time_enter.value * 1.3 - #30% buffer added + self.__total_time = ( + self.__n_x * self.__n_y * self.image_time_enter.value * 1.3 + ) + # 30% buffer added # Show minutes self.update_total_time_label() @Slot() def _on_evaluate_clicked(self): if self.__beamline_state != BeamlineStateEnum.SampleAlignment: - logger.error(f"Beamline state {self.__beamline_state} is not Sample Alignment") - QMessageBox.critical(None, "Error", "Beamline state is not Sample Alignment") + logger.error( + f"Beamline state {self.__beamline_state} is not Sample Alignment" + ) + QMessageBox.critical( + None, "Error", "Beamline state is not Sample Alignment" + ) return - if self.check_before_run(scan_kind = "raster"): + if self.check_before_run(scan_kind="raster"): self.evaluate_grid.emit() @Slot() def _on_evaluate_auto_clicked(self): if self.__beamline_state != BeamlineStateEnum.SampleAlignment: - logger.error(f"Beamline state {self.__beamline_state} is not Sample Alignment") - QMessageBox.critical(None, "Error", "Beamline state is not Sample Alignment") + logger.error( + f"Beamline state {self.__beamline_state} is not Sample Alignment" + ) + QMessageBox.critical( + None, "Error", "Beamline state is not Sample Alignment" + ) return - if self.check_before_run(scan_kind = "raster"): - self.evaluate_grid_auto.emit() \ No newline at end of file + if self.check_before_run(scan_kind="raster"): + self.evaluate_grid_auto.emit() diff --git a/src/aare/gui/panels/reference_tools_panel.py b/src/aare/gui/panels/reference_tools_panel.py index 03cbd1d8..4933ceb7 100644 --- a/src/aare/gui/panels/reference_tools_panel.py +++ b/src/aare/gui/panels/reference_tools_panel.py @@ -1,25 +1,30 @@ # reference_tools_panel.py from typing import Optional -from PySide6.QtCore import Qt, QAbstractTableModel, QModelIndex, Slot, Signal +from aarecommon.config.logger import setup_logger +from aarecommon.models.models import ( + DAQStatusModel, + SampleShortInfo, + SampleShortInfoList, +) +from PySide6.QtCore import QAbstractTableModel, QModelIndex, Qt, Signal, Slot from PySide6.QtGui import QBrush, QColor from PySide6.QtWidgets import ( + QAbstractItemView, QFrame, QGridLayout, - QTableView, QHeaderView, QLabel, - QPushButton, QMenu, - QAbstractItemView, + QPushButton, + QTableView, ) -from aare.common.logger_config import setup_logger -from aare.common.models import SampleShortInfoList, SampleShortInfo, DAQStatusModel from aare.gui.widgets.title_label import TitleLabel logger = setup_logger("aareGUI") + def get_entry(sample: SampleShortInfo, column: int): if column == 0: return sample.loc_str() @@ -35,13 +40,26 @@ def get_entry(sample: SampleShortInfo, column: int): return sample.screening_count return "" + class ReferenceToolsModel(QAbstractTableModel): - def __init__(self, rows: Optional[list[SampleShortInfo]] | None = None, parent=None, current_reference: int | None = None,): + def __init__( + self, + rows: Optional[list[SampleShortInfo]] | None = None, + parent=None, + current_reference: int | None = None, + ): super().__init__(parent) self.samples: list[SampleShortInfo] = rows or [] self.current_reference = current_reference - self.header = ["Position", "Sample name", "Mount count", "Raster count", "Rotation count", "Screening count",] + self.header = [ + "Position", + "Sample name", + "Mount count", + "Raster count", + "Rotation count", + "Screening count", + ] self.__sort_col = 0 self.__sort_order = Qt.SortOrder.AscendingOrder self.__sorted_samples: list[SampleShortInfo] = [] @@ -78,10 +96,13 @@ class ReferenceToolsModel(QAbstractTableModel): if role != Qt.ItemDataRole.DisplayRole: return None if orientation == Qt.Orientation.Horizontal: - return self.header[section] if section < len(self.header) else f"Column {section + 1}" + return ( + self.header[section] + if section < len(self.header) + else f"Column {section + 1}" + ) return str(section + 1) - def update_rows(self, rows: list[SampleShortInfo]): self.beginResetModel() try: @@ -93,7 +114,6 @@ class ReferenceToolsModel(QAbstractTableModel): finally: self.endResetModel() - def sort(self, column, order): self.layoutAboutToBeChanged.emit() self.__sort_order = order @@ -112,22 +132,30 @@ class ReferenceToolsModel(QAbstractTableModel): self.__sorted_samples = sorted( self.samples, key=lambda row: row.loc_str_sort(), - reverse=(self.__sort_order == Qt.SortOrder.DescendingOrder) + reverse=(self.__sort_order == Qt.SortOrder.DescendingOrder), ) elif self.__sort_col == 2: # Numeric sort for Mount count; place None last on ascending, first on descending - none_sentinel = float('inf') if self.__sort_order == Qt.SortOrder.AscendingOrder else float('-inf') + none_sentinel = ( + float("inf") + if self.__sort_order == Qt.SortOrder.AscendingOrder + else float("-inf") + ) self.__sorted_samples = sorted( self.samples, - key=lambda row: (row.mount_count if isinstance(row.mount_count, (int, float)) else none_sentinel), - reverse=(self.__sort_order == Qt.SortOrder.DescendingOrder) + key=lambda row: ( + row.mount_count + if isinstance(row.mount_count, (int, float)) + else none_sentinel + ), + reverse=(self.__sort_order == Qt.SortOrder.DescendingOrder), ) else: # String sort with empty fallback self.__sorted_samples = sorted( self.samples, - key=lambda row: (get_entry(row, self.__sort_col) or ""), - reverse=(self.__sort_order == Qt.SortOrder.DescendingOrder) + key=lambda row: get_entry(row, self.__sort_col) or "", + reverse=(self.__sort_order == Qt.SortOrder.DescendingOrder), ) def get_item(self, row: int) -> Optional[SampleShortInfo]: @@ -144,7 +172,7 @@ class ReferenceToolsModel(QAbstractTableModel): self.dataChanged.emit( self.index(0, 0), self.index(self.rowCount() - 1, self.columnCount() - 1), - [Qt.ItemDataRole.BackgroundRole] + [Qt.ItemDataRole.BackgroundRole], ) @@ -152,7 +180,12 @@ class ReferenceToolsPanel(QFrame): mount = Signal(SampleShortInfo, bool) unmount = Signal() - def __init__(self, samples: SampleShortInfoList | None = None, parent=None, refresh_interval_ms: int = 5000): + def __init__( + self, + samples: SampleShortInfoList | None = None, + parent=None, + refresh_interval_ms: int = 5000, + ): """ :param samples: optional initial SampleShortInfoList to populate the table :param parent: Qt parent @@ -207,9 +240,7 @@ class ReferenceToolsPanel(QFrame): self.table_view.setSelectionBehavior( QAbstractItemView.SelectionBehavior.SelectRows ) - self.table_view.setSelectionMode( - QTableView.SelectionMode.SingleSelection - ) + self.table_view.setSelectionMode(QTableView.SelectionMode.SingleSelection) def _selected_item(self) -> Optional[SampleShortInfo]: idx = self.table_view.currentIndex() @@ -264,4 +295,4 @@ class ReferenceToolsPanel(QFrame): f"Current sample: {sample.sample_name} ({sample.location.segment}{sample.location.pos}-{sample.pin})" ) except Exception as e: - self.curr_sample_label.setText(f"Confusing information :/ {e}") \ No newline at end of file + self.curr_sample_label.setText(f"Confusing information :/ {e}") diff --git a/src/aare/gui/panels/rotation_data_collection.py b/src/aare/gui/panels/rotation_data_collection.py index 7f32e858..83fc9e2e 100644 --- a/src/aare/gui/panels/rotation_data_collection.py +++ b/src/aare/gui/panels/rotation_data_collection.py @@ -1,37 +1,49 @@ from pathlib import Path -from PySide6.QtCore import Slot, Signal, Qt -from PySide6.QtWidgets import QLabel, QComboBox, QPushButton, QMessageBox +from aarecommon.config.logger import setup_logger +from aarecommon.math.diffraction_geometry import DiffractionGeometry +from aarecommon.models.models import BeamlineStateEnum, DAQStatusModel +from aarecommon.models.rotation_scan import RotationScanRequest +from PySide6.QtCore import Qt, Signal, Slot +from PySide6.QtWidgets import QComboBox, QLabel, QMessageBox, QPushButton -from aare.common.diffraction_geometry import DiffractionGeometry -from aare.common.logger_config import setup_logger -from aare.common.models import DAQStatusModel, BeamlineStateEnum -from aare.common.rotation_scan import RotationScanRequest from aare.gui.panels.scan_settings_panel import ScanSettingsPanel -from aare.gui.widgets.number_line_edit import NumberLineEdit, CheckedLineEdit, DbOverrideLineEdit +from aare.gui.widgets.number_line_edit import ( + CheckedLineEdit, + DbOverrideLineEdit, + NumberLineEdit, +) logger = setup_logger("aareGUI") + def add_screening_to_path(path): p = Path(path) return "screening" / p + def add_data_to_path(path): p = Path(path) return "data" / p + class RotationDataCollectionPanel(ScanSettingsPanel): rotation_scan = Signal(RotationScanRequest) viewer_track_online = Signal() - def __init__(self, diffraction: DiffractionGeometry, - default_dtz: float = 200.0, - default_transmission: float = 1.0, - parent=None): - super().__init__(parent=parent, - diffraction=diffraction, - default_dtz=default_dtz, - default_transmission=default_transmission) + def __init__( + self, + diffraction: DiffractionGeometry, + default_dtz: float = 200.0, + default_transmission: float = 1.0, + parent=None, + ): + super().__init__( + parent=parent, + diffraction=diffraction, + default_dtz=default_dtz, + default_transmission=default_transmission, + ) self.__beamline_state = None self.__curr_pgroup = "p11206" @@ -52,10 +64,14 @@ class RotationDataCollectionPanel(ScanSettingsPanel): self.omega_button.clicked.connect(self.update_omega_start) self._layout.addWidget(self.omega_button, 3, 4) - self._layout.addWidget(QLabel("
Screening
", parent=self), 4, 0, 1,6) + self._layout.addWidget( + QLabel("
Screening
", parent=self), 4, 0, 1, 6 + ) self._layout.addWidget(QLabel("Image angle", parent=self), 5, 0) - self.screening_image_angle = NumberLineEdit(0, 90.0, 0.5, decimals=3, parent=self) + self.screening_image_angle = NumberLineEdit( + 0, 90.0, 0.5, decimals=3, parent=self + ) self._layout.addWidget(self.screening_image_angle, 5, 1, 1, 3) self._layout.addWidget(QLabel("°", parent=self), 5, 4) @@ -68,11 +84,21 @@ class RotationDataCollectionPanel(ScanSettingsPanel): self.screening_type = QComboBox(parent=self) self.screening_type.addItem("1 image", {"steps": 1, "omega_step_deg": 0}) - self.screening_type.addItem("2 images every 90°", {"steps": 2, "omega_step_deg": 90}) - self.screening_type.addItem("4 images every 90°", {"steps": 4, "omega_step_deg": 90}) - self.screening_type.addItem("3 images every 60°", {"steps": 3, "omega_step_deg": 60}) - self.screening_type.addItem("2 images every 45°", {"steps": 2, "omega_step_deg": 45}) - self.screening_type.addItem("4 images every 45°", {"steps": 4, "omega_step_deg": 45}) + self.screening_type.addItem( + "2 images every 90°", {"steps": 2, "omega_step_deg": 90} + ) + self.screening_type.addItem( + "4 images every 90°", {"steps": 4, "omega_step_deg": 90} + ) + self.screening_type.addItem( + "3 images every 60°", {"steps": 3, "omega_step_deg": 60} + ) + self.screening_type.addItem( + "2 images every 45°", {"steps": 2, "omega_step_deg": 45} + ) + self.screening_type.addItem( + "4 images every 45°", {"steps": 4, "omega_step_deg": 45} + ) self._layout.addWidget(self.screening_type, 7, 0, 1, 6) @@ -81,20 +107,26 @@ class RotationDataCollectionPanel(ScanSettingsPanel): self.screening_button.clicked.connect(self.run_screening) self._layout.addWidget(self.screening_button, 8, 0, 1, 6) - self._layout.addWidget(QLabel("
Rotation
", parent=self), 9, 0, 1, 6) + self._layout.addWidget( + QLabel("
Rotation
", parent=self), 9, 0, 1, 6 + ) self._layout.addWidget(QLabel("Total angle", parent=self), 10, 0) - self.total_angle = DbOverrideLineEdit(0, 9999.0, default=360.0, decimals=3, parent=self) + self.total_angle = DbOverrideLineEdit( + 0, 9999.0, default=360.0, decimals=3, parent=self + ) self._layout.addWidget(self.total_angle, 10, 1, 1, 3) self._layout.addWidget(QLabel("°", parent=self), 10, 4) self._register_override_field(self.total_angle) self._layout.addWidget(QLabel("Image angle", parent=self), 11, 0) - self.image_angle = DbOverrideLineEdit(0, 10.0, default=0.2, decimals=3, parent=self) + self.image_angle = DbOverrideLineEdit( + 0, 10.0, default=0.2, decimals=3, parent=self + ) self._layout.addWidget(self.image_angle, 11, 1, 1, 3) self._layout.addWidget(QLabel("°", parent=self), 11, 4) self._register_override_field(self.image_angle) - #TODO add protection on X10SA to prevent too short exposure time/ too high detector rep rate + # TODO add protection on X10SA to prevent too short exposure time/ too high detector rep rate self._layout.addWidget(QLabel("Image time", parent=self), 12, 0) self.image_time_enter = DbOverrideLineEdit( 0.0005, 10.0, default=0.01, decimals=4, parent=self @@ -105,7 +137,9 @@ class RotationDataCollectionPanel(ScanSettingsPanel): self._layout.addWidget(QLabel("Total measurement time", parent=self), 13, 0) self.total_time = QLabel(f"{self.__total_time} min 0 s") - self.total_time.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.total_time.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.total_time, 13, 1, 1, 3) self.total_angle.valueChanged.connect(self.calculate_measurement_time) @@ -116,7 +150,9 @@ class RotationDataCollectionPanel(ScanSettingsPanel): self._layout.addWidget(QLabel("Dose", parent=self), 14, 0) self.dose = QLabel(f"{self.__dose_mgy}") - self.dose.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.dose.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.dose, 14, 1, 1, 3) self._layout.addWidget(QLabel("MGy", parent=self), 14, 4) @@ -132,10 +168,14 @@ class RotationDataCollectionPanel(ScanSettingsPanel): @Slot() def run_screening(self): if self.__beamline_state != BeamlineStateEnum.SampleAlignment: - logger.error(f"Beamline state {self.__beamline_state} is not Sample Alignment") - QMessageBox.critical(None, "Error", "Beamline state is not Sample Alignment") + logger.error( + f"Beamline state {self.__beamline_state} is not Sample Alignment" + ) + QMessageBox.critical( + None, "Error", "Beamline state is not Sample Alignment" + ) return - if not self.check_before_run(scan_kind = "screening"): + if not self.check_before_run(scan_kind="screening"): logger.error("Cannot run measurement because of check") return screening_settings = self.screening_type.currentData() @@ -176,10 +216,14 @@ class RotationDataCollectionPanel(ScanSettingsPanel): @Slot() def run_measurement(self): if self.__beamline_state != BeamlineStateEnum.SampleAlignment: - logger.error(f"Beamline state {self.__beamline_state} is not Sample Alignment") - QMessageBox.critical(None, "Error", "Beamline state is not Sample Alignment") + logger.error( + f"Beamline state {self.__beamline_state} is not Sample Alignment" + ) + QMessageBox.critical( + None, "Error", "Beamline state is not Sample Alignment" + ) return - if not self.check_before_run(scan_kind = "rotation"): + if not self.check_before_run(scan_kind="rotation"): logger.error("Cannot run measurement because of check") return r = RotationScanRequest( @@ -190,7 +234,8 @@ class RotationDataCollectionPanel(ScanSettingsPanel): dtz=self._dtz, transmission=self._transmission, screening=False, - exp_time_s=self.image_time_enter.value) + exp_time_s=self.image_time_enter.value, + ) self.rotation_scan.emit(r) self.viewer_track_online.emit() @@ -214,16 +259,20 @@ class RotationDataCollectionPanel(ScanSettingsPanel): secs = 0 self.total_time.setText(f"{mins} min {secs} s") - def calculate_measurement_time(self): - if self.image_angle.value <= 0 or self.total_angle.value <= 0 or self.image_time_enter.value < 0: + if ( + self.image_angle.value <= 0 + or self.total_angle.value <= 0 + or self.image_time_enter.value < 0 + ): self.__total_time = 0.0 self.update_total_time_label() return - self.__total_time = (self.total_angle.value / self.image_angle.value) * self.image_time_enter.value + self.__total_time = ( + self.total_angle.value / self.image_angle.value + ) * self.image_time_enter.value self.update_total_time_label() - @Slot(str) def update_filename(self, filename: str): self._filename = filename @@ -234,8 +283,14 @@ class RotationDataCollectionPanel(ScanSettingsPanel): self.__omega = s.geom.omega_deg can_edit = getattr(self, "_can_edit_params", False) - for w in (self.start_angle, self.screening_image_angle, self.screening_image_time_enter, - self.total_angle, self.image_angle, self.image_time_enter): + for w in ( + self.start_angle, + self.screening_image_angle, + self.screening_image_time_enter, + self.total_angle, + self.image_angle, + self.image_time_enter, + ): if hasattr(w, "set_busy"): w.set_busy(not can_edit) else: @@ -245,8 +300,12 @@ class RotationDataCollectionPanel(ScanSettingsPanel): # The DB-sourced fields (total_angle, image_angle, image_time_enter) # are reset by the base panel (_reset_to_defaults); only reset the # screening/start widgets that are not part of the source toggle. - for w in (self.start_angle, self.screening_image_angle, self.screening_image_time_enter): - w.reset_to_default() + for w in ( + self.start_angle, + self.screening_image_angle, + self.screening_image_time_enter, + ): + w.reset_to_default() self._previous_sample_was_none_rotation = True elif s.sample is not None: self._previous_sample_was_none_rotation = False @@ -259,7 +318,9 @@ class RotationDataCollectionPanel(ScanSettingsPanel): beam_area = s.geom.beam_size_mm.x * s.geom.beam_size_mm.y * 1e6 time = self.image_number() * self.image_time_enter.value - self.__dose_mgy = (time * s.bl.flux_ph_s * self._transmission) / (beam_area * kdose) + self.__dose_mgy = (time * s.bl.flux_ph_s * self._transmission) / ( + beam_area * kdose + ) self.dose.setText(f"{(self.__dose_mgy / 1e6):.1f}") self.__beamline_state = s.state @@ -271,9 +332,7 @@ class RotationDataCollectionPanel(ScanSettingsPanel): def get_parameter_mappings(self): """Return rotation-specific parameter mappings.""" return [ - ('totalrange', self.total_angle, lambda v: float(v)), - ('oscillation', self.image_angle, None), - ('exposure', self.image_time_enter, None), + ("totalrange", self.total_angle, lambda v: float(v)), + ("oscillation", self.image_angle, None), + ("exposure", self.image_time_enter, None), ] - - diff --git a/src/aare/gui/panels/samcam_panel.py b/src/aare/gui/panels/samcam_panel.py index 1cc30dfa..51c78645 100644 --- a/src/aare/gui/panels/samcam_panel.py +++ b/src/aare/gui/panels/samcam_panel.py @@ -1,8 +1,17 @@ -from PySide6.QtWidgets import QWidget, QVBoxLayout, QHBoxLayout, QLabel, QDoubleSpinBox, QCheckBox, QLineEdit, \ - QPushButton, QComboBox +from aarecommon.models.models import DAQStatusModel, SampleCameraSettings from PySide6.QtCore import Signal, Slot +from PySide6.QtWidgets import ( + QCheckBox, + QComboBox, + QDoubleSpinBox, + QHBoxLayout, + QLabel, + QLineEdit, + QPushButton, + QVBoxLayout, + QWidget, +) -from aare.common.models import SampleCameraSettings, DAQStatusModel from aare.gui.widgets.title_label import TitleLabel @@ -34,7 +43,9 @@ class SamcamPanel(QWidget): self.exposure_spinbox = QDoubleSpinBox() self.exposure_spinbox.setRange(0, 1.0) # Adjust range as needed self.exposure_spinbox.setSingleStep(0.001) - self.exposure_spinbox.setStyleSheet("QDoubleSpinBox { background-color: white; }") + self.exposure_spinbox.setStyleSheet( + "QDoubleSpinBox { background-color: white; }" + ) self.exposure_spinbox.setDecimals(3) self.exposure_spinbox.valueChanged.connect(self.__changed) @@ -56,13 +67,17 @@ class SamcamPanel(QWidget): # Persist the current gain/exposure as the beam-location preset for the # current zoom (only meaningful in beam-location mode). self.save_beam_location_button = QPushButton("Save samcam settings (bl)") - self.save_beam_location_button.clicked.connect(self.save_beam_location_setting.emit) + self.save_beam_location_button.clicked.connect( + self.save_beam_location_setting.emit + ) screenshot_filename_layout = QHBoxLayout() screenshot_filename_label = QLabel("Filename:") self.screenshot_filename_edit = QLineEdit() self.screenshot_filename_edit.setPlaceholderText("optional") - self.screenshot_filename_edit.setStyleSheet("QLineEdit { background-color: white; }") + self.screenshot_filename_edit.setStyleSheet( + "QLineEdit { background-color: white; }" + ) screenshot_filename_layout.addWidget(screenshot_filename_label) screenshot_filename_layout.addWidget(self.screenshot_filename_edit) @@ -70,7 +85,9 @@ class SamcamPanel(QWidget): screenshot_message_label = QLabel("Message:") self.screenshot_message_edit = QLineEdit() self.screenshot_message_edit.setPlaceholderText("optional") - self.screenshot_message_edit.setStyleSheet("QLineEdit { background-color: white; }") + self.screenshot_message_edit.setStyleSheet( + "QLineEdit { background-color: white; }" + ) screenshot_message_layout.addWidget(screenshot_message_label) screenshot_message_layout.addWidget(self.screenshot_message_edit) @@ -84,39 +101,49 @@ class SamcamPanel(QWidget): self.show_detections_checkbox.toggled.connect(self.show_detections_changed.emit) detections_layout.addWidget(self.show_detections_checkbox) - #Show detection polygons checkbox + # Show detection polygons checkbox detection_polygons_layout = QHBoxLayout() self.show_detection_polygons_checkbox = QCheckBox("Show ML polygons") self.show_detection_polygons_checkbox.setChecked(True) - self.show_detection_polygons_checkbox.toggled.connect(self.show_detection_polygons_changed.emit) + self.show_detection_polygons_checkbox.toggled.connect( + self.show_detection_polygons_changed.emit + ) detection_polygons_layout.addWidget(self.show_detection_polygons_checkbox) # Show target point checkbox target_point_layout = QHBoxLayout() self.show_target_point_checkbox = QCheckBox("Show target point") self.show_target_point_checkbox.setChecked(True) - self.show_target_point_checkbox.toggled.connect(self.show_target_point_changed.emit) + self.show_target_point_checkbox.toggled.connect( + self.show_target_point_changed.emit + ) target_point_layout.addWidget(self.show_target_point_checkbox) # Show target coordinates checkbox target_coords_layout = QHBoxLayout() self.show_target_coordinates_checkbox = QCheckBox("Show target coordinates") self.show_target_coordinates_checkbox.setChecked(True) - self.show_target_coordinates_checkbox.toggled.connect(self.show_target_coordinates_changed.emit) + self.show_target_coordinates_checkbox.toggled.connect( + self.show_target_coordinates_changed.emit + ) target_coords_layout.addWidget(self.show_target_coordinates_checkbox) # Show legend checkbox legend_layout = QHBoxLayout() self.show_overlay_legend_checkbox = QCheckBox("Show overlay legend") self.show_overlay_legend_checkbox.setChecked(True) - self.show_overlay_legend_checkbox.toggled.connect(self.show_overlay_legend_changed.emit) + self.show_overlay_legend_checkbox.toggled.connect( + self.show_overlay_legend_changed.emit + ) legend_layout.addWidget(self.show_overlay_legend_checkbox) # Compact legend checkbox compact_legend_layout = QHBoxLayout() self.compact_overlay_legend_checkbox = QCheckBox("Compact legend") self.compact_overlay_legend_checkbox.setChecked(False) - self.compact_overlay_legend_checkbox.toggled.connect(self.compact_overlay_legend_changed.emit) + self.compact_overlay_legend_checkbox.toggled.connect( + self.compact_overlay_legend_changed.emit + ) compact_legend_layout.addWidget(self.compact_overlay_legend_checkbox) # Target color @@ -125,7 +152,9 @@ class SamcamPanel(QWidget): self.target_color_combo = QComboBox() self.target_color_combo.addItems(["Cyan", "Dark Blue", "Dark Red"]) self.target_color_combo.setCurrentText("Cyan") - self.target_color_combo.currentTextChanged.connect(self.target_color_changed.emit) + self.target_color_combo.currentTextChanged.connect( + self.target_color_changed.emit + ) target_color_layout.addWidget(target_color_label) target_color_layout.addWidget(self.target_color_combo) @@ -146,8 +175,11 @@ class SamcamPanel(QWidget): self.setLayout(layout) def __changed(self): - self.changed.emit(SampleCameraSettings(gain=self.gain_spinbox.value(), - exposure=self.exposure_spinbox.value())) + self.changed.emit( + SampleCameraSettings( + gain=self.gain_spinbox.value(), exposure=self.exposure_spinbox.value() + ) + ) def __request_screenshot(self): self.screenshot_requested.emit( @@ -156,15 +188,15 @@ class SamcamPanel(QWidget): ) def apply_overlay_settings( - self, - *, - show_detections: bool, - show_detection_polygons: bool, - show_target_point: bool, - show_target_coordinates: bool, - show_overlay_legend: bool, - compact_overlay_legend: bool, - target_color: str, + self, + *, + show_detections: bool, + show_detection_polygons: bool, + show_target_point: bool, + show_target_coordinates: bool, + show_overlay_legend: bool, + compact_overlay_legend: bool, + target_color: str, ) -> None: for widget in ( self.show_detections_checkbox, diff --git a/src/aare/gui/panels/sample_queue_panel.py b/src/aare/gui/panels/sample_queue_panel.py index 629c4cc5..ed0087cb 100644 --- a/src/aare/gui/panels/sample_queue_panel.py +++ b/src/aare/gui/panels/sample_queue_panel.py @@ -1,18 +1,32 @@ -from PySide6.QtCore import Signal, Slot, Qt, QTimer -from PySide6.QtWidgets import QFrame, QVBoxLayout, QHBoxLayout, QPushButton, QHeaderView, QSizePolicy, \ - QTableView, QMessageBox, QCheckBox +from aarecommon.config.logger import setup_logger +from aarecommon.models.models import ( + BeamlineStateEnum, + DAQStatusModel, + SampleShortInfo, + SampleShortInfoList, + SessionsStateEnum, +) +from PySide6.QtCore import Qt, QTimer, Signal, Slot from PySide6.QtGui import QKeySequence, QShortcut -from aare.common.models import SampleShortInfoList, SampleShortInfo, BeamlineStateEnum, SessionsStateEnum +from PySide6.QtWidgets import ( + QCheckBox, + QFrame, + QHBoxLayout, + QHeaderView, + QMessageBox, + QPushButton, + QSizePolicy, + QTableView, + QVBoxLayout, +) from aare.gui.models.sample_queue_model import SampleQueueSpreadsheet from aare.gui.widgets.message_box import LOW_CURRENT_THRESHOLD, conditions_auto_check from aare.gui.widgets.title_label import TitleLabel -from aare.common.models import DAQStatusModel - -from aare.common.logger_config import setup_logger logger = setup_logger("aareGUI") + class SampleQueuePanel(QFrame): auto_scan = Signal(SampleShortInfo) unmount = Signal() @@ -22,7 +36,12 @@ class SampleQueuePanel(QFrame): automation_running_changed = Signal(bool) step_through_changed = Signal(bool) - def __init__(self, parent=None, samples: SampleShortInfoList | None = None, show_user: bool = False): + def __init__( + self, + parent=None, + samples: SampleShortInfoList | None = None, + show_user: bool = False, + ): super().__init__(parent) self.__baton_holder = None self.__busy = None @@ -57,17 +76,25 @@ class SampleQueuePanel(QFrame): header = self.table_view.horizontalHeader() header.setSectionResizeMode(QHeaderView.ResizeMode.Stretch) - self.table_view.setSizePolicy(QSizePolicy.Policy.Expanding, QSizePolicy.Policy.Expanding) + self.table_view.setSizePolicy( + QSizePolicy.Policy.Expanding, QSizePolicy.Policy.Expanding + ) self.table_view.setAcceptDrops(True) self.table_view.setDropIndicatorShown(True) self.table_view.setSelectionBehavior(QTableView.SelectionBehavior.SelectRows) - self.delete_shortcut = QShortcut(QKeySequence.StandardKey.Delete, self.table_view) + self.delete_shortcut = QShortcut( + QKeySequence.StandardKey.Delete, self.table_view + ) self.delete_shortcut.setContext(Qt.ShortcutContext.WidgetWithChildrenShortcut) self.delete_shortcut.activated.connect(self.remove_selected_samples) - self.run_pause_shortcut = QShortcut(QKeySequence(Qt.Key.Key_Space), self.table_view) - self.run_pause_shortcut.setContext(Qt.ShortcutContext.WidgetWithChildrenShortcut) + self.run_pause_shortcut = QShortcut( + QKeySequence(Qt.Key.Key_Space), self.table_view + ) + self.run_pause_shortcut.setContext( + Qt.ShortcutContext.WidgetWithChildrenShortcut + ) self.run_pause_shortcut.activated.connect(self.run) layout.addWidget(self.table_view) @@ -82,10 +109,14 @@ class SampleQueuePanel(QFrame): self.clear_button = QPushButton("✖ Clear list", self) self.clear_button.clicked.connect(self.clear) - self.park_and_dry_when_cleared = QCheckBox("Park and dry when automation finishes", self) + self.park_and_dry_when_cleared = QCheckBox( + "Park and dry when automation finishes", self + ) self.park_and_dry_when_cleared.setChecked(True) - self.pause_on_conditions_cb = QCheckBox("Pause on bad conditions (beam/shutter/door/robot)", self) + self.pause_on_conditions_cb = QCheckBox( + "Pause on bad conditions (beam/shutter/door/robot)", self + ) self.pause_on_conditions_cb.setChecked(True) self.pause_on_conditions_cb.setToolTip( "When on, automation refuses to start and pauses between samples if the beam, " @@ -120,14 +151,18 @@ class SampleQueuePanel(QFrame): def is_step_through(self) -> bool: return self._step_through - def queue_preview(self) -> tuple[SampleShortInfo | None, SampleShortInfo | None, SampleShortInfo | None]: + def queue_preview( + self, + ) -> tuple[SampleShortInfo | None, SampleShortInfo | None, SampleShortInfo | None]: samples = self.table_model.samples current = samples[0] if len(samples) > 0 else None next_sample = samples[1] if len(samples) > 1 else None next_next_sample = samples[2] if len(samples) > 2 else None return current, next_sample, next_next_sample - def queue_samples(self, samples: list[SampleShortInfo], replace: bool = True) -> None: + def queue_samples( + self, samples: list[SampleShortInfo], replace: bool = True + ) -> None: if replace: self.table_model.updateData(list(samples)) else: @@ -194,7 +229,11 @@ class SampleQueuePanel(QFrame): """Beamline conditions that currently block/should pause automation.""" problems: list[str] = [] if self.ring_current is None or self.ring_current < LOW_CURRENT_THRESHOLD: - rc = "unknown" if self.ring_current is None else f"{round(self.ring_current, 2)} mA" + rc = ( + "unknown" + if self.ring_current is None + else f"{round(self.ring_current, 2)} mA" + ) problems.append(f"beam (ring current {rc})") if not self._experiment_shutter_state: problems.append("experiment shutter closed") @@ -224,13 +263,17 @@ class SampleQueuePanel(QFrame): self.__warning_msg_box = QMessageBox(self) self.__warning_msg_box.setIcon(QMessageBox.Icon.Warning) self.__warning_msg_box.setWindowTitle("TELL Warning") - self.__warning_msg_box.setText("TELL reported a warning. Please see console for details.") + self.__warning_msg_box.setText( + "TELL reported a warning. Please see console for details." + ) self.__warning_msg_box.setInformativeText( "Automation paused for 10 minutes to allow sufficient dry time." "\nClick 'Continue Now' to resume immediately, or wait for auto-resume." ) - continue_btn = self.__warning_msg_box.addButton("Continue Now", QMessageBox.ButtonRole.AcceptRole) + continue_btn = self.__warning_msg_box.addButton( + "Continue Now", QMessageBox.ButtonRole.AcceptRole + ) self.__warning_msg_box.setWindowModality(Qt.WindowModality.NonModal) continue_btn.clicked.connect(self.manual_resume_from_warning) self.__warning_msg_box.show() @@ -297,12 +340,14 @@ class SampleQueuePanel(QFrame): return if self.__busy: - logger.error(f"Cannot run automation while beamline is busy. Busy flag = {self.__busy}") + logger.error( + f"Cannot run automation while beamline is busy. Busy flag = {self.__busy}" + ) self.show_error_dialog( title="Beamline is busy", msg="Cannot run automation while beamline is busy", info="Please wait until beamline is idle " - "or contact your local contact for support", + "or contact your local contact for support", ) return @@ -311,7 +356,7 @@ class SampleQueuePanel(QFrame): title="Session is vacant", msg="Starting automation while session is vacant is not currently implemented", info="Please grab the baton before continuing " - "or contact your local contact for support", + "or contact your local contact for support", ) return @@ -320,8 +365,8 @@ class SampleQueuePanel(QFrame): title="You do not hold the baton", msg="You do not hold the baton.", info="Please request the baton if it is your shift." - "If your baton request is denied and it should be the start of your shift," - "please contact your local contact for support", + "If your baton request is denied and it should be the start of your shift," + "please contact your local contact for support", ) return @@ -329,8 +374,10 @@ class SampleQueuePanel(QFrame): self.show_error_dialog( title="Maintenance mode", msg="Cannot run automation while beamline is in maintenance mode", - info=("Change to safe state such as Sample Exchange before trying to continue. " - "If this issue persists please contact your local contact for support."), + info=( + "Change to safe state such as Sample Exchange before trying to continue. " + "If this issue persists please contact your local contact for support." + ), ) return @@ -338,7 +385,9 @@ class SampleQueuePanel(QFrame): if checks_enabled: bad = self._bad_conditions() if bad: - logger.warning(f"Cannot start automation; beamline not ready: {bad}") + logger.warning( + f"Cannot start automation; beamline not ready: {bad}" + ) self.show_error_dialog( title="Beamline not ready", msg="Cannot start automation:\n- " + "\n- ".join(bad), @@ -401,7 +450,9 @@ class SampleQueuePanel(QFrame): msg = "Automation paused — beamline not ready:\n- " + "\n- ".join(bad) if conditions_auto_check(self, msg, self._conditions_ok): - logger.debug("Conditions recovered or user chose to continue; resuming automation") + logger.debug( + "Conditions recovered or user chose to continue; resuming automation" + ) self.table_model.set_running(True) self.__pause = False self.play_button.setText("⏸ Pause") @@ -441,16 +492,23 @@ class SampleQueuePanel(QFrame): self._current_db_id = None self.pause_automation() logger.critical("TELL Critical Error: Stopping automation.") - self.show_error_dialog(title=reply, msg="Critical Error", info="Repeated mount failuers.\n" - "Please check to see if pins" - " are loaded in these postions." - "\n If not please continue with new " - "samples.\nOtherwise contact your" - " local contact for support") + self.show_error_dialog( + title=reply, + msg="Critical Error", + info="Repeated mount failuers.\n" + "Please check to see if pins" + " are loaded in these postions." + "\n If not please continue with new " + "samples.\nOtherwise contact your" + " local contact for support", + ) elif reply == "Authentication Error": self.pause_automation() - self.show_error_dialog(title=reply, msg="Authentication Error", - info="Please take the session to continue") + self.show_error_dialog( + title=reply, + msg="Authentication Error", + info="Please take the session to continue", + ) else: self.table_model.remove_sample(db_id) self._current_db_id = None @@ -465,4 +523,4 @@ class SampleQueuePanel(QFrame): else: self._finish_empty_queue() - self._emit_samples_in_queue_changed() \ No newline at end of file + self._emit_samples_in_queue_changed() diff --git a/src/aare/gui/panels/scan_settings_panel.py b/src/aare/gui/panels/scan_settings_panel.py index 0c6a1ceb..34c4ad08 100644 --- a/src/aare/gui/panels/scan_settings_panel.py +++ b/src/aare/gui/panels/scan_settings_panel.py @@ -1,34 +1,38 @@ -from PySide6.QtCore import Slot, Signal +from aarecommon.config.logger import setup_logger +from aarecommon.math.diffraction_geometry import DiffractionGeometry +from aarecommon.models.models import DAQStatusModel, SessionsStateEnum +from PySide6.QtCore import Signal, Slot from PySide6.QtWidgets import ( - QWidget, - QVBoxLayout, - QHBoxLayout, + QButtonGroup, QGridLayout, + QHBoxLayout, QLabel, QPushButton, QRadioButton, - QButtonGroup, + QVBoxLayout, + QWidget, ) -from aare.common.diffraction_geometry import DiffractionGeometry -from aare.common.logger_config import setup_logger -from aare.common.models import DAQStatusModel, SessionsStateEnum from aare.gui.widgets.message_box import precondition_check from aare.gui.widgets.number_line_edit import DbOverrideLineEdit logger = setup_logger("aareGUI") + class ScanSettingsPanel(QWidget): dtz_updated = Signal(float) dtz_move = Signal(float) - #TODO min and max dtz is set by beamline add max - MIN_DTZ = 108.0 # this is beamline dependent + # TODO min and max dtz is set by beamline add max + MIN_DTZ = 108.0 # this is beamline dependent transmission_updated = Signal(float) - def __init__(self, diffraction: DiffractionGeometry, - default_dtz: float = 200.0, - default_transmission: float = 1.0, - parent=None): + def __init__( + self, + diffraction: DiffractionGeometry, + default_dtz: float = 200.0, + default_transmission: float = 1.0, + parent=None, + ): super().__init__(parent) self.__diffraction = diffraction @@ -99,9 +103,13 @@ class ScanSettingsPanel(QWidget): self._register_override_field(self.transmission_enter) self.reload_params_button = QPushButton("Reload DB params") - self.reload_params_button.setToolTip("Reload data collection parameters from database") + self.reload_params_button.setToolTip( + "Reload data collection parameters from database" + ) self.reload_params_button.clicked.connect(self.reload_parameters) - self.reload_params_button.setVisible(False) # Child classes should make it visible + self.reload_params_button.setVisible( + False + ) # Child classes should make it visible # -- source toggle ----------------------------------------------------- def _build_source_toggle(self) -> QWidget: @@ -159,22 +167,29 @@ class ScanSettingsPanel(QWidget): @Slot(DAQStatusModel) def update_daq_status(self, s: DAQStatusModel): self.dtz_enter.update_limits(s.bl.dtz_min, s.bl.dtz_max) - self.high_res_enter.update_limits(self.__diffraction.resolution_angstrom(s.bl.dtz_min), - self.__diffraction.resolution_angstrom(s.bl.dtz_max)) + self.high_res_enter.update_limits( + self.__diffraction.resolution_angstrom(s.bl.dtz_min), + self.__diffraction.resolution_angstrom(s.bl.dtz_max), + ) self.__diffraction = s.diffraction self._ring_current = s.bl.ring_current_mA self._experiment_shutter_state = s.bl.exp_shutter_open self._door_prohibited = getattr(s.bl, "pss_prohibited", None) - can_edit = (not s.busy) and (s.session.session in (SessionsStateEnum.OwnedByYou, SessionsStateEnum.PendingElseToYou)) + can_edit = (not s.busy) and ( + s.session.session + in (SessionsStateEnum.OwnedByYou, SessionsStateEnum.PendingElseToYou) + ) self._can_edit_params = can_edit # Lock/unlock the override fields for w in (self.dtz_enter, self.high_res_enter, self.transmission_enter): - w.set_busy(not can_edit) + w.set_busy(not can_edit) # Update sample and parameters if s.sample is not None: self._sample = s.sample - self._params = s.sample.aaredb_params if hasattr(s.sample, 'aaredb_params') else None + self._params = ( + s.sample.aaredb_params if hasattr(s.sample, "aaredb_params") else None + ) # Enable reload button only if sample has parameters self._previous_sample_was_none = False self.reload_params_button.setEnabled(self._params is not None) @@ -288,21 +303,23 @@ class ScanSettingsPanel(QWidget): widget.update_value(converted) # Handle transmission with conversion (common to all panels) - if (transmission := getattr(self._params, 'transmission', None)) is not None: - transmission_value = transmission / 100.0 if transmission > 1.0 else transmission + if (transmission := getattr(self._params, "transmission", None)) is not None: + transmission_value = ( + transmission / 100.0 if transmission > 1.0 else transmission + ) self._transmission = transmission_value self.transmission_enter.set_db_value(transmission_value) # Handle target resolution (common to all panels). dtz is derived from # the resolution so its database value is kept consistent here. - if (target_res := getattr(self._params, 'targetresolution', None)) is not None: + if (target_res := getattr(self._params, "targetresolution", None)) is not None: self._apply_db_resolution(float(target_res)) # Store metadata (common to all panels) - self._sample_space_group = getattr(self._params, 'spacegroupnumber', None) - self._sample_cell_parameters = getattr(self._params, 'cellparameters', None) - self._sample_pdb_id = getattr(self._params, 'pdbid', None) - self._target_dose = getattr(self._params, 'dose', None) + self._sample_space_group = getattr(self._params, "spacegroupnumber", None) + self._sample_cell_parameters = getattr(self._params, "cellparameters", None) + self._sample_pdb_id = getattr(self._params, "pdbid", None) + self._target_dose = getattr(self._params, "dose", None) def _apply_db_resolution(self, target_res: float): """Set the database resolution and the matching database dtz so the @@ -321,7 +338,7 @@ class ScanSettingsPanel(QWidget): """ return [] - def check_before_run (self, scan_kind: str): + def check_before_run(self, scan_kind: str): if not precondition_check( self, ring_current=self._ring_current, diff --git a/src/aare/gui/panels/smargon_panel.py b/src/aare/gui/panels/smargon_panel.py index d9b4849f..9bd66726 100644 --- a/src/aare/gui/panels/smargon_panel.py +++ b/src/aare/gui/panels/smargon_panel.py @@ -1,12 +1,12 @@ +from aarecommon.math.coordinate import Coordinate, SmargonCoordinate +from aarecommon.math.sample_geometry import SampleGeometryModel +from aarecommon.models.models import DAQStatusModel from PySide6.QtCore import Signal, Slot -from PySide6.QtWidgets import QWidget, QGridLayout, QLabel, QPushButton +from PySide6.QtWidgets import QGridLayout, QLabel, QPushButton, QWidget -from aare.common.models import DAQStatusModel from aare.gui.widgets.button_with_payload import ButtonWithPayload from aare.gui.widgets.number_line_edit import NumberLineEdit from aare.gui.widgets.title_label import TitleLabel -from aare.common.coordinate import SmargonCoordinate, Coordinate -from aare.common.sample_geometry import SampleGeometryModel class SmargonMoveWidget(QWidget): @@ -43,11 +43,11 @@ class SmargonMoveWidget(QWidget): grid_layout.addWidget(self.button_in, 3, 2) grid_layout.addWidget(self.button_out, 3, 0) - - @Slot(dict) def smargon_button(self, payload: dict): - self.smargon_rel.emit(Coordinate(x=payload["x"], y=payload["y"],z=payload["z"])) + self.smargon_rel.emit( + Coordinate(x=payload["x"], y=payload["y"], z=payload["z"]) + ) class SmargonPanel(QWidget): @@ -83,7 +83,6 @@ class SmargonPanel(QWidget): grid_layout.addWidget(self.move_panel, 3, 0, 1, 6) self.move_panel.smargon_rel.connect(self.smargon_rel) - grid_layout.addWidget(QLabel("Step", parent=self), 4, 0) self.step = NumberLineEdit(1, 1000, 100, 0, parent=self) grid_layout.addWidget(self.step, 4, 1) @@ -96,8 +95,10 @@ class SmargonPanel(QWidget): @Slot() def home(self): - #TODO move SMARGON_HOME to REDIS, allow GUI to read this value - self.smargon.emit(SmargonCoordinate(sh_mm=Coordinate(x=0, y=0, z=18), phi_deg=0, chi_deg=0)) + # TODO move SMARGON_HOME to REDIS, allow GUI to read this value + self.smargon.emit( + SmargonCoordinate(sh_mm=Coordinate(x=0, y=0, z=18), phi_deg=0, chi_deg=0) + ) @Slot(float) def phi(self, f: float): diff --git a/src/aare/gui/panels/smart_rotation_panel.py b/src/aare/gui/panels/smart_rotation_panel.py index bf88477e..156e6d04 100644 --- a/src/aare/gui/panels/smart_rotation_panel.py +++ b/src/aare/gui/panels/smart_rotation_panel.py @@ -1,16 +1,24 @@ import math -from PySide6.QtCore import Slot, Qt, Signal -from PySide6.QtWidgets import QWidget, QGridLayout, QLabel, QPushButton, QSpacerItem, QSizePolicy +from aarecommon.config.logger import setup_logger +from aarecommon.models.models import CrystalSize, DAQStatusModel, SimpleScanParameters +from aarecommon.models.rotation_scan import RotationScanRequest +from PySide6.QtCore import Qt, Signal, Slot +from PySide6.QtWidgets import ( + QGridLayout, + QLabel, + QPushButton, + QSizePolicy, + QSpacerItem, + QWidget, +) -from aare.common.logger_config import setup_logger -from aare.common.models import DAQStatusModel, SimpleScanParameters, CrystalSize -from aare.common.rotation_scan import RotationScanRequest from aare.gui.panels.rotation_data_collection import add_data_to_path from aare.gui.widgets.number_line_edit import NumberLineEdit logger = setup_logger("aareGUI") + class SimpleRotationSettingsPanel(QWidget): rotation_scan = Signal(RotationScanRequest) viewer_track_online = Signal() @@ -22,7 +30,7 @@ class SimpleRotationSettingsPanel(QWidget): self.n_images = 1 self.xtal_size_dose_rate_MGy_s = None - self.xtal_size = CrystalSize(x=0,y=0,z=0) + self.xtal_size = CrystalSize(x=0, y=0, z=0) self.xtal_x = None self.xtal_y = None self.xtal_z = None @@ -52,7 +60,9 @@ class SimpleRotationSettingsPanel(QWidget): self.visible_res_enter.newValue.connect(self.set_visible_resolution) self._layout.addWidget(QLabel("Start angle", parent=self), 1, 0) - self.start_angle_enter = NumberLineEdit(-720, 720.0, 0.0, decimals=3, parent=self) + self.start_angle_enter = NumberLineEdit( + -720, 720.0, 0.0, decimals=3, parent=self + ) self._layout.addWidget(self.start_angle_enter, 1, 1, 1, 3) self._layout.addWidget(QLabel("°", parent=self), 1, 4) @@ -79,7 +89,9 @@ class SimpleRotationSettingsPanel(QWidget): self.image_angle_enter.newValue.connect(self.set_image_angle) self._layout.addWidget(QLabel("Temperature", parent=self), 4, 0) - self.temp_enter = NumberLineEdit(80, 330, decimals=2, default=100.0, parent=self) + self.temp_enter = NumberLineEdit( + 80, 330, decimals=2, default=100.0, parent=self + ) self._layout.addWidget(self.temp_enter, 4, 1, 1, 3) self._layout.addWidget(QLabel("K", parent=self), 4, 4) self.temp_enter.newValue.connect(self.set_temperature) @@ -87,101 +99,143 @@ class SimpleRotationSettingsPanel(QWidget): # Calculated labels self._layout.addWidget(QLabel("Target resolution", parent=self), 5, 0) self.target_res_label = QLabel("--", parent=self) - self.target_res_label.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.target_res_label.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.target_res_label, 5, 1, 1, 3) self._layout.addWidget(QLabel("Å", parent=self), 5, 4) self._layout.addWidget(QLabel("Image time", parent=self), 6, 0) self.image_time_label = QLabel("--", parent=self) - self.image_time_label.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.image_time_label.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.image_time_label, 6, 1, 1, 3) self._layout.addWidget(QLabel("s", parent=self), 6, 4) self._layout.addWidget(QLabel("Transmission", parent=self), 7, 0) self.transmission_label = QLabel("--", parent=self) - self.transmission_label.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.transmission_label.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.transmission_label, 7, 1, 1, 3) self._layout.addWidget(QLabel("%", parent=self), 7, 4) self._layout.addWidget(QLabel("Detector distance", parent=self), 8, 0) self.dtz_label = QLabel(f"--", parent=self) - self.dtz_label.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.dtz_label.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.dtz_label, 8, 1, 1, 3) self._layout.addWidget(QLabel("mm", parent=self), 8, 4) self._layout.addWidget(QLabel("Target Dose", parent=self), 9, 0) self.target_dose_label = QLabel(f"--", parent=self) - self.target_dose_label.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.target_dose_label.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.target_dose_label, 9, 1, 1, 3) self._layout.addWidget(QLabel("MGy", parent=self), 9, 4) self._layout.addWidget(QLabel("Calculated Dose Rate", parent=self), 10, 0) self.calculated_dose_rate_label = QLabel(f"--", parent=self) - self.calculated_dose_rate_label.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.calculated_dose_rate_label.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.calculated_dose_rate_label, 10, 1, 1, 3) self._layout.addWidget(QLabel("MGy s-1", parent=self), 10, 4) self._layout.addWidget(QLabel("Wilson B Factor", parent=self), 11, 0) self.wilson_b_label = QLabel(f"--", parent=self) - self.wilson_b_label.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.wilson_b_label.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.wilson_b_label, 11, 1, 1, 3) self._layout.addWidget(QLabel("Å2", parent=self), 11, 4) self._layout.addWidget(QLabel("Crystal Size x", parent=self), 12, 0) self.xtal_x_label = QLabel(f"--", parent=self) - self.xtal_x_label.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.xtal_x_label.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.xtal_x_label, 12, 1, 1, 3) self._layout.addWidget(QLabel("um", parent=self), 12, 4) self._layout.addWidget(QLabel("Crystal Size y", parent=self), 13, 0) self.xtal_y_label = QLabel(f"--", parent=self) - self.xtal_y_label.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.xtal_y_label.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.xtal_y_label, 13, 1, 1, 3) self._layout.addWidget(QLabel("um", parent=self), 13, 4) self._layout.addWidget(QLabel("Crystal Size z", parent=self), 14, 0) self.xtal_z_label = QLabel(f"--", parent=self) - self.xtal_z_label.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.xtal_z_label.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.xtal_z_label, 14, 1, 1, 3) self._layout.addWidget(QLabel("um", parent=self), 14, 4) - self._layout.addWidget(QLabel("Calculated Dose (xtal size)", parent=self), 15, 0) + self._layout.addWidget( + QLabel("Calculated Dose (xtal size)", parent=self), 15, 0 + ) self.xtal_size_dose_label = QLabel(f"--", parent=self) - self.xtal_size_dose_label.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.xtal_size_dose_label.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.xtal_size_dose_label, 15, 1, 1, 3) self._layout.addWidget(QLabel("MGy", parent=self), 15, 4) self._layout.addWidget(QLabel("X-ray Wavelength", parent=self), 16, 0) self.wavelength_label = QLabel("--", parent=self) - self.wavelength_label.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.wavelength_label.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.wavelength_label, 16, 1, 1, 3) self._layout.addWidget(QLabel("Å", parent=self), 16, 4) self._layout.addWidget(QLabel("Flux", parent=self), 17, 0) self.flux_label = QLabel(f"--", parent=self) - self.flux_label.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.flux_label.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.flux_label, 17, 1, 1, 3) - self._layout.addWidget(QLabel("x 109 ph s-1", parent=self), 17, 4) + self._layout.addWidget( + QLabel("x 109 ph s-1", parent=self), 17, 4 + ) self._layout.addWidget(QLabel("Beam Size", parent=self), 18, 0) self.beam_size_label = QLabel(f"--", parent=self) - self.beam_size_label.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.beam_size_label.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.beam_size_label, 18, 1, 1, 3) self._layout.addWidget(QLabel("um2", parent=self), 18, 4) self._layout.addWidget(QLabel("Calculated Dose", parent=self), 19, 0) self.calculated_dose_label = QLabel(f"--", parent=self) - self.calculated_dose_label.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.calculated_dose_label.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.calculated_dose_label, 19, 1, 1, 3) self._layout.addWidget(QLabel("MGy", parent=self), 19, 4) self._layout.addWidget(QLabel("Total measurement time", parent=self), 20, 0) self.total_time = QLabel(f"{self.total_time_s} min 0 s") - self.total_time.setAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter) + self.total_time.setAlignment( + Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter + ) self._layout.addWidget(self.total_time, 20, 1, 1, 3) # add vertical stretch between detector distance and the run button - self._layout.addItem(QSpacerItem(0, 0, QSizePolicy.Policy.Minimum, QSizePolicy.Policy.Expanding), 21, 0, 1, 6) + self._layout.addItem( + QSpacerItem(0, 0, QSizePolicy.Policy.Minimum, QSizePolicy.Policy.Expanding), + 21, + 0, + 1, + 6, + ) # Run rotation button self.run_rotation_button = QPushButton("Run rotation", parent=self) @@ -192,9 +246,9 @@ class SimpleRotationSettingsPanel(QWidget): @Slot(DAQStatusModel) def update_daq_status(self, s: DAQStatusModel): self.__d = s - #TODO only update best_res after raster finished otherwise ask to update. or have toggle to overwrite with user value - #TODO only take best res from flat face scan - #TODO identify flat face!!!! + # TODO only update best_res after raster finished otherwise ask to update. or have toggle to overwrite with user value + # TODO only take best res from flat face scan + # TODO identify flat face!!!! best_res = s.last_best_res self.__omega = s.geom.omega_deg @@ -205,7 +259,6 @@ class SimpleRotationSettingsPanel(QWidget): self.visible_res_enter.update_value(v) self.set_visible_resolution(v) - self._wilson_b = s.last_best_b_factor if self._wilson_b is not None: @@ -221,7 +274,7 @@ class SimpleRotationSettingsPanel(QWidget): self.xtal_y_label.setText(f"{self.xtal_y:.2f}") self.xtal_z_label.setText(f"{self.xtal_z:.2f}") - #self.omega.new_value(s.geom.omega_deg) + # self.omega.new_value(s.geom.omega_deg) self.beam_size_label.setText( f"{s.geom.beam_size_mm.x * 1000.0} x {s.geom.beam_size_mm.y * 1000.0}" ) @@ -230,7 +283,9 @@ class SimpleRotationSettingsPanel(QWidget): if s.diffraction.wavelength_angstrom is None: self.wavelength_label.setText("N/A") else: - self.wavelength_label.setText(f"{s.diffraction.wavelength_angstrom:.3f}") + self.wavelength_label.setText( + f"{s.diffraction.wavelength_angstrom:.3f}" + ) self.update_calculated_labels() @Slot(float) @@ -268,39 +323,48 @@ class SimpleRotationSettingsPanel(QWidget): def update_calculated_labels(self): if self.__d is None: return - flux = 2.5e11 #TODO link flux - #TODO add start angle - #TODO link beam energy - #TODO read resolution estiamtion from jfjoch + flux = 2.5e11 # TODO link flux + # TODO add start angle + # TODO link beam energy + # TODO read resolution estiamtion from jfjoch total_angle = self.angular_range_enter.value d_vis = self.visible_res_enter.value image_angle = self.image_angle_enter.value - d_vis = d_vis - 0.2 #additional fudge factor that weights more towards high_res + d_vis = ( + d_vis - 0.2 + ) # additional fudge factor that weights more towards high_res if d_vis <= 0.0: d_vis = 1.3 - d_tar = 1/(1/d_vis + 0.1) + d_tar = 1 / (1 / d_vis + 0.1) self.target_res_label.setText(f"{d_tar:.2f}") Kdose = 2000 / (self.__d.diffraction.wavelength_angstrom**2) - beam_size_um_y = self.__d.geom.beam_size_mm.y * 1000 + beam_size_um_y = self.__d.geom.beam_size_mm.y * 1000 beam_size_um_x = self.__d.geom.beam_size_mm.x * 1000 self.dose_rate_MGy_s = (flux / (beam_size_um_x * beam_size_um_y * Kdose)) / 1e6 - if self.xtal_y is None or self.xtal_z is None or beam_size_um_x is None or beam_size_um_y is None: + if ( + self.xtal_y is None + or self.xtal_z is None + or beam_size_um_x is None + or beam_size_um_y is None + ): self.xtal_size_dose_rate_MGy_s = self.dose_rate_MGy_s - - elif self.xtal_y > beam_size_um_y or self.xtal_z > beam_size_um_y : + + elif self.xtal_y > beam_size_um_y or self.xtal_z > beam_size_um_y: multiplier_1 = max(self.xtal_y, beam_size_um_y) - multiplier_2 = max( self.xtal_z,beam_size_um_y) - new_beam_um_y = math.sqrt(multiplier_1*multiplier_2) - self.xtal_size_dose_rate_MGy_s = (flux / (beam_size_um_x * new_beam_um_y * Kdose)) / 1e6 + multiplier_2 = max(self.xtal_z, beam_size_um_y) + new_beam_um_y = math.sqrt(multiplier_1 * multiplier_2) + self.xtal_size_dose_rate_MGy_s = ( + flux / (beam_size_um_x * new_beam_um_y * Kdose) + ) / 1e6 else: self.xtal_size_dose_rate_MGy_s = self.dose_rate_MGy_s - + if self.temp_enter.value > 250: - self.target_dose_MGy = 0.1 * d_tar / 2.0 #TODO add user input + self.target_dose_MGy = 0.1 * d_tar / 2.0 # TODO add user input else: self.target_dose_MGy = 10 * d_tar / 2.0 @@ -311,11 +375,13 @@ class SimpleRotationSettingsPanel(QWidget): self.total_time_s = self.target_dose_MGy / self.dose_rate_MGy_s self.update_total_time_label() - self.calculated_dose_label.setText(f"{self.xtal_size_dose_rate_MGy_s*self.total_time_s:.2f}") + self.calculated_dose_label.setText( + f"{self.xtal_size_dose_rate_MGy_s * self.total_time_s:.2f}" + ) if image_angle == 0.0: image_angle = 0.001 self.n_images = round(total_angle / image_angle) - self.image_time_s = self.total_time_s / self.n_images + self.image_time_s = self.total_time_s / self.n_images if self.image_time_s < 0.0011: self.transmission = self.image_time_s / 0.0011 self.image_time_s = 0.0011 @@ -332,7 +398,9 @@ class SimpleRotationSettingsPanel(QWidget): if self.dtz <= 0.0: self.dtz_label.setText(f"""-""") elif self.dtz < self.__d.bl.dtz_min: - self.dtz_label.setText(f"""{self.dtz:.2f}""") + self.dtz_label.setText( + f"""{self.dtz:.2f}""" + ) self.dtz = self.__d.bl.dtz_min else: self.dtz_label.setText(f"{self.dtz:.2f}") @@ -344,23 +412,23 @@ class SimpleRotationSettingsPanel(QWidget): incr_omega_deg=image_angle, steps=self.n_images, transmission=self.transmission, - last_best_b_factor = self._wilson_b, - crystal_size = self.xtal_size, - flux_ph_s = None, - calculated_dose_rate_MGy_s = None, - xtal_size_dose_rate_MGy_s = None, - target_dose_MGy = None, - beam_size_x_um = beam_size_um_x, - beam_size_y_um = beam_size_um_y, - d_vis = d_vis, - d_tar = d_tar + last_best_b_factor=self._wilson_b, + crystal_size=self.xtal_size, + flux_ph_s=None, + calculated_dose_rate_MGy_s=None, + xtal_size_dose_rate_MGy_s=None, + target_dose_MGy=None, + beam_size_x_um=beam_size_um_x, + beam_size_y_um=beam_size_um_y, + d_vis=d_vis, + d_tar=d_tar, ) - #logger.debug("updated labels") + # logger.debug("updated labels") if self._prev_params != self.parameters: self._prev_params = self.parameters self.parameters_changed.emit(self.parameters) - #logger.debug("emitted parameters") + # logger.debug("emitted parameters") @Slot() def run_measurement(self): @@ -372,6 +440,7 @@ class SimpleRotationSettingsPanel(QWidget): dtz=self.dtz, transmission=self.transmission, screening=False, - exp_time_s=self.image_time_s) + exp_time_s=self.image_time_s, + ) self.rotation_scan.emit(r) - self.viewer_track_online.emit() \ No newline at end of file + self.viewer_track_online.emit() diff --git a/src/aare/gui/panels/status_panel.py b/src/aare/gui/panels/status_panel.py index 67051080..971f2efc 100644 --- a/src/aare/gui/panels/status_panel.py +++ b/src/aare/gui/panels/status_panel.py @@ -1,8 +1,8 @@ +from aarecommon.math.sample_geometry import SampleGeometryModel +from aarecommon.models.models import DAQStatusModel from PySide6.QtCore import Qt, Slot from PySide6.QtGui import QPixmap -from PySide6.QtWidgets import QFrame, QGridLayout, QLabel, QSpacerItem, QSizePolicy -from aare.common.models import DAQStatusModel -from aare.common.sample_geometry import SampleGeometryModel +from PySide6.QtWidgets import QFrame, QGridLayout, QLabel, QSizePolicy, QSpacerItem from aare.gui.widgets.status_label import StatusLabel from aare.gui.widgets.title_label import TitleLabel @@ -64,7 +64,9 @@ class StatusPanel(QFrame): ) if s.bl.ring_current_mA < 390.0: - self.ring_current.setText(f"{s.bl.ring_current_mA:.1f}") + self.ring_current.setText( + f'{s.bl.ring_current_mA:.1f}' + ) else: self.ring_current.setText(f"{s.bl.ring_current_mA:.1f}") @@ -72,4 +74,4 @@ class StatusPanel(QFrame): if s.diffraction.wavelength_angstrom is None: self.wavelength.setText("N/A") else: - self.wavelength.setText(f"{s.diffraction.wavelength_angstrom:.3f}") \ No newline at end of file + self.wavelength.setText(f"{s.diffraction.wavelength_angstrom:.3f}") diff --git a/src/aare/gui/panels/target_stability_panel.py b/src/aare/gui/panels/target_stability_panel.py index 51346a28..9ac5bfa7 100644 --- a/src/aare/gui/panels/target_stability_panel.py +++ b/src/aare/gui/panels/target_stability_panel.py @@ -1,28 +1,27 @@ import csv import math -import time import threading +import time from collections import deque +from aarecommon.config.logger import setup_logger from PySide6.QtCharts import QChart, QChartView, QLineSeries, QValueAxis -from PySide6.QtCore import Qt, Slot, QTimer, QPointF, QPoint -from PySide6.QtGui import QPainter, QPen, QColor, QMouseEvent, QWheelEvent +from PySide6.QtCore import QPoint, QPointF, Qt, QTimer, Slot +from PySide6.QtGui import QColor, QMouseEvent, QPainter, QPen, QWheelEvent from PySide6.QtWidgets import ( - QWidget, - QVBoxLayout, - QHBoxLayout, - QLabel, - QPushButton, - QDoubleSpinBox, - QFileDialog, QCheckBox, QComboBox, - QMessageBox, + QDoubleSpinBox, + QFileDialog, QGroupBox, + QHBoxLayout, + QLabel, + QMessageBox, + QPushButton, + QVBoxLayout, + QWidget, ) -from aare.common.logger_config import setup_logger - logger = setup_logger("aareGUI") @@ -77,7 +76,9 @@ class InteractiveChartView(QChartView): pos = event.position().toPoint() delta = pos - self._last_pos self._last_pos = pos - self._panel._pan_selected_axes(delta.x(), delta.y(), self.viewport().width(), self.viewport().height()) + self._panel._pan_selected_axes( + delta.x(), delta.y(), self.viewport().width(), self.viewport().height() + ) event.accept() return @@ -214,43 +215,61 @@ class TargetStabilityPanel(QWidget): self.help_button.clicked.connect(self._show_metrics_help) self.score_basis_combo = QComboBox() - self.score_basis_combo.addItems([self.SCORE_FROM_STEP_XY, self.SCORE_FROM_SIGMA_XY]) + self.score_basis_combo.addItems( + [self.SCORE_FROM_STEP_XY, self.SCORE_FROM_SIGMA_XY] + ) self.score_basis_combo.setCurrentText(self.SCORE_FROM_STEP_XY) self.score_basis_combo.currentTextChanged.connect(self._on_score_basis_changed) self.wheel_mode_combo = QComboBox() - self.wheel_mode_combo.addItems([ - self.WHEEL_X, - self.WHEEL_LEFT_Y, - self.WHEEL_RIGHT_Y, - self.WHEEL_SCORE, - ]) + self.wheel_mode_combo.addItems( + [ + self.WHEEL_X, + self.WHEEL_LEFT_Y, + self.WHEEL_RIGHT_Y, + self.WHEEL_SCORE, + ] + ) self.wheel_mode_combo.setCurrentText(self.WHEEL_X) - self.wheel_mode_combo.currentTextChanged.connect(lambda _text: self._update_status_label()) + self.wheel_mode_combo.currentTextChanged.connect( + lambda _text: self._update_status_label() + ) self.pan_mode_combo = QComboBox() - self.pan_mode_combo.addItems([ - self.PAN_ALL, - self.PAN_X, - self.PAN_LEFT_Y, - self.PAN_RIGHT_Y, - self.PAN_SCORE, - ]) + self.pan_mode_combo.addItems( + [ + self.PAN_ALL, + self.PAN_X, + self.PAN_LEFT_Y, + self.PAN_RIGHT_Y, + self.PAN_SCORE, + ] + ) self.pan_mode_combo.setCurrentText(self.PAN_ALL) - self.pan_mode_combo.currentTextChanged.connect(lambda _text: self._update_status_label()) + self.pan_mode_combo.currentTextChanged.connect( + lambda _text: self._update_status_label() + ) self.status_label = QLabel("Waiting for target-point data...") - self.status_label.setAlignment(Qt.AlignmentFlag.AlignLeft | Qt.AlignmentFlag.AlignVCenter) + self.status_label.setAlignment( + Qt.AlignmentFlag.AlignLeft | Qt.AlignmentFlag.AlignVCenter + ) - self.metrics_label = QLabel("Target: (-, -) | Beam: (-, -) | Distance: - px | Std dev: - px") - self.metrics_label.setAlignment(Qt.AlignmentFlag.AlignLeft | Qt.AlignmentFlag.AlignVCenter) + self.metrics_label = QLabel( + "Target: (-, -) | Beam: (-, -) | Distance: - px | Std dev: - px" + ) + self.metrics_label.setAlignment( + Qt.AlignmentFlag.AlignLeft | Qt.AlignmentFlag.AlignVCenter + ) self.metrics_label.setTextFormat(Qt.TextFormat.RichText) self.metrics_label.setWordWrap(True) self.controls_legend_label = QLabel( "Mouse: wheel=selected axis zoom | Shift+wheel=score | left-drag=selected pan mode | right-click=reset | click a trace point to target its axis" ) - self.controls_legend_label.setAlignment(Qt.AlignmentFlag.AlignLeft | Qt.AlignmentFlag.AlignVCenter) + self.controls_legend_label.setAlignment( + Qt.AlignmentFlag.AlignLeft | Qt.AlignmentFlag.AlignVCenter + ) self.controls_legend_label.setWordWrap(True) controls.addWidget(self.pause_button) @@ -327,7 +346,9 @@ class TargetStabilityPanel(QWidget): self.chart.addSeries(series) self.chart.legend().setVisible(True) - self.chart.setTitle("Target stability, distance, score, and step motion (last 60 s)") + self.chart.setTitle( + "Target stability, distance, score, and step motion (last 60 s)" + ) self.axis_x = QValueAxis() self.axis_x.setTitleText("Time [s ago]") @@ -350,7 +371,12 @@ class TargetStabilityPanel(QWidget): self.chart.addAxis(self.axis_y_distance, Qt.AlignmentFlag.AlignRight) self.chart.addAxis(self.axis_y_score, Qt.AlignmentFlag.AlignRight) - for series in (self.series, self.sigma_x_series, self.sigma_y_series, self.step_series): + for series in ( + self.series, + self.sigma_x_series, + self.sigma_y_series, + self.step_series, + ): series.attachAxis(self.axis_x) series.attachAxis(self.axis_y) @@ -460,7 +486,9 @@ class TargetStabilityPanel(QWidget): ) def _connect_series_selection(self, series: QLineSeries, target_name: str) -> None: - series.clicked.connect(lambda _point, target=target_name: self._select_trace_target(target)) + series.clicked.connect( + lambda _point, target=target_name: self._select_trace_target(target) + ) def _try_select_series_at_point(self, view_pos: QPoint) -> bool: """Check if click is near a visible series point and select it. Returns True if found.""" @@ -517,7 +545,9 @@ class TargetStabilityPanel(QWidget): def _set_paused(self, paused: bool) -> None: self._paused = paused if paused: - self.pause_button.setText("Resume (Live)" if self._collected_samples else "Resume") + self.pause_button.setText( + "Resume (Live)" if self._collected_samples else "Resume" + ) else: self._unfreeze_plot() self.pause_button.setText("Pause") @@ -613,9 +643,15 @@ class TargetStabilityPanel(QWidget): } return { - "rms_step_dx": math.sqrt(self._rolling_step_dx2_sum / self._rolling_step_count), - "rms_step_dy": math.sqrt(self._rolling_step_dy2_sum / self._rolling_step_count), - "rms_step_xy": math.sqrt(self._rolling_step_xy2_sum / self._rolling_step_count), + "rms_step_dx": math.sqrt( + self._rolling_step_dx2_sum / self._rolling_step_count + ), + "rms_step_dy": math.sqrt( + self._rolling_step_dy2_sum / self._rolling_step_count + ), + "rms_step_xy": math.sqrt( + self._rolling_step_xy2_sum / self._rolling_step_count + ), } def _current_score_value(self) -> float: @@ -628,7 +664,10 @@ class TargetStabilityPanel(QWidget): def _update_metrics_label(self, force: bool = False) -> None: now = time.monotonic() - if not force and (now - self._last_metrics_update_ts) < self._metrics_update_interval_s: + if ( + not force + and (now - self._last_metrics_update_ts) < self._metrics_update_interval_s + ): return self._last_metrics_update_ts = now @@ -777,16 +816,23 @@ class TargetStabilityPanel(QWidget): def hideEvent(self, event) -> None: super().hideEvent(event) - if self._track_only_when_visible and self._collecting_until is None and self._refresh_timer is not None: + if ( + self._track_only_when_visible + and self._collecting_until is None + and self._refresh_timer is not None + ): self._refresh_timer.stop() - @Slot(dict) def update_target_point(self, payload: dict) -> None: if self._paused or self._beam_center is None: return - if self._track_only_when_visible and not self.isVisible() and self._collecting_until is None: + if ( + self._track_only_when_visible + and not self.isVisible() + and self._collecting_until is None + ): return raw = payload.get("target_point") @@ -899,17 +945,19 @@ class TargetStabilityPanel(QWidget): else self._stability_score(step_xy) ) - self._plot_points.append({ - "ts": float(sample["ts"]), - "sigma": sigma_stats["sigma"], - "sigma_x": sigma_stats["std_dx"], - "sigma_y": sigma_stats["std_dy"], - "distance": float(sample["distance"]), - "dx": dx, - "dy": dy, - "step_xy": step_xy, - "score": score, - }) + self._plot_points.append( + { + "ts": float(sample["ts"]), + "sigma": sigma_stats["sigma"], + "sigma_x": sigma_stats["std_dx"], + "sigma_y": sigma_stats["std_dy"], + "distance": float(sample["distance"]), + "dx": dx, + "dy": dy, + "step_xy": step_xy, + "score": score, + } + ) def _remove_oldest_sigma_sample(self) -> None: if not self._sigma_samples: @@ -985,22 +1033,28 @@ class TargetStabilityPanel(QWidget): while self._plot_points and float(self._plot_points[0]["ts"]) < cutoff: self._plot_points.popleft() - def _data_limits(self) -> tuple[float, float, float, float, float, float, float, float]: + def _data_limits( + self, + ) -> tuple[float, float, float, float, float, float, float, float]: left_values: list[float] = [] right_values: list[float] = [] for point in self._plot_points: - left_values.extend([ - float(point["sigma"]), - float(point["sigma_x"]), - float(point["sigma_y"]), - float(point["step_xy"]), - ]) - right_values.extend([ - float(point["distance"]), - float(point["dx"]), - float(point["dy"]), - ]) + left_values.extend( + [ + float(point["sigma"]), + float(point["sigma_x"]), + float(point["sigma_y"]), + float(point["step_xy"]), + ] + ) + right_values.extend( + [ + float(point["distance"]), + float(point["dx"]), + float(point["dy"]), + ] + ) left_min = min(left_values, default=0.0) left_max = max(left_values, default=1.0) @@ -1010,8 +1064,12 @@ class TargetStabilityPanel(QWidget): left_span = max(1.0, left_max - left_min) right_span = max(1.0, right_max - right_min) - score_min = min((float(point["score"]) for point in self._plot_points), default=0.0) - score_max = max((float(point["score"]) for point in self._plot_points), default=100.0) + score_min = min( + (float(point["score"]) for point in self._plot_points), default=0.0 + ) + score_max = max( + (float(point["score"]) for point in self._plot_points), default=100.0 + ) score_lo = max(0.0, score_min - 5.0) score_hi = min(100.0, max(score_lo + 10.0, score_max + 5.0)) @@ -1028,7 +1086,9 @@ class TargetStabilityPanel(QWidget): def _reset_view(self) -> None: self._auto_scale_enabled = True - xmin, xmax, left_ymin, left_ymax, right_ymin, right_ymax, score_lo, score_hi = self._data_limits() + xmin, xmax, left_ymin, left_ymax, right_ymin, right_ymax, score_lo, score_hi = ( + self._data_limits() + ) self.axis_x.setRange(xmin, xmax) self.axis_y.setRange(left_ymin, left_ymax) self.axis_y_distance.setRange(right_ymin, right_ymax) @@ -1080,7 +1140,9 @@ class TargetStabilityPanel(QWidget): self._clamp_axes() self._update_status_label() - def _pan_selected_axes(self, dx_pixels: int, dy_pixels: int, width: int, height: int) -> None: + def _pan_selected_axes( + self, dx_pixels: int, dy_pixels: int, width: int, height: int + ) -> None: if width <= 0 or height <= 0: return @@ -1104,7 +1166,9 @@ class TargetStabilityPanel(QWidget): mode = self.pan_mode_combo.currentText() if mode == self.PAN_X: - self.axis_x.setRange(self.axis_x.min() + x_shift, self.axis_x.max() + x_shift) + self.axis_x.setRange( + self.axis_x.min() + x_shift, self.axis_x.max() + x_shift + ) elif mode == self.PAN_LEFT_Y: if ( @@ -1113,7 +1177,9 @@ class TargetStabilityPanel(QWidget): or self.show_sigma_y_cb.isChecked() or self.show_step_cb.isChecked() ): - self.axis_y.setRange(self.axis_y.min() + y_shift, self.axis_y.max() + y_shift) + self.axis_y.setRange( + self.axis_y.min() + y_shift, self.axis_y.max() + y_shift + ) elif mode == self.PAN_RIGHT_Y: if ( @@ -1134,7 +1200,9 @@ class TargetStabilityPanel(QWidget): ) else: - self.axis_x.setRange(self.axis_x.min() + x_shift, self.axis_x.max() + x_shift) + self.axis_x.setRange( + self.axis_x.min() + x_shift, self.axis_x.max() + x_shift + ) if ( self.show_sigma_cb.isChecked() @@ -1142,7 +1210,9 @@ class TargetStabilityPanel(QWidget): or self.show_sigma_y_cb.isChecked() or self.show_step_cb.isChecked() ): - self.axis_y.setRange(self.axis_y.min() + y_shift, self.axis_y.max() + y_shift) + self.axis_y.setRange( + self.axis_y.min() + y_shift, self.axis_y.max() + y_shift + ) if ( self.show_distance_cb.isChecked() @@ -1203,7 +1273,11 @@ class TargetStabilityPanel(QWidget): self._update_metrics_label() return - if self._paused and self._collecting_until is None and self._frozen_plot_points is None: + if ( + self._paused + and self._collecting_until is None + and self._frozen_plot_points is None + ): self._update_status_label() self._update_metrics_label() return @@ -1217,14 +1291,38 @@ class TargetStabilityPanel(QWidget): self._chart_dirty = False - sigma_points = [QPointF(float(point["ts"]) - ref_time, float(point["sigma"])) for point in plot_data] - sigma_x_points = [QPointF(float(point["ts"]) - ref_time, float(point["sigma_x"])) for point in plot_data] - sigma_y_points = [QPointF(float(point["ts"]) - ref_time, float(point["sigma_y"])) for point in plot_data] - distance_points = [QPointF(float(point["ts"]) - ref_time, float(point["distance"])) for point in plot_data] - dx_points = [QPointF(float(point["ts"]) - ref_time, float(point["dx"])) for point in plot_data] - dy_points = [QPointF(float(point["ts"]) - ref_time, float(point["dy"])) for point in plot_data] - score_points = [QPointF(float(point["ts"]) - ref_time, float(point["score"])) for point in plot_data] - step_points = [QPointF(float(point["ts"]) - ref_time, float(point["step_xy"])) for point in plot_data] + sigma_points = [ + QPointF(float(point["ts"]) - ref_time, float(point["sigma"])) + for point in plot_data + ] + sigma_x_points = [ + QPointF(float(point["ts"]) - ref_time, float(point["sigma_x"])) + for point in plot_data + ] + sigma_y_points = [ + QPointF(float(point["ts"]) - ref_time, float(point["sigma_y"])) + for point in plot_data + ] + distance_points = [ + QPointF(float(point["ts"]) - ref_time, float(point["distance"])) + for point in plot_data + ] + dx_points = [ + QPointF(float(point["ts"]) - ref_time, float(point["dx"])) + for point in plot_data + ] + dy_points = [ + QPointF(float(point["ts"]) - ref_time, float(point["dy"])) + for point in plot_data + ] + score_points = [ + QPointF(float(point["ts"]) - ref_time, float(point["score"])) + for point in plot_data + ] + step_points = [ + QPointF(float(point["ts"]) - ref_time, float(point["step_xy"])) + for point in plot_data + ] self.series.replace(sigma_points) self.sigma_x_series.replace(sigma_x_points) @@ -1236,7 +1334,16 @@ class TargetStabilityPanel(QWidget): self.step_series.replace(step_points) if self._auto_scale_enabled: - xmin, xmax, left_ymin, left_ymax, right_ymin, right_ymax, score_lo, score_hi = self._data_limits() + ( + xmin, + xmax, + left_ymin, + left_ymax, + right_ymin, + right_ymax, + score_lo, + score_hi, + ) = self._data_limits() self.axis_x.setRange(xmin, xmax) self.axis_y.setRange(left_ymin, left_ymax) self.axis_y_distance.setRange(right_ymin, right_ymax) @@ -1251,7 +1358,9 @@ class TargetStabilityPanel(QWidget): seconds = float(self.seconds_spin.value()) now = time.monotonic() rows = [s for s in self._samples if float(s["ts"]) >= now - seconds] - self._save_rows(rows, suggested_name=f"target_stability_last_{self._seconds_text()}s.csv") + self._save_rows( + rows, suggested_name=f"target_stability_last_{self._seconds_text()}s.csv" + ) def _collect_next_x_seconds(self) -> None: if self._paused: @@ -1304,8 +1413,14 @@ class TargetStabilityPanel(QWidget): if temp_rolling_count >= 2: mean_dx = temp_rolling_sum_dx / temp_rolling_count mean_dy = temp_rolling_sum_dy / temp_rolling_count - var_dx = max(0.0, (temp_rolling_sum_dx2 / temp_rolling_count) - (mean_dx * mean_dx)) - var_dy = max(0.0, (temp_rolling_sum_dy2 / temp_rolling_count) - (mean_dy * mean_dy)) + var_dx = max( + 0.0, + (temp_rolling_sum_dx2 / temp_rolling_count) - (mean_dx * mean_dx), + ) + var_dy = max( + 0.0, + (temp_rolling_sum_dy2 / temp_rolling_count) - (mean_dy * mean_dy), + ) std_dx = math.sqrt(var_dx) std_dy = math.sqrt(var_dy) sigma = math.hypot(std_dx, std_dy) @@ -1329,19 +1444,24 @@ class TargetStabilityPanel(QWidget): step_xy = math.sqrt(step_sum / count) score = self._stability_score( - sigma if self.score_basis_combo.currentText() == self.SCORE_FROM_SIGMA_XY else step_xy) + sigma + if self.score_basis_combo.currentText() == self.SCORE_FROM_SIGMA_XY + else step_xy + ) - self._frozen_plot_points.append({ - "ts": float(sample["ts"]), - "sigma": sigma, - "sigma_x": std_dx, - "sigma_y": std_dy, - "distance": float(sample["distance"]), - "dx": dx, - "dy": dy, - "step_xy": step_xy, - "score": score, - }) + self._frozen_plot_points.append( + { + "ts": float(sample["ts"]), + "sigma": sigma, + "sigma_x": std_dx, + "sigma_y": std_dy, + "distance": float(sample["distance"]), + "dx": dx, + "dy": dy, + "step_xy": step_xy, + "score": score, + } + ) self._paused = True self.pause_button.setChecked(True) @@ -1353,7 +1473,10 @@ class TargetStabilityPanel(QWidget): # Schedule the save dialog to run after the current event processing # This avoids blocking while any locks might be held file_name = f"target_stability_collected_{self._seconds_text()}s.csv" - QTimer.singleShot(0, lambda: self._save_rows(self._collected_samples, suggested_name=file_name)) + QTimer.singleShot( + 0, + lambda: self._save_rows(self._collected_samples, suggested_name=file_name), + ) self._update_status_label() def _unfreeze_plot(self) -> None: @@ -1380,12 +1503,14 @@ class TargetStabilityPanel(QWidget): writer = csv.writer(f) writer.writerow(["t_monotonic_s", "dx_px", "dy_px", "distance_px"]) for row in rows: - writer.writerow([ - f"{float(row['ts']):.6f}", - f"{float(row['dx']):.6f}", - f"{float(row['dy']):.6f}", - f"{float(row['distance']):.6f}", - ]) + writer.writerow( + [ + f"{float(row['ts']):.6f}", + f"{float(row['dx']):.6f}", + f"{float(row['dy']):.6f}", + f"{float(row['distance']):.6f}", + ] + ) @staticmethod def _coerce_target_point(raw) -> tuple[float, float] | None: @@ -1397,4 +1522,4 @@ class TargetStabilityPanel(QWidget): return float(raw[0]), float(raw[1]) except Exception as e: logger.warning(f"Failed to parse target point {raw}: {e}") - return None \ No newline at end of file + return None diff --git a/src/aare/gui/panels/tell_sample_panel.py b/src/aare/gui/panels/tell_sample_panel.py index 32228876..ddd5ae84 100644 --- a/src/aare/gui/panels/tell_sample_panel.py +++ b/src/aare/gui/panels/tell_sample_panel.py @@ -1,23 +1,28 @@ +from aarecommon.config.logger import setup_logger +from aarecommon.models.models import ( + BeamlineStateEnum, + DAQStatusModel, + SampleShortInfo, + SampleShortInfoList, +) from PySide6.QtCore import Qt, Signal, Slot from PySide6.QtWidgets import ( + QAbstractItemView, QFrame, QGridLayout, - QTableView, QHeaderView, - QMenu, QLabel, + QMenu, QPushButton, - QAbstractItemView, + QTableView, ) -from aare.common.logger_config import setup_logger -from aare.common.models import BeamlineStateEnum, SampleShortInfo, SampleShortInfoList, DAQStatusModel - from aare.gui.models.user_sample_model import UserSampleSpreadsheet from aare.gui.widgets.title_label import TitleLabel logger = setup_logger("aareGUI") + class TellSamplePanel(QFrame): mount = Signal(SampleShortInfo) unmount = Signal() @@ -134,7 +139,10 @@ class TellSamplePanel(QFrame): for p in presets: act = prefix_menu.addAction(p) act.triggered.connect( - lambda checked=False, vv=p: self.table_model.set_column_filter(logical_index, vv)) + lambda checked=False, vv=p: self.table_model.set_column_filter( + logical_index, vv + ) + ) elif logical_index == 3: segs, segpos = self.table_model.suggested_prefixes_for_location() if segs: @@ -142,13 +150,19 @@ class TellSamplePanel(QFrame): for s in segs: act = seg_menu.addAction(s) act.triggered.connect( - lambda checked=False, vv=s: self.table_model.set_column_filter(logical_index, vv)) + lambda checked=False, vv=s: self.table_model.set_column_filter( + logical_index, vv + ) + ) if segpos: sp_menu = menu.addMenu("Filter by segment+position (e.g. B3)") for sp in segpos: act = sp_menu.addAction(sp) act.triggered.connect( - lambda checked=False, vv=sp: self.table_model.set_column_filter(logical_index, vv)) + lambda checked=False, vv=sp: self.table_model.set_column_filter( + logical_index, vv + ) + ) else: values = self.table_model.unique_values_for_column(logical_index) if values: @@ -156,7 +170,10 @@ class TellSamplePanel(QFrame): for v in values: act = choose_menu.addAction(v) act.triggered.connect( - lambda checked=False, vv=v: self.table_model.set_column_filter(logical_index, vv)) + lambda checked=False, vv=v: self.table_model.set_column_filter( + logical_index, vv + ) + ) # Manual entry and clear options set_filter_action = menu.addAction(f"Filter column: {col_name}...") @@ -168,8 +185,12 @@ class TellSamplePanel(QFrame): toggle_all = menu.addAction("Show all pgroups (ignore current p-group)") toggle_all.setCheckable(True) toggle_all.setChecked(self.table_model.show_all_pgroups) + def _toggle_all(): - self.table_model.set_show_all_pgroups(not self.table_model.show_all_pgroups) + self.table_model.set_show_all_pgroups( + not self.table_model.show_all_pgroups + ) + toggle_all.triggered.connect(_toggle_all) action = menu.exec_(header.mapToGlobal(pos)) @@ -178,7 +199,10 @@ class TellSamplePanel(QFrame): if action == set_filter_action: from PySide6.QtWidgets import QInputDialog - text, ok = QInputDialog.getText(self, "Set filter", f"Filter for '{col_name}':") + + text, ok = QInputDialog.getText( + self, "Set filter", f"Filter for '{col_name}':" + ) if ok: self.table_model.set_column_filter(logical_index, text) elif action == clear_filter_action: @@ -194,7 +218,9 @@ class TellSamplePanel(QFrame): tell_details = "" if tell_state is not None: activity = tell_state.activity.display_name() - phase = tell_state.phase.display_name() if tell_state.phase is not None else "" + phase = ( + tell_state.phase.display_name() if tell_state.phase is not None else "" + ) message = (tell_state.message or "").strip() tell_parts = [activity] @@ -211,7 +237,9 @@ class TellSamplePanel(QFrame): else: try: if sample.location is None: - base_text = f"Current sample: {sample.sample_name} (Manual mount)" + base_text = ( + f"Current sample: {sample.sample_name} (Manual mount)" + ) else: base_text = ( f"Current sample: {sample.sample_name} " @@ -224,10 +252,12 @@ class TellSamplePanel(QFrame): base_text = f"Confusing information :/ {e}" if tell_details: - self.curr_sample_label.setText(f"{base_text}
TELL: {tell_details}") + self.curr_sample_label.setText( + f"{base_text}
TELL: {tell_details}" + ) else: self.curr_sample_label.setText(base_text) if status.session.current_pgroup is not None: self.__current_pgroup = status.session.current_pgroup - self.table_model.set_default_user_filter(self.__current_pgroup) \ No newline at end of file + self.table_model.set_default_user_filter(self.__current_pgroup) diff --git a/src/aare/gui/panels/zoom_panel.py b/src/aare/gui/panels/zoom_panel.py index 3911a901..678cf26e 100644 --- a/src/aare/gui/panels/zoom_panel.py +++ b/src/aare/gui/panels/zoom_panel.py @@ -1,7 +1,7 @@ -from PySide6.QtCore import Slot, Signal -from PySide6.QtWidgets import QWidget, QGridLayout +from aarecommon.models.models import DAQStatusModel +from PySide6.QtCore import Signal, Slot +from PySide6.QtWidgets import QGridLayout, QWidget -from aare.common.models import DAQStatusModel from aare.gui.widgets.button_with_payload import ButtonWithPayload from aare.gui.widgets.title_label import TitleLabel @@ -38,14 +38,13 @@ class ZoomPanel(QWidget): def set_zoom(self, payload: dict): self.zoom.emit(payload["zoom"]) - @Slot(DAQStatusModel) def update_daq_status(self, status: DAQStatusModel): self.__zoom = status.bl.zoom for z in self.__zoom_settings: - if z["value"] == self.__zoom: - self.__buttons[self.__zoom_settings.index(z)].setStyleSheet("font-weight: bold;") + self.__buttons[self.__zoom_settings.index(z)].setStyleSheet( + "font-weight: bold;" + ) else: self.__buttons[self.__zoom_settings.index(z)].setStyleSheet("") - diff --git a/src/aare/gui/scan_logic/raster_grid_manager.py b/src/aare/gui/scan_logic/raster_grid_manager.py index 19db4913..cfbbb2a8 100644 --- a/src/aare/gui/scan_logic/raster_grid_manager.py +++ b/src/aare/gui/scan_logic/raster_grid_manager.py @@ -1,20 +1,21 @@ import math from enum import Enum - -import numpy as np -from PySide6.QtCore import QObject, Signal, Slot, QPointF, QRectF, QLineF -from PySide6.QtGui import QPainter, QPen, QColor, QBrush, QImage -from PySide6.QtCore import Qt, QRect from typing import List, Tuple -from aare.common.beamline import mx_beamline, get_jfjoch_url -from aare.common.coordinate import Coordinate, SmargonCoordinate -from aare.common.models import DAQStatusModel -from aare.common.find_xtal import compute_crystal_score_array -from aare.common.raster_grid import RasterGridRequest, CompletedRasterGrid, CompletedRasterGridElem -from aare.common.sample_geometry import SampleGeometryModel - -from aare.common.logger_config import setup_logger +import numpy as np +from aarecommon.config.beamline import get_jfjoch_url, mx_beamline +from aarecommon.config.logger import setup_logger +from aarecommon.math.coordinate import Coordinate, SmargonCoordinate +from aarecommon.math.find_xtal import compute_crystal_score_array +from aarecommon.math.sample_geometry import SampleGeometryModel +from aarecommon.models.models import DAQStatusModel +from aarecommon.models.raster_grid import ( + CompletedRasterGrid, + CompletedRasterGridElem, + RasterGridRequest, +) +from PySide6.QtCore import QLineF, QObject, QPointF, QRect, QRectF, Qt, Signal, Slot +from PySide6.QtGui import QBrush, QColor, QImage, QPainter, QPen logger = setup_logger("aareGUI") @@ -23,7 +24,7 @@ class RasterGridMetric(Enum): BKG = 0 SPOTS = 1 INDEXING = 2 - PR = 3 # JD: repurposing PR for testing scoring + PR = 3 # JD: repurposing PR for testing scoring BFACTOR = 4 SPOTS_LOW_RES = 5 RES = 6 @@ -32,6 +33,7 @@ class RasterGridMetric(Enum): SPOTS_ICE_LOW_RES = 9 RASTER_SCORE = 10 + # Viridis colormap colors (from matplotlib), dark purple -> yellow _VIRIDIS_COLORS = [ (68, 1, 84), # Dark purple @@ -44,7 +46,7 @@ _VIRIDIS_COLORS = [ (68, 190, 112), (121, 209, 81), (189, 223, 38), - (253, 231, 37) # Yellow + (253, 231, 37), # Yellow ] # Float (N, 3) lookup table for vectorised colour mapping. @@ -92,6 +94,7 @@ def float_to_viridis_brush(value: float, alpha: int = 127) -> QBrush: return QBrush(QColor(color[0], color[1], color[2], alpha)) + def normalize_angle(angle_deg: float) -> float: normalized_angle = math.fmod(angle_deg + 180, 360) - 180.0 return normalized_angle @@ -128,25 +131,28 @@ class RasterGridManager(QObject): self.__loaded_image_prefix = None self.__loaded_image_index = None - self.__start_point : QPointF = QPointF(0, 0) - self.__active_grid : RasterGridRequest = RasterGridRequest( - n_x= 0, - n_y= 0, - smargon_top_left = self.__geom.smargon, - grid_size_mm=Coordinate(x= 0.8 * self.__beam_size_mm.x, - y= 0.8 * self.__beam_size_mm.y), + self.__start_point: QPointF = QPointF(0, 0) + self.__active_grid: RasterGridRequest = RasterGridRequest( + n_x=0, + n_y=0, + smargon_top_left=self.__geom.smargon, + grid_size_mm=Coordinate( + x=0.8 * self.__beam_size_mm.x, y=0.8 * self.__beam_size_mm.y + ), omega_deg=self.__geom.omega_deg, exp_time_s=0.02, transmission=1.0, - dtz=200.0 + dtz=200.0, ) - self.__completed_grids : List[CompletedRasterGridElem] = [] + self.__completed_grids: List[CompletedRasterGridElem] = [] # Cache of pre-rendered heatmap bitmaps, keyed by (id(grid_elem), metric). # Each entry is (QImage, backing ndarray); the ndarray must be kept alive # because QImage shares its buffer without copying. Rebuilt only when the # data or metric changes, not on every repaint (sample move / zoom). - self.__heatmap_cache: dict[tuple[int, "RasterGridMetric"], tuple[QImage, np.ndarray]] = {} + self.__heatmap_cache: dict[ + tuple[int, "RasterGridMetric"], tuple[QImage, np.ndarray] + ] = {} @property def active_grid(self) -> RasterGridRequest: @@ -164,7 +170,9 @@ class RasterGridManager(QObject): return True return False - def _grid_pixel_geometry(self, grid: RasterGridRequest) -> tuple[float, float, float, float] | None: + def _grid_pixel_geometry( + self, grid: RasterGridRequest + ) -> tuple[float, float, float, float] | None: if not self._is_grid_visible(grid): return None @@ -206,9 +214,13 @@ class RasterGridManager(QObject): return 0, grid.n_x, 0, grid.n_y min_x = max(0, int(math.floor((visible_rect.left() - start_x) / cell_w)) - 1) - max_x = min(grid.n_x, int(math.ceil((visible_rect.right() - start_x) / cell_w)) + 1) + max_x = min( + grid.n_x, int(math.ceil((visible_rect.right() - start_x) / cell_w)) + 1 + ) min_y = max(0, int(math.floor((visible_rect.top() - start_y) / cell_h)) - 1) - max_y = min(grid.n_y, int(math.ceil((visible_rect.bottom() - start_y) / cell_h)) + 1) + max_y = min( + grid.n_y, int(math.ceil((visible_rect.bottom() - start_y) / cell_h)) + 1 + ) return min_x, max_x, min_y, max_y @@ -234,7 +246,9 @@ class RasterGridManager(QObject): self.clear_active_grid() if self.__beam_size_mm != s.geom.beam_size_mm: self.__beam_size_mm = s.geom.beam_size_mm - self.update_grid_size(0.8 * self.__beam_size_mm.x, 0.8 * self.__beam_size_mm.y) + self.update_grid_size( + 0.8 * self.__beam_size_mm.x, 0.8 * self.__beam_size_mm.y + ) self.__geom = s.geom def resize_active_grid(self, end_point: QPointF): @@ -298,7 +312,9 @@ class RasterGridManager(QObject): self.__active_grid.grid_size_mm.y, ) - def get_grid_coord(self, grid: RasterGridRequest, point: QPointF) -> Tuple[int, int]: + def get_grid_coord( + self, grid: RasterGridRequest, point: QPointF + ) -> Tuple[int, int]: point_bl = self.__geom.picture_to_sample(Coordinate(x=point.x(), y=point.y())) delta = point_bl - self.__geom.smargon_to_beamline(grid.smargon_top_left.sh_mm) @@ -315,7 +331,6 @@ class RasterGridManager(QObject): def load_image(self, point: QPointF): for grid in self.__completed_grids: if self._is_grid_visible(grid.request): - n_x = grid.request.n_x n_y = grid.request.n_y cell_x, cell_y = self.get_grid_coord(grid.request, point) @@ -323,25 +338,35 @@ class RasterGridManager(QObject): if 0 <= cell_x < n_x and 0 <= cell_y < n_y: cell = cell_y * n_x + cell_x - if (self.__loaded_image_prefix != grid.result.file_prefix or - self.__loaded_image_index != grid.result.images[cell].number): + if ( + self.__loaded_image_prefix != grid.result.file_prefix + or self.__loaded_image_index != grid.result.images[cell].number + ): self.__loaded_image_prefix = grid.result.file_prefix self.__loaded_image_index = grid.result.images[cell].number - logger.debug(f"Load {grid.result.file_prefix} {grid.result.images[cell].number}") - self.image_selected.emit(grid.result.file_prefix, grid.result.images[cell].number) - #TODO if not in the same PGroup, user can stream from last run only but never load from a file. - #If loading a prior run, this shoudl throw an error, - #if user in same pgroup, load will always be fine. - #self.image_selected.emit(self.__detector_url, grid.result.images[cell].number) + logger.debug( + f"Load {grid.result.file_prefix} {grid.result.images[cell].number}" + ) + self.image_selected.emit( + grid.result.file_prefix, grid.result.images[cell].number + ) + # TODO if not in the same PGroup, user can stream from last run only but never load from a file. + # If loading a prior run, this shoudl throw an error, + # if user in same pgroup, load will always be fine. + # self.image_selected.emit(self.__detector_url, grid.result.images[cell].number) def is_part_of_active_grid(self, point: QPointF) -> bool: if not self._is_grid_visible(self.__active_grid): return False point_bl = self.__geom.picture_to_sample(Coordinate(x=point.x(), y=point.y())) - delta = point_bl - self.__geom.smargon_to_beamline(self.__active_grid.smargon_top_left.sh_mm) - return (0 <= delta.x < self.__active_grid.n_x * self.__active_grid.grid_size_mm.x) and ( - 0 <= delta.y < self.__active_grid.n_y * self.__active_grid.grid_size_mm.y + delta = point_bl - self.__geom.smargon_to_beamline( + self.__active_grid.smargon_top_left.sh_mm + ) + return ( + 0 <= delta.x < self.__active_grid.n_x * self.__active_grid.grid_size_mm.x + ) and ( + 0 <= delta.y < self.__active_grid.n_y * self.__active_grid.grid_size_mm.y ) def is_part_of_completed_grid(self, point: QPointF) -> str | None: @@ -354,54 +379,92 @@ class RasterGridManager(QObject): if 0 <= cell_x < n_x and 0 <= cell_y < n_y: cell = cell_y * n_x + cell_x txt = "" - if cell < len(grid.result.images) and grid.result.images[cell].efficiency == 1.0: + if ( + cell < len(grid.result.images) + and grid.result.images[cell].efficiency == 1.0 + ): txt += f"Image {grid.result.images[cell].number}
" - if grid.result.images[cell].bkg is not None and grid.result.images[cell].bkg >= 0: + if ( + grid.result.images[cell].bkg is not None + and grid.result.images[cell].bkg >= 0 + ): txt += f"Background estimate {grid.result.images[cell].bkg:.2f}
" - if grid.result.images[cell].spots is not None and grid.result.images[cell].spots >= 0: + if ( + grid.result.images[cell].spots is not None + and grid.result.images[cell].spots >= 0 + ): txt += f"Spot count {grid.result.images[cell].spots}
" - if grid.result.images[cell].spots_ice is not None and grid.result.images[cell].spots_ice >= 0: + if ( + grid.result.images[cell].spots_ice is not None + and grid.result.images[cell].spots_ice >= 0 + ): txt += f"Spot count (ice) {grid.result.images[cell].spots_ice}
" - if self.spot_ice_ratio(grid.result.images[cell]) is not None and self.spot_ice_ratio(grid.result.images[cell]) >= 0: + if ( + self.spot_ice_ratio(grid.result.images[cell]) is not None + and self.spot_ice_ratio(grid.result.images[cell]) >= 0 + ): txt += f"Spot ratio (ice/low res.) {self.spot_ice_ratio(grid.result.images[cell])}
" - if grid.result.images[cell].spots_low_res is not None and grid.result.images[cell].spots_low_res >= 0: + if ( + grid.result.images[cell].spots_low_res is not None + and grid.result.images[cell].spots_low_res >= 0 + ): txt += f"Spot count (low res.) {grid.result.images[cell].spots_low_res}
" - if grid.result.images[cell].index is not None and grid.result.images[cell].index > 0 and grid.result.images[cell].uc is not None: + if ( + grid.result.images[cell].index is not None + and grid.result.images[cell].index > 0 + and grid.result.images[cell].uc is not None + ): txt += f"Indexed {grid.result.images[cell].uc.a:.1f} {grid.result.images[cell].uc.b:.1f} {grid.result.images[cell].uc.c:.1f} {grid.result.images[cell].uc.alpha:.1f} {grid.result.images[cell].uc.beta:.1f} {grid.result.images[cell].uc.gamma:.1f}
" - if grid.result.images[cell].b is not None and grid.result.images[cell].b >= 0: + if ( + grid.result.images[cell].b is not None + and grid.result.images[cell].b >= 0 + ): txt += f"B-factor {grid.result.images[cell].b:.2f}
" - if grid.result.images[cell].pr is not None and grid.result.images[cell].pr >= 0: - txt += f"Profile Radius {grid.result.images[cell].pr:.2f}
" - if grid.result.images[cell].nx is not None and grid.result.images[cell].ny is not None: + if ( + grid.result.images[cell].pr is not None + and grid.result.images[cell].pr >= 0 + ): + txt += ( + f"Profile Radius {grid.result.images[cell].pr:.2f}
" + ) + if ( + grid.result.images[cell].nx is not None + and grid.result.images[cell].ny is not None + ): score = compute_crystal_score_array(grid.result.images) txt += f"Raster score {score[grid.result.images[cell].nx, grid.result.images[cell].ny]:.2f}" return txt return None - @Slot(float, float) def update_grid_size(self, grid_size_mm_x: float, grid_size_mm_y: float): if (grid_size_mm_x <= 0) or (grid_size_mm_y <= 0): raise ValueError("Grid size must be positive") - self.__active_grid.n_x = round(self.__active_grid.n_x * self.__active_grid.grid_size_mm.x / grid_size_mm_x) - self.__active_grid.n_y = round(self.__active_grid.n_y * self.__active_grid.grid_size_mm.y / grid_size_mm_y) + self.__active_grid.n_x = round( + self.__active_grid.n_x * self.__active_grid.grid_size_mm.x / grid_size_mm_x + ) + self.__active_grid.n_y = round( + self.__active_grid.n_y * self.__active_grid.grid_size_mm.y / grid_size_mm_y + ) self.__active_grid.grid_size_mm = Coordinate(x=grid_size_mm_x, y=grid_size_mm_y) - self.grid_scan_size_changed.emit(self.__active_grid.n_x, - self.__active_grid.n_y, - self.__active_grid.grid_size_mm.x, - self.__active_grid.grid_size_mm.y) + self.grid_scan_size_changed.emit( + self.__active_grid.n_x, + self.__active_grid.n_y, + self.__active_grid.grid_size_mm.x, + self.__active_grid.grid_size_mm.y, + ) @Slot() def run_grid_scan(self): @@ -415,7 +478,11 @@ class RasterGridManager(QObject): n_y=ag.n_y, grid_size_mm=Coordinate(x=ag.grid_size_mm.x, y=ag.grid_size_mm.y), smargon_top_left=SmargonCoordinate( - sh_mm=Coordinate(x=ag.smargon_top_left.sh_mm.x, y=ag.smargon_top_left.sh_mm.y, z=ag.smargon_top_left.sh_mm.z), + sh_mm=Coordinate( + x=ag.smargon_top_left.sh_mm.x, + y=ag.smargon_top_left.sh_mm.y, + z=ag.smargon_top_left.sh_mm.z, + ), phi_deg=ag.smargon_top_left.phi_deg, chi_deg=ag.smargon_top_left.chi_deg, ), @@ -470,10 +537,17 @@ class RasterGridManager(QObject): case RasterGridMetric.INDEXING: v = [obj.index for obj in i.result.images] case RasterGridMetric.PR: - v = [obj.index / max(obj.spots_low_res, 1) for obj in i.result.images] + v = [ + obj.index / max(obj.spots_low_res, 1) for obj in i.result.images + ] case RasterGridMetric.RASTER_SCORE: score = compute_crystal_score_array(i.result.images) - v = [score[obj.nx, obj.ny] if obj.nx is not None and obj.ny is not None else None for obj in i.result.images] + v = [ + score[obj.nx, obj.ny] + if obj.nx is not None and obj.ny is not None + else None + for obj in i.result.images + ] case RasterGridMetric.BFACTOR: v = [obj.b for obj in i.result.images] case RasterGridMetric.RES: @@ -551,9 +625,7 @@ class RasterGridManager(QObject): # (n_y rows, n_x cols, RGBA); contiguous so QImage can share the buffer. rgba = np.ascontiguousarray(rgba.reshape(n_y, n_x, 4)) - image = QImage( - rgba.data, n_x, n_y, 4 * n_x, QImage.Format.Format_RGBA8888 - ) + image = QImage(rgba.data, n_x, n_y, 4 * n_x, QImage.Format.Format_RGBA8888) return image, rgba def _draw_completed_heatmap( @@ -601,12 +673,14 @@ class RasterGridManager(QObject): painter.drawImage(bounds, image) painter.setOpacity(1.0) - painter.setPen(QPen(QColor(114, 159, 207, min(255, alpha + 40)), 1, Qt.PenStyle.SolidLine)) + painter.setPen( + QPen(QColor(114, 159, 207, min(255, alpha + 40)), 1, Qt.PenStyle.SolidLine) + ) painter.setBrush(Qt.BrushStyle.NoBrush) painter.drawRect(bounds) painter.restore() - #TODO make sure draw_grid is visualising the grid correctly, correct orientation, correct x/y labelling!!!! + # TODO make sure draw_grid is visualising the grid correctly, correct orientation, correct x/y labelling!!!! def _draw_grid( self, painter: QPainter, @@ -633,11 +707,17 @@ class RasterGridManager(QObject): if bounds is None: return - if visible_rect is not None and not visible_rect.isEmpty() and not bounds.intersects(visible_rect): + if ( + visible_rect is not None + and not visible_rect.isEmpty() + and not bounds.intersects(visible_rect) + ): return if values is not None and (cell_w < 3.0 or cell_h < 3.0): - if self._draw_completed_grid_fast(painter, grid, values, alpha, visible_rect): + if self._draw_completed_grid_fast( + painter, grid, values, alpha, visible_rect + ): return painter.save() @@ -645,8 +725,12 @@ class RasterGridManager(QObject): if values is not None: painter.setPen(Qt.PenStyle.NoPen) - min_value = min((x for x in values if x is not None and not math.isnan(x)), default=0) - max_value = max((x for x in values if x is not None and not math.isnan(x)), default=1) + min_value = min( + (x for x in values if x is not None and not math.isnan(x)), default=0 + ) + max_value = max( + (x for x in values if x is not None and not math.isnan(x)), default=1 + ) diff = 1 if min_value == max_value else (max_value - min_value) else: painter.setPen(QPen(QColor(114, 159, 207), 1, Qt.PenStyle.SolidLine)) @@ -665,7 +749,9 @@ class RasterGridManager(QObject): if values is None: painter.setBrush(Qt.BrushStyle.NoBrush) - painter.drawRect(QRect(round(px), round(py), round(cell_w), round(cell_h))) + painter.drawRect( + QRect(round(px), round(py), round(cell_w), round(cell_h)) + ) continue idx = x + y * grid.n_x @@ -677,13 +763,23 @@ class RasterGridManager(QObject): ): continue - brush = float_to_viridis_brush((values[idx] - min_value) / diff, alpha=alpha) - painter.fillRect(QRectF(px, py, max(1.0, cell_w), max(1.0, cell_h)), brush) + brush = float_to_viridis_brush( + (values[idx] - min_value) / diff, alpha=alpha + ) + painter.fillRect( + QRectF(px, py, max(1.0, cell_w), max(1.0, cell_h)), brush + ) if values is None: painter.setBrush(Qt.BrushStyle.NoBrush) else: - painter.setPen(QPen(QColor(114, 159, 207, min(255, alpha + 40)), 1, Qt.PenStyle.SolidLine)) + painter.setPen( + QPen( + QColor(114, 159, 207, min(255, alpha + 40)), + 1, + Qt.PenStyle.SolidLine, + ) + ) painter.setBrush(Qt.BrushStyle.NoBrush) painter.drawRect(bounds) @@ -704,7 +800,11 @@ class RasterGridManager(QObject): if bounds is None: return False - if visible_rect is not None and not visible_rect.isEmpty() and not bounds.intersects(visible_rect): + if ( + visible_rect is not None + and not visible_rect.isEmpty() + and not bounds.intersects(visible_rect) + ): return True painter.save() @@ -762,10 +862,16 @@ class RasterGridManager(QObject): if bounds is None: return False - if visible_rect is not None and not visible_rect.isEmpty() and not bounds.intersects(visible_rect): + if ( + visible_rect is not None + and not visible_rect.isEmpty() + and not bounds.intersects(visible_rect) + ): return True - valid_values = [x for x in values if x is not None and not math.isnan(x) and x >= 0] + valid_values = [ + x for x in values if x is not None and not math.isnan(x) and x >= 0 + ] min_value = min(valid_values, default=0) max_value = max(valid_values, default=1) diff = 1 if min_value == max_value else (max_value - min_value) @@ -804,7 +910,9 @@ class RasterGridManager(QObject): brush = float_to_viridis_brush((value - min_value) / diff, alpha=alpha) painter.fillRect(QRectF(px, py, draw_w, draw_h), brush) - painter.setPen(QPen(QColor(114, 159, 207, min(255, alpha + 40)), 1, Qt.PenStyle.SolidLine)) + painter.setPen( + QPen(QColor(114, 159, 207, min(255, alpha + 40)), 1, Qt.PenStyle.SolidLine) + ) painter.setBrush(Qt.BrushStyle.NoBrush) painter.drawRect(bounds) @@ -894,10 +1002,16 @@ class RasterGridManager(QObject): def completed_grid_scan_goto(self, row: int): if 0 <= row < len(self.__completed_grids): self.omega.emit(self.__completed_grids[row].request.omega_deg) - self.smargon.emit(SmargonCoordinate( - phi_deg=self.__completed_grids[row].request.smargon_top_left.phi_deg, - chi_deg=self.__completed_grids[row].request.smargon_top_left.chi_deg, - )) + self.smargon.emit( + SmargonCoordinate( + phi_deg=self.__completed_grids[ + row + ].request.smargon_top_left.phi_deg, + chi_deg=self.__completed_grids[ + row + ].request.smargon_top_left.chi_deg, + ) + ) @Slot(int) def completed_grid_redo(self, row: int): @@ -905,9 +1019,15 @@ class RasterGridManager(QObject): src = self.__completed_grids[row].request self.__active_grid.n_x = src.n_x self.__active_grid.n_y = src.n_y - self.__active_grid.grid_size_mm = Coordinate(x=src.grid_size_mm.x, y=src.grid_size_mm.y) + self.__active_grid.grid_size_mm = Coordinate( + x=src.grid_size_mm.x, y=src.grid_size_mm.y + ) self.__active_grid.smargon_top_left = SmargonCoordinate( - sh_mm=Coordinate(x=src.smargon_top_left.sh_mm.x, y=src.smargon_top_left.sh_mm.y, z=src.smargon_top_left.sh_mm.z), + sh_mm=Coordinate( + x=src.smargon_top_left.sh_mm.x, + y=src.smargon_top_left.sh_mm.y, + z=src.smargon_top_left.sh_mm.z, + ), phi_deg=src.smargon_top_left.phi_deg, chi_deg=src.smargon_top_left.chi_deg, ) @@ -916,11 +1036,13 @@ class RasterGridManager(QObject): self.__active_grid.visible = True self.__completed_grids[row].request.visible = False - self.grid_scan_size_changed.emit(self.__active_grid.n_x, - self.__active_grid.n_y, - self.__active_grid.grid_size_mm.x, - self.__active_grid.grid_size_mm.y) + self.grid_scan_size_changed.emit( + self.__active_grid.n_x, + self.__active_grid.n_y, + self.__active_grid.grid_size_mm.x, + self.__active_grid.grid_size_mm.y, + ) self.completed_grid_updated.emit() def get_completed_grids(self) -> List[CompletedRasterGridElem]: - return self.__completed_grids \ No newline at end of file + return self.__completed_grids diff --git a/src/aare/gui/scan_logic/rotation_scan_manager.py b/src/aare/gui/scan_logic/rotation_scan_manager.py index 50a08be2..2d9c41d7 100644 --- a/src/aare/gui/scan_logic/rotation_scan_manager.py +++ b/src/aare/gui/scan_logic/rotation_scan_manager.py @@ -1,7 +1,6 @@ +from aarecommon.config.beamline import get_jfjoch_url, mx_beamline +from aarecommon.models.rotation_scan import CompletedRotationScan from PySide6.QtCore import QObject, Signal, Slot -from aare.common.beamline import mx_beamline, get_jfjoch_url - -from aare.common.rotation_scan import CompletedRotationScan class RotationScanManager(QObject): @@ -16,5 +15,5 @@ class RotationScanManager(QObject): @Slot(CompletedRotationScan) def scan_completed(self, r: CompletedRotationScan): if r.result.file_prefix is not None: - #self.file_ready.emit(r.result.file_prefix, 0) + # self.file_ready.emit(r.result.file_prefix, 0) self.file_ready.emit(self.__detector_url, 0) diff --git a/src/aare/gui/scan_logic/sample_mount_logic.py b/src/aare/gui/scan_logic/sample_mount_logic.py index 60bdccb3..5b077f02 100644 --- a/src/aare/gui/scan_logic/sample_mount_logic.py +++ b/src/aare/gui/scan_logic/sample_mount_logic.py @@ -1,6 +1,5 @@ -from PySide6.QtCore import QObject, Slot, Signal - -from aare.common.models import DAQStatusModel, SampleShortInfo +from aarecommon.models.models import DAQStatusModel, SampleShortInfo +from PySide6.QtCore import QObject, Signal, Slot class SampleMountLogic(QObject): diff --git a/src/aare/gui/threads/camera_thread.py b/src/aare/gui/threads/camera_thread.py index 7cc709cd..cb92b1af 100644 --- a/src/aare/gui/threads/camera_thread.py +++ b/src/aare/gui/threads/camera_thread.py @@ -4,11 +4,11 @@ import time import cv2 import numpy as np import zmq +from aarecommon.math.autofocus import focus_measure_edges +from aarecommon.models.models import DAQStatusModel from PySide6.QtCore import QThread, Signal, Slot from PySide6.QtGui import QImage, QPixmap -from aare.common.autofocus_tools import focus_measure_edges -from aare.common.models import DAQStatusModel class SampleCameraThread(QThread): # Define a signal to communicate messages from the thread to the main GUI @@ -106,7 +106,10 @@ class SampleCameraThread(QThread): encoded = np.frombuffer(data, dtype=np.uint8) bgr = cv2.imdecode(encoded, cv2.IMREAD_COLOR) if bgr is None: - self.__set_camera_available(False, "Sample camera feed unavailable: failed to decode JPEG frame") + self.__set_camera_available( + False, + "Sample camera feed unavailable: failed to decode JPEG frame", + ) continue rgb = cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB) elif "shape" in header: @@ -114,26 +117,41 @@ class SampleCameraThread(QThread): raw = np.frombuffer(data, np.uint8).reshape((h, w)) rgb = cv2.cvtColor(raw, cv2.COLOR_BAYER_GB2RGB) else: - self.__set_camera_available(False, f"Sample camera feed unavailable: unsupported frame header {header}") + self.__set_camera_available( + False, + f"Sample camera feed unavailable: unsupported frame header {header}", + ) continue self.__set_camera_available(True) if self.__measure_focus: gray = cv2.cvtColor(rgb, cv2.COLOR_RGB2GRAY) - if self.__focus_mask is None or self.__focus_mask.shape != gray.shape: + if ( + self.__focus_mask is None + or self.__focus_mask.shape != gray.shape + ): height, width = gray.shape y, x = np.ogrid[:height, :width] self.__focus_mask = (x - self.__beam_x) ** 2 + ( - y - self.__beam_y) ** 2 <= self.__radius ** 2 + y - self.__beam_y + ) ** 2 <= self.__radius**2 sharpness = focus_measure_edges(gray, self.__focus_mask) self.focus_measure.emit(sharpness) - qimage = QImage(rgb.data, rgb.shape[1], rgb.shape[0], QImage.Format.Format_RGB888).copy() + qimage = QImage( + rgb.data, + rgb.shape[1], + rgb.shape[0], + QImage.Format.Format_RGB888, + ).copy() self.camera_image.emit(QPixmap.fromImage(qimage)) else: - self.__set_camera_available(False, "Sample camera feed unavailable: no frame header in zmq stream") + self.__set_camera_available( + False, + "Sample camera feed unavailable: no frame header in zmq stream", + ) except zmq.Again: # Timeout occurred now = time.perf_counter() elapsed = now - self.__fps_window_start @@ -147,10 +165,14 @@ class SampleCameraThread(QThread): self.__fps_frame_count = 0 if no_frames_long: - self.__set_camera_available(False, "Sample camera feed unavailable") + self.__set_camera_available( + False, "Sample camera feed unavailable" + ) continue # Check self.running again except Exception as e: - self.__set_camera_available(False, f"Sample camera feed unavailable: {e}") + self.__set_camera_available( + False, f"Sample camera feed unavailable: {e}" + ) self.running = False def stop(self): diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py index eb59c86a..bfe5d8a3 100644 --- a/src/aare/gui/threads/daq_worker.py +++ b/src/aare/gui/threads/daq_worker.py @@ -1,41 +1,44 @@ import copy +import json +import logging import os import random import re import time -import json -import logging -from datetime import datetime -from dataclasses import dataclass from collections import deque -from typing import cast, Literal +from dataclasses import dataclass +from datetime import datetime +from typing import Literal, cast -from PySide6.QtCore import Signal, QUrl, Slot, QTimer, QObject, QByteArray -from PySide6.QtNetwork import QNetworkAccessManager, QNetworkRequest, QNetworkReply, QSslError -from jfjoch_client import ScanResult, ScanResultImagesInner - -from aare.common.auth_models import BatonStatus -from aare.common.coordinate import SmargonCoordinate, AerotechCoordinate -from aare.common.error_codes import export_error_codes, AareErrorCode, AuthErrorCode -from aare.common.models import ( - DAQStatusModel, - SampleShortInfoList, - SampleShortInfo, - SampleCameraSettings, - AutofocusSettings, - SimpleScanParameters, - FluorescenceSpectrumParameterModel, - FluorescenceSpectrumOutputModel, - OpenGuiSessionInfo, -) -from aare.common.automation_models import ( +from aarecommon.config.logger import setup_logger +from aarecommon.errors.codes import AareErrorCode, AuthErrorCode, export_error_codes +from aarecommon.math.coordinate import AerotechCoordinate, SmargonCoordinate +from aarecommon.models.auth import BatonStatus +from aarecommon.models.automation import ( AutomationProgress, LogEvent, StepState, StepStatus, WorkflowStateKind, ) -from aare.common.recurrence_watcher import ( +from aarecommon.models.models import ( + AutofocusSettings, + DAQStatusModel, + FluorescenceSpectrumOutputModel, + FluorescenceSpectrumParameterModel, + OpenGuiSessionInfo, + SampleCameraSettings, + SampleShortInfo, + SampleShortInfoList, + SimpleScanParameters, +) +from aarecommon.models.raster_grid import ( + CompletedRasterGrid, + CompletedRasterGridElem, + RasterGridRequest, +) +from aarecommon.models.rotation_scan import CompletedRotationScan, RotationScanRequest +from aarecommon.recurrence_watcher import ( DEFAULT_WATCHERS, WatcherTrip, create_default_watchers, @@ -43,10 +46,14 @@ from aare.common.recurrence_watcher import ( redis_key_to_env_var, resolve_exception_class, ) -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 +from jfjoch_client import ScanResult, ScanResultImagesInner +from PySide6.QtCore import QByteArray, QObject, QTimer, QUrl, Signal, Slot +from PySide6.QtNetwork import ( + QNetworkAccessManager, + QNetworkReply, + QNetworkRequest, + QSslError, +) logger = setup_logger("aareGUI") @@ -56,10 +63,12 @@ SPREADHSEET_FREQUENCY = 25 # Every 5 seconds # triggered an action while the beamline/robot was busy). These are blocking but # benign, so they are surfaced quietly (log + status bar) rather than as a modal # pop-up, regardless of the exception's critical flag. -QUIET_OPERATION_ERROR_CODES = frozenset({ - AareErrorCode.BEAMLINE_BUSY_EXCEPTION.value, - AareErrorCode.TELL_COMMAND_WHILE_BUSY_EXCEPTION.value, -}) +QUIET_OPERATION_ERROR_CODES = frozenset( + { + AareErrorCode.BEAMLINE_BUSY_EXCEPTION.value, + AareErrorCode.TELL_COMMAND_WHILE_BUSY_EXCEPTION.value, + } +) @dataclass(frozen=True) @@ -159,7 +168,10 @@ class DAQWorker(QObject): self._device_error_log_min_interval_s = 10.0 self._last_device_error_log_ts: dict[str, float] = {"tell": 0.0, "smargon": 0.0} - self._last_device_error_log_key: dict[str, str | None] = {"tell": None, "smargon": None} + self._last_device_error_log_key: dict[str, str | None] = { + "tell": None, + "smargon": None, + } self._last_status_error = None self._smargon_error_active = False @@ -244,7 +256,9 @@ class DAQWorker(QObject): 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: + 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. @@ -338,7 +352,7 @@ class DAQWorker(QObject): Send an asynchronous HTTP GET request to the /status endpoint. """ if self.__base_url is None: - return + return now = time.monotonic() if now - self._last_status_request_ts < self._status_request_min_interval: @@ -361,7 +375,6 @@ class DAQWorker(QObject): if self_signed: reply.ignoreSslErrors(self_signed) - @staticmethod def handle_response(reply: QNetworkReply): """ @@ -386,14 +399,14 @@ class DAQWorker(QObject): raise RuntimeError(reply.errorString()) def _compose_device_status_message( - self, - *, - tell_conn: bool, - smargon_conn: bool, - aerotech_conn: bool, - tell_changed: bool, - smargon_changed: bool, - aerotech_changed: bool, + self, + *, + tell_conn: bool, + smargon_conn: bool, + aerotech_conn: bool, + tell_changed: bool, + smargon_changed: bool, + aerotech_changed: bool, ) -> tuple[str | None, str | None, bool]: disconnected: list[str] = [] restored: list[str] = [] @@ -455,7 +468,9 @@ class DAQWorker(QObject): parsed_response = DAQStatusModel.model_validate_json(response_data) self.update.emit(parsed_response) - pss_alarm = bool(getattr(getattr(parsed_response, "bl", None), "pss_alarm", False)) + pss_alarm = bool( + getattr(getattr(parsed_response, "bl", None), "pss_alarm", False) + ) if pss_alarm != self._last_pss_alarm: self._last_pss_alarm = pss_alarm self.pss_alarm_changed.emit(pss_alarm) @@ -484,11 +499,17 @@ class DAQWorker(QObject): 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() + 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: + 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 @@ -499,16 +520,16 @@ class DAQWorker(QObject): return tell_changed = ( - 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 != 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 != 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 @@ -533,19 +554,27 @@ class DAQWorker(QObject): 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) + self._log_device_error_throttled( + device="aerotech", message=aerotech_err + ) except Exception as e: - status=None + status = None err_msg = str(e) try: - status = reply.attribute(QNetworkRequest.Attribute.HttpStatusCodeAttribute) + status = reply.attribute( + QNetworkRequest.Attribute.HttpStatusCodeAttribute + ) 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") or body_json.get("detail") or err_msg + err_msg = ( + body_json.get("message") + or body_json.get("detail") + or err_msg + ) if body_json.get("code") == "AEROTECH_UNAVAILABLE": extra = body_json.get("extra") or {} if isinstance(extra, dict): @@ -627,15 +656,19 @@ class DAQWorker(QObject): net_err_name = getattr(net_err, "name", None) net_err_value = getattr(net_err, "value", None) - self._set_last_error_payload({ - "url": url, - "http_status": int(status) if status is not None else None, - "network_error": net_err_name or str(net_err), - "network_error_value": int(net_err_value) if isinstance(net_err_value, int) else None, - "error_string": str(reply.errorString()), - "body_raw": raw_body, - "body_json": body_json, - }) + self._set_last_error_payload( + { + "url": url, + "http_status": int(status) if status is not None else None, + "network_error": net_err_name or str(net_err), + "network_error_value": int(net_err_value) + if isinstance(net_err_value, int) + else None, + "error_string": str(reply.errorString()), + "body_raw": raw_body, + "body_json": body_json, + } + ) if self._is_auth_error_code(error_info.code) or status == 401: now = time.monotonic() @@ -653,7 +686,9 @@ class DAQWorker(QObject): else: logger.error(f"{error_info.message}") title = self._operation_error_title(error_info.exception_class) - self.operation_failed.emit(title, error_info.message, error_info.critical) + self.operation_failed.emit( + title, error_info.message, error_info.critical + ) reply.deleteLater() @@ -697,7 +732,9 @@ class DAQWorker(QObject): logger.error(f"Hardware metadata resync failed: {e}") self.http_error.emit(str(e)) - def _handle_recovery_action_response(self, reply: QNetworkReply, default_message: str): + def _handle_recovery_action_response( + self, reply: QNetworkReply, default_message: str + ): try: response_data = self.handle_response(reply) payload = json.loads(response_data) if response_data else {} @@ -868,7 +905,9 @@ class DAQWorker(QObject): body = json.dumps({"confirmation_code": confirmation_code}) reply = self.__net_manager.post(request, QByteArray(body.encode("utf-8"))) reply.finished.connect( - lambda: self._handle_recovery_action_response(reply, "Beamline busy flag cleared.") + lambda: self._handle_recovery_action_response( + reply, "Beamline busy flag cleared." + ) ) @Slot(str) @@ -883,7 +922,9 @@ class DAQWorker(QObject): body = json.dumps({"confirmation_code": confirmation_code}) reply = self.__net_manager.post(request, QByteArray(body.encode("utf-8"))) reply.finished.connect( - lambda: self._handle_recovery_action_response(reply, "Beamline session taken over.") + lambda: self._handle_recovery_action_response( + reply, "Beamline session taken over." + ) ) @Slot(str) @@ -898,7 +939,9 @@ class DAQWorker(QObject): body = json.dumps({"confirmation_code": confirmation_code}) reply = self.__net_manager.post(request, QByteArray(body.encode("utf-8"))) reply.finished.connect( - lambda: self._handle_recovery_action_response(reply, "Beamline recovered to Maintenance.") + lambda: self._handle_recovery_action_response( + reply, "Beamline recovered to Maintenance." + ) ) @Slot(str) @@ -913,7 +956,9 @@ class DAQWorker(QObject): body = json.dumps({"confirmation_code": confirmation_code}) reply = self.__net_manager.post(request, QByteArray(body.encode("utf-8"))) reply.finished.connect( - lambda: self._handle_recovery_action_response(reply, "Recovery unmount completed.") + lambda: self._handle_recovery_action_response( + reply, "Recovery unmount completed." + ) ) @Slot(str) @@ -933,6 +978,7 @@ class DAQWorker(QObject): try: response_data = self.handle_response(reply) import json + arr = json.loads(response_data) if response_data else [] if not isinstance(arr, list): raise RuntimeError("Invalid all_pgroups payload") @@ -976,8 +1022,13 @@ class DAQWorker(QObject): if self._is_detector_state_failure_message(err_msg): self._emit_detector_message(f"JFJoch: {err_msg}", is_error=True) else: - code = body_json.get("code", "") if isinstance(body_json, dict) else "" - if code == AareErrorCode.JF_JOCH_COMMUNICATION_ERROR or status == 503: + code = ( + body_json.get("code", "") if isinstance(body_json, dict) else "" + ) + if ( + code == AareErrorCode.JF_JOCH_COMMUNICATION_ERROR + or status == 503 + ): self._emit_detector_message(f"JFJoch: {err_msg}", is_error=True) if self._is_critical_detector_failure(status, body_json, err_msg): @@ -986,10 +1037,17 @@ class DAQWorker(QObject): "Please call your local contact.\n\n" f"Details: {err_msg}" ) - self._emit_detector_message(f"Detector error during manual collection: {err_msg}", is_error=True) + self._emit_detector_message( + f"Detector error during manual collection: {err_msg}", + is_error=True, + ) self._emit_manual_collection_critical_failure(critical_msg) - short_msg = err_msg.split("input':", 1)[0].strip() if "input':" in err_msg else err_msg + short_msg = ( + err_msg.split("input':", 1)[0].strip() + if "input':" in err_msg + else err_msg + ) logger.error(f"Rotation scan failed: {short_msg}") self.http_error.emit(short_msg) reply.deleteLater() @@ -1030,8 +1088,13 @@ class DAQWorker(QObject): if self._is_detector_state_failure_message(err_msg): self._emit_detector_message(f"JFJoch: {err_msg}", is_error=True) else: - code = body_json.get("code", "") if isinstance(body_json, dict) else "" - if code == AareErrorCode.JF_JOCH_COMMUNICATION_ERROR or status == 503: + code = ( + body_json.get("code", "") if isinstance(body_json, dict) else "" + ) + if ( + code == AareErrorCode.JF_JOCH_COMMUNICATION_ERROR + or status == 503 + ): self._emit_detector_message(f"JFJoch: {err_msg}", is_error=True) if self._is_critical_detector_failure(status, body_json, err_msg): @@ -1040,7 +1103,10 @@ class DAQWorker(QObject): "Please call your local contact.\n\n" f"Details: {err_msg}" ) - self._emit_detector_message(f"Detector error during manual collection: {err_msg}", is_error=True) + self._emit_detector_message( + f"Detector error during manual collection: {err_msg}", + is_error=True, + ) self._emit_manual_collection_critical_failure(critical_msg) logger.error(f"Raster scan failed: {err_msg}") @@ -1074,24 +1140,28 @@ class DAQWorker(QObject): images = [] for i in range(image_number): - images.append(ScanResultImagesInner( - number=i, - efficiency=1.0, - bkg = random.gauss(3.0, 0.1), - spots= random.randint(0, 250), - index= random.randint(0, 1), - b= random.uniform(15.0, 80.0) - )) + images.append( + ScanResultImagesInner( + number=i, + efficiency=1.0, + bkg=random.gauss(3.0, 0.1), + spots=random.randint(0, 250), + index=random.randint(0, 1), + b=random.uniform(15.0, 80.0), + ) + ) raster_elem = CompletedRasterGridElem( - request=new_copy, - result=ScanResult(file_prefix=r.file_prefix, images=images), - centre_of_mass = None, - ) + 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 - request = QNetworkRequest(QUrl(f"{self.__base_url}/scan/raster?auto_center=false")) + request = QNetworkRequest( + QUrl(f"{self.__base_url}/scan/raster?auto_center=false") + ) request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8")) request.setRawHeader(b"Content-Type", b"application/json") body = r.model_dump_json() @@ -1114,31 +1184,34 @@ class DAQWorker(QObject): images = [] for i in range(image_number): - images.append(ScanResultImagesInner( - number=i, - efficiency=1.0, - bkg=random.gauss(3.0, 0.1), - spots=random.randint(0, 250), - index=random.randint(0, 1), - b=random.uniform(15.0, 80.0) - )) + images.append( + ScanResultImagesInner( + number=i, + efficiency=1.0, + bkg=random.gauss(3.0, 0.1), + spots=random.randint(0, 250), + index=random.randint(0, 1), + b=random.uniform(15.0, 80.0), + ) + ) raster_elem = CompletedRasterGridElem( - request=new_copy, - result=ScanResult(file_prefix=r.file_prefix, images=images), - centre_of_mass = None, - ) + 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 - request = QNetworkRequest(QUrl(f"{self.__base_url}/scan/raster?auto_center=true")) + request = QNetworkRequest( + QUrl(f"{self.__base_url}/scan/raster?auto_center=true") + ) request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8")) request.setRawHeader(b"Content-Type", b"application/json") body = r.model_dump_json() reply = self.__net_manager.post(request, QByteArray(body.encode("utf-8"))) reply.finished.connect(lambda: self.handle_raster_scan_response(reply)) - @Slot() def load_spreadsheet(self): if self.__base_url is None: @@ -1183,7 +1256,9 @@ class DAQWorker(QObject): for watcher in self._recurrence_watchers: watcher.reset() - def _observe_recurrence(self, exception_class_name: str | None) -> WatcherTrip | None: + def _observe_recurrence( + self, exception_class_name: str | None + ) -> WatcherTrip | None: exception_class = resolve_exception_class(exception_class_name) for watcher in self._recurrence_watchers: trip = watcher.maybe_trip(exception_class) @@ -1198,7 +1273,9 @@ class DAQWorker(QObject): self.automation_critical_failure.emit(message) @staticmethod - def _extract_reply_error_details(reply: QNetworkReply) -> tuple[int | None, str, dict | None]: + def _extract_reply_error_details( + reply: QNetworkReply, + ) -> tuple[int | None, str, dict | None]: status = reply.attribute(QNetworkRequest.Attribute.HttpStatusCodeAttribute) err_msg = reply.errorString() body_json = None @@ -1233,19 +1310,23 @@ class DAQWorker(QObject): @staticmethod def _classify_reply_error( - status: int | None, - err_msg: str, - body_json: dict | None, + status: int | None, + err_msg: str, + body_json: dict | None, ) -> ErrorInfo: code = body_json.get("code") if isinstance(body_json, dict) else None - exception_class = body_json.get("exception_class") if isinstance(body_json, dict) else None + exception_class = ( + body_json.get("exception_class") if isinstance(body_json, dict) else None + ) context = body_json.get("context") if isinstance(body_json, dict) else {} if not isinstance(context, dict): context = {} return ErrorInfo( critical=DAQWorker._is_critical(body_json), code=str(code) if code is not None else None, - exception_class=str(exception_class) if exception_class is not None else None, + exception_class=str(exception_class) + if exception_class is not None + else None, message=str(err_msg or ""), context=context, ) @@ -1270,9 +1351,9 @@ class DAQWorker(QObject): @staticmethod def _is_critical_detector_failure( - status: int | None, - body_json: dict | None = None, - err_msg: str | None = None, + status: int | None, + body_json: dict | None = None, + err_msg: str | None = None, ) -> bool: try: status_int = int(status) if status is not None else None @@ -1288,16 +1369,22 @@ class DAQWorker(QObject): if status_int == 500: return True - if code == AareErrorCode.JF_JOCH_COMMUNICATION_ERROR and status_int is not None and 500 <= status_int < 600: + if ( + code == AareErrorCode.JF_JOCH_COMMUNICATION_ERROR + and status_int is not None + and 500 <= status_int < 600 + ): return True - if "daq state error" in text or "must be idle to start measurement" in text or "must be idle" in text: + if ( + "daq state error" in text + or "must be idle to start measurement" in text + or "must be idle" in text + ): return True return False - - def handle_auto_scan_response(self, reply, sample_id: int): if reply.error() == QNetworkReply.NetworkError.NoError: resp = reply.readAll().data().decode("utf-8") @@ -1310,7 +1397,9 @@ class DAQWorker(QObject): watcher_trip = self._observe_recurrence(error_info.exception_class) if self._is_detector_state_failure_message(error_info.message): - self._emit_detector_message(f"JFJoch: {error_info.message}", is_error=True) + self._emit_detector_message( + f"JFJoch: {error_info.message}", is_error=True + ) if self._is_auth_error_code(error_info.code) or status == 401: logger.error(f"Error in auto scan: {error_info.message}") @@ -1411,17 +1500,21 @@ class DAQWorker(QObject): logger.info("POST /local_contact/resync/detector_metadata") return - request = QNetworkRequest(QUrl(f"{self.__base_url}/local_contact/resync/detector_metadata")) + request = QNetworkRequest( + QUrl(f"{self.__base_url}/local_contact/resync/detector_metadata") + ) 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_detector_metadata_resync_response(reply)) + reply.finished.connect( + lambda: self._handle_detector_metadata_resync_response(reply) + ) @Slot(float) def anneal(self, time_s: float): self.generic_post(f"beamline/anneal?time_s={time_s:.1f}") - #TODO combine dry and park and dry + # TODO combine dry and park and dry @Slot() def park_and_dry(self): self.generic_post("tell/park_and_dry") @@ -1435,26 +1528,26 @@ class DAQWorker(QObject): self.generic_post("tell/toggle_blower") def _disable_local_contact_metadata_polling(self, message: str) -> None: - if self._local_contact_metadata_poll_enabled is False and self._local_contact_metadata_error == message: + if ( + self._local_contact_metadata_poll_enabled is False + and self._local_contact_metadata_error == message + ): return self._local_contact_metadata_poll_enabled = False self._local_contact_metadata_error = message logger.error(f"Disabling Local Contact metadata polling: {message}") self.local_contact_transfer_error.emit(message) - def _build_local_contact_error_message(self, context: str, reply: QNetworkReply, exc: Exception) -> str: + def _build_local_contact_error_message( + self, context: str, reply: QNetworkReply, exc: Exception + ) -> str: status = reply.attribute(QNetworkRequest.Attribute.HttpStatusCodeAttribute) try: url = reply.request().url().toString() except Exception: url = "unknown-url" - return ( - f"{context}\n\n" - f"URL: {url}\n" - f"HTTP status: {status}\n" - f"Error: {exc}" - ) + return f"{context}\n\nURL: {url}\nHTTP status: {status}\nError: {exc}" @Slot() def load_local_contact_simulation_state(self): @@ -1463,14 +1556,24 @@ class DAQWorker(QObject): if self.__base_url is None: self.local_contact_simulation_state_loaded.emit( - {"bec": False, "detector": False, "tell": False, "aerotech": False, "smargon": False} + { + "bec": False, + "detector": False, + "tell": False, + "aerotech": False, + "smargon": False, + } ) return - request = QNetworkRequest(QUrl(f"{self.__base_url}/local_contact/simulation_state")) + request = QNetworkRequest( + QUrl(f"{self.__base_url}/local_contact/simulation_state") + ) request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8")) reply = self.__net_manager.get(request) - reply.finished.connect(lambda: self._handle_local_contact_simulation_state_response(reply)) + reply.finished.connect( + lambda: self._handle_local_contact_simulation_state_response(reply) + ) def _handle_local_contact_simulation_state_response(self, reply: QNetworkReply): try: @@ -1499,7 +1602,9 @@ class DAQWorker(QObject): request = QNetworkRequest(QUrl(f"{self.__base_url}/local_contact/device_state")) request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8")) reply = self.__net_manager.get(request) - reply.finished.connect(lambda: self._handle_local_contact_device_state_response(reply)) + reply.finished.connect( + lambda: self._handle_local_contact_device_state_response(reply) + ) def _handle_local_contact_device_state_response(self, reply: QNetworkReply): try: @@ -1548,17 +1653,21 @@ class DAQWorker(QObject): @Slot() def load_local_contact_config(self): if self.__base_url is None: - self.local_contact_config_loaded.emit({ - "mount_to_center_sleep_s": 0.0, - "line_scan_loop_face_y_padding_fraction_each_side": 0.5, - "line_scan_loop_all_y_padding_fraction_each_side": 0.5, - }) + self.local_contact_config_loaded.emit( + { + "mount_to_center_sleep_s": 0.0, + "line_scan_loop_face_y_padding_fraction_each_side": 0.5, + "line_scan_loop_all_y_padding_fraction_each_side": 0.5, + } + ) return request = QNetworkRequest(QUrl(f"{self.__base_url}/local_contact/config")) request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8")) reply = self.__net_manager.get(request) - reply.finished.connect(lambda: self._handle_local_contact_config_response(reply)) + reply.finished.connect( + lambda: self._handle_local_contact_config_response(reply) + ) def _handle_local_contact_config_response(self, reply: QNetworkReply): try: @@ -1586,7 +1695,9 @@ class DAQWorker(QObject): request.setRawHeader(b"Content-Type", b"application/json") body = QByteArray(json.dumps(payload).encode("utf-8")) reply = self.__net_manager.put(request, body) - reply.finished.connect(lambda: self._handle_set_local_contact_config_response(reply)) + reply.finished.connect( + lambda: self._handle_set_local_contact_config_response(reply) + ) def _handle_set_local_contact_config_response(self, reply: QNetworkReply): try: @@ -1607,7 +1718,9 @@ class DAQWorker(QObject): @Slot(str, bool) def set_local_contact_simulation(self, device: str, enabled: bool): - self.generic_post(f"local_contact/simulate/{device}?enabled={str(enabled).lower()}") + self.generic_post( + f"local_contact/simulate/{device}?enabled={str(enabled).lower()}" + ) @Slot(str) def restart_local_contact_device(self, device: str): @@ -1665,7 +1778,9 @@ class DAQWorker(QObject): @Slot(str) def bec_reinitialise_planner_and_position_devices(self, method: str = "auto"): - self.generic_post(f"bec/reinitialise_planner_and_position_devices?method={method}") + self.generic_post( + f"bec/reinitialise_planner_and_position_devices?method={method}" + ) @Slot() def bec_save_current_bs_pos(self): @@ -1689,7 +1804,9 @@ class DAQWorker(QObject): @Slot() def initialise_aerotech(self): - logger.info("initisalisation does not initisalise aareSCAN but runs homing script") + logger.info( + "initisalisation does not initisalise aareSCAN but runs homing script" + ) self.generic_post("aerotech/initialize") @Slot() @@ -1747,6 +1864,7 @@ class DAQWorker(QObject): raise RuntimeError(reply.errorString()) payload = reply.readAll().data().decode("utf-8") or "{}" import json + data = json.loads(payload) self.face_detection_result.emit(data) except Exception as e: @@ -1774,13 +1892,17 @@ class DAQWorker(QObject): status = None if reply is not None: try: - status = reply.attribute(QNetworkRequest.Attribute.HttpStatusCodeAttribute) + status = reply.attribute( + QNetworkRequest.Attribute.HttpStatusCodeAttribute + ) except Exception: status = None if status == 403: self._face_detection_stream_blocked_403 = True - logger.info("Face detection SSE access denied; waiting for access change before retrying.") + logger.info( + "Face detection SSE access denied; waiting for access change before retrying." + ) return if self.__base_url is not None: @@ -1850,7 +1972,9 @@ class DAQWorker(QObject): finished=bool(progress_payload.get("finished", False)), success=progress_payload.get("success"), samples_in_queue=int(progress_payload.get("samples_in_queue", 0) or 0), - avg_time_per_sample=float(progress_payload.get("avg_time_per_sample", 0.0) or 0.0), + avg_time_per_sample=float( + progress_payload.get("avg_time_per_sample", 0.0) or 0.0 + ), current_sample_name=str(progress_payload.get("current_sample_name") or ""), ) @@ -1912,7 +2036,9 @@ class DAQWorker(QObject): def _process_automation_progress_buffer(self) -> None: while "\n\n" in self._automation_progress_buffer: - event_data, self._automation_progress_buffer = self._automation_progress_buffer.split("\n\n", 1) + event_data, self._automation_progress_buffer = ( + self._automation_progress_buffer.split("\n\n", 1) + ) data_lines: list[str] = [] for line in event_data.splitlines(): @@ -1944,13 +2070,17 @@ class DAQWorker(QObject): status = None if reply is not None: try: - status = reply.attribute(QNetworkRequest.Attribute.HttpStatusCodeAttribute) + status = reply.attribute( + QNetworkRequest.Attribute.HttpStatusCodeAttribute + ) except Exception: status = None if status == 403: self._automation_progress_stream_blocked_403 = True - logger.info("Automation progress SSE access denied; waiting for access change before retrying.") + logger.info( + "Automation progress SSE access denied; waiting for access change before retrying." + ) return if self.__base_url is not None: @@ -2005,7 +2135,11 @@ class DAQWorker(QObject): logger.info(f"POST /face_detection/run?steps={steps}&step_size={step_size}") return - request = QNetworkRequest(QUrl(f"{self.__base_url}/face_detection/run?steps={steps}&step_size={step_size}")) + request = QNetworkRequest( + QUrl( + f"{self.__base_url}/face_detection/run?steps={steps}&step_size={step_size}" + ) + ) 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"")) @@ -2027,13 +2161,20 @@ class DAQWorker(QObject): try: response_data = self.handle_response(reply) import json + data = json.loads(response_data) if response_data else [] if emit_status: # fetch status and bkg in parallel (simple sequential here) - status_req = QNetworkRequest(QUrl(f"{self.__base_url}/fluorimeter/status")) - status_req.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8")) + status_req = QNetworkRequest( + QUrl(f"{self.__base_url}/fluorimeter/status") + ) + status_req.setRawHeader( + b"Authorization", f"Bearer {self.__token}".encode("utf-8") + ) status_reply = self.__net_manager.get(status_req) - status_reply.finished.connect(lambda: self._handle_fluorimeter_status_and_emit(data, status_reply)) + status_reply.finished.connect( + lambda: self._handle_fluorimeter_status_and_emit(data, status_reply) + ) else: self.fluorimeter_update.emit(data, [], -1) except Exception as e: @@ -2045,9 +2186,13 @@ class DAQWorker(QObject): s_payload = self.handle_response(status_reply) s = int(s_payload) if s_payload not in ("", "null") else -1 b_req = QNetworkRequest(QUrl(f"{self.__base_url}/fluorimeter/background")) - b_req.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8")) + b_req.setRawHeader( + b"Authorization", f"Bearer {self.__token}".encode("utf-8") + ) b_reply = self.__net_manager.get(b_req) - b_reply.finished.connect(lambda: self._emit_fluorimeter_with_bkg(data, s, b_reply)) + b_reply.finished.connect( + lambda: self._emit_fluorimeter_with_bkg(data, s, b_reply) + ) except Exception as e: logger.error(f"Fluorimeter status error: {e}") self.http_error.emit(str(e)) @@ -2056,6 +2201,7 @@ class DAQWorker(QObject): try: bkg_json = self.handle_response(b_reply) import json + bkg = json.loads(bkg_json) if bkg_json else [] self.fluorimeter_update.emit(data, bkg, s) except Exception as e: @@ -2076,7 +2222,9 @@ class DAQWorker(QObject): def _handle_fluorimeter_spectrum(self, reply: QNetworkReply): try: response_data = self.handle_response(reply) - parsed_response = FluorescenceSpectrumOutputModel.model_validate_json(response_data) + parsed_response = FluorescenceSpectrumOutputModel.model_validate_json( + response_data + ) self.fluorimeter_spectrum_update.emit(parsed_response) except Exception as e: logger.error(f"Exception from fluorimeter spectrum: {e}") @@ -2089,7 +2237,9 @@ class DAQWorker(QObject): req = QNetworkRequest(QUrl(f"{self.__base_url}/fluorimeter/data")) req.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8")) reply = self.__net_manager.get(req) - reply.finished.connect(lambda: self._handle_fluorimeter_data(reply, emit_status=True)) + reply.finished.connect( + lambda: self._handle_fluorimeter_data(reply, emit_status=True) + ) # Optional: SSE listener for live updates def start_fluorimeter_stream(self): @@ -2186,7 +2336,9 @@ class DAQWorker(QObject): @Slot(str, str) def send_screenshot_db(self, filename: str = "", message: str = ""): if self.__base_url is None: - logger.info(f"POST /samcam/send_screenshot_db?filename={filename}&message={message}") + logger.info( + f"POST /samcam/send_screenshot_db?filename={filename}&message={message}" + ) return from urllib.parse import quote @@ -2239,13 +2391,20 @@ class DAQWorker(QObject): if self._baton_timeout_timer.isActive(): self._baton_timeout_timer.stop() - if (status.incoming_request and - (self._last_baton_status is None or - not self._last_baton_status.incoming_request)): - self.baton_incoming_request.emit({ - "requester": status.pending_request.requester_username if status.pending_request else "Unknown", - "timeout": status.pending_request.timeout_seconds if status.pending_request else 30 - }) + 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) @@ -2280,7 +2439,7 @@ class DAQWorker(QObject): elif result.get("pending"): self.status_message.emit( f"Request sent - waiting for response ({result.get('timeout_seconds', 30)}s timeout)", - False + False, ) elif result.get("error"): self.status_message.emit(result.get("message", "Request failed"), True) @@ -2300,7 +2459,9 @@ class DAQWorker(QObject): logger.info(f"POST /baton/respond?accept={accept}") return - request = QNetworkRequest(QUrl(f"{self.__base_url}/baton/respond?accept={str(accept).lower()}")) + 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"")) @@ -2369,12 +2530,16 @@ class DAQWorker(QObject): - clear the active beamline session if this GUI owns it """ try: - if hasattr(self, "_baton_timeout_timer") and self._baton_timeout_timer is not None: + if ( + hasattr(self, "_baton_timeout_timer") + and self._baton_timeout_timer is not None + ): self._baton_timeout_timer.stop() self.end_session() from PySide6.QtCore import QEventLoop, QTimer + loop = QEventLoop() QTimer.singleShot(500, loop.quit) loop.exec() @@ -2407,16 +2572,22 @@ class DAQWorker(QObject): @Slot(int, int) def request_gui_close(self, session_id: int, grace_seconds: int = 60): if self.__base_url is None: - logger.info(f"POST /admin/gui_sessions/{session_id}/request_close?grace_seconds={grace_seconds}") + logger.info( + f"POST /admin/gui_sessions/{session_id}/request_close?grace_seconds={grace_seconds}" + ) return request = QNetworkRequest( - QUrl(f"{self.__base_url}/admin/gui_sessions/{session_id}/request_close?grace_seconds={grace_seconds}") + QUrl( + f"{self.__base_url}/admin/gui_sessions/{session_id}/request_close?grace_seconds={grace_seconds}" + ) ) 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_gui_session_mutation_response(reply)) + reply.finished.connect( + lambda: self._handle_gui_session_mutation_response(reply) + ) @Slot(int) def force_remove_gui_session(self, session_id: int): @@ -2424,10 +2595,14 @@ class DAQWorker(QObject): logger.info(f"DELETE /admin/gui_sessions/{session_id}") return - request = QNetworkRequest(QUrl(f"{self.__base_url}/admin/gui_sessions/{session_id}")) + request = QNetworkRequest( + QUrl(f"{self.__base_url}/admin/gui_sessions/{session_id}") + ) request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8")) reply = self.__net_manager.deleteResource(request) - reply.finished.connect(lambda: self._handle_gui_session_mutation_response(reply)) + reply.finished.connect( + lambda: self._handle_gui_session_mutation_response(reply) + ) def _handle_gui_session_mutation_response(self, reply: QNetworkReply): try: @@ -2448,7 +2623,9 @@ class DAQWorker(QObject): if self.__base_url is None: return - request = QNetworkRequest(QUrl(f"{self.__base_url}/admin/gui_sessions/{session_id}/interaction")) + request = QNetworkRequest( + QUrl(f"{self.__base_url}/admin/gui_sessions/{session_id}/interaction") + ) 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"")) @@ -2460,7 +2637,10 @@ class DAQWorker(QObject): self._cleanup_done = True try: - if hasattr(self, "_baton_timeout_timer") and self._baton_timeout_timer is not None: + if ( + hasattr(self, "_baton_timeout_timer") + and self._baton_timeout_timer is not None + ): self._baton_timeout_timer.stop() except Exception as e: logger.warning(f"Failed to stop _baton_timeout_timer: {e}") @@ -2494,4 +2674,4 @@ class DAQWorker(QObject): self._face_detection_stream_blocked_403 = False self._automation_progress_stream_blocked_403 = False - self._automation_progress_buffer = "" \ No newline at end of file + self._automation_progress_buffer = "" diff --git a/src/aare/gui/threads/jfjoch_viewer.py b/src/aare/gui/threads/jfjoch_viewer.py index fd3455d0..250ff853 100644 --- a/src/aare/gui/threads/jfjoch_viewer.py +++ b/src/aare/gui/threads/jfjoch_viewer.py @@ -1,6 +1,6 @@ +from aarecommon.config.beamline import get_jfjoch_url, mx_beamline from PySide6.QtCore import QObject, Slot from PySide6.QtDBus import QDBusConnection, QDBusInterface -from aare.common.beamline import MXBeamline, mx_beamline, get_jfjoch_url class JFJochDBusClient(QObject): @@ -28,7 +28,6 @@ class JFJochDBusClient(QObject): except Exception as e: print(f"D-Bus not available: {e}.") - def _ensure_interface(self): """Ensure we have a valid interface, creating one if needed.""" if not self.__dbus_available: @@ -44,7 +43,7 @@ class JFJochDBusClient(QObject): "ch.psi.jfjoch_viewer", # Service name "/", # Object path "ch.psi.jfjoch_viewer", # Interface name - self.__session_bus + self.__session_bus, ) if not self.__interface.isValid(): @@ -75,4 +74,3 @@ class JFJochDBusClient(QObject): if self._ensure_interface(): reply = self.__interface.call("LoadFile", self.__detector_url, -1, 1) - diff --git a/src/aare/gui/threads/prediction_subscriber.py b/src/aare/gui/threads/prediction_subscriber.py index 175b66e9..81430ca2 100644 --- a/src/aare/gui/threads/prediction_subscriber.py +++ b/src/aare/gui/threads/prediction_subscriber.py @@ -1,20 +1,18 @@ import json import time +import cv2 import numpy as np import zmq - +from aarecommon.config.logger import setup_logger +from aarecommon.math.autofocus import focus_measure_edges +from aarecommon.models.models import DAQStatusModel from PySide6.QtCore import QThread, Signal, Slot from PySide6.QtGui import QImage, QPixmap -import cv2 - -from aare.common.autofocus_tools import focus_measure_edges -from aare.common.logger_config import setup_logger -from aare.common.models import DAQStatusModel - logger = setup_logger("aareGUI") + class PredictionSubscriber(QThread): prediction = Signal(dict) target_point = Signal(dict) @@ -144,9 +142,9 @@ class PredictionSubscriber(QThread): if self._focus_mask is None or self._focus_mask.shape != gray.shape: height, width = gray.shape y, x = np.ogrid[:height, :width] - self._focus_mask = ( - (x - self._beam_x) ** 2 + (y - self._beam_y) ** 2 <= self._focus_radius ** 2 - ) + self._focus_mask = (x - self._beam_x) ** 2 + ( + y - self._beam_y + ) ** 2 <= self._focus_radius**2 sharpness = focus_measure_edges(gray, self._focus_mask) self.focus_measure.emit(sharpness) @@ -171,7 +169,9 @@ class PredictionSubscriber(QThread): self._fps_frame_count = 0 if no_frames_long: - self._set_camera_available(False, "Sample camera feed unavailable") + self._set_camera_available( + False, "Sample camera feed unavailable" + ) continue if not parts: @@ -200,7 +200,8 @@ class PredictionSubscriber(QThread): header = next( ( - d for d in json_dicts + d + for d in json_dicts if d.get("encoding") == "jpeg" or ("shape" in d and d.get("type") == "uint8") ), @@ -213,14 +214,20 @@ class PredictionSubscriber(QThread): if self._emit_images and header and image_bytes: rgb = self._decode_rgb_image(header, image_bytes) if rgb is None: - self._set_camera_available(False, "Sample camera feed unavailable: failed to decode frame") + self._set_camera_available( + False, + "Sample camera feed unavailable: failed to decode frame", + ) else: self._set_camera_available(True) self._emit_focus_measure_if_enabled(rgb) if self.running: self.image.emit(self._rgb_to_pixmap(rgb)) elif self._emit_images: - self._set_camera_available(False, "Sample camera feed unavailable: no frame header in zmq stream") + self._set_camera_available( + False, + "Sample camera feed unavailable: no frame header in zmq stream", + ) if detections and self.running: self.prediction.emit(detections) @@ -230,7 +237,9 @@ class PredictionSubscriber(QThread): except Exception as e: if self.running: - self._set_camera_available(False, f"Sample camera feed unavailable: {e}") + self._set_camera_available( + False, f"Sample camera feed unavailable: {e}" + ) logger.exception(f"PredictionSubscriber error: {e}") finally: try: @@ -260,4 +269,4 @@ class PredictionSubscriber(QThread): pass if not self.wait(1500): - logger.warning("PredictionSubscriber did not stop within timeout") \ No newline at end of file + logger.warning("PredictionSubscriber did not stop within timeout") diff --git a/src/aare/gui/widgets/alert_banner.py b/src/aare/gui/widgets/alert_banner.py index dd5c8716..3bb27e42 100644 --- a/src/aare/gui/widgets/alert_banner.py +++ b/src/aare/gui/widgets/alert_banner.py @@ -1,8 +1,13 @@ -from PySide6.QtCore import QTimer, Slot, Qt +from aarecommon.config.logger import setup_logger +from PySide6.QtCore import Qt, QTimer, Slot from PySide6.QtGui import QColor -from PySide6.QtWidgets import QFrame, QHBoxLayout, QLabel, QSizePolicy, QGraphicsDropShadowEffect - -from aare.common.logger_config import setup_logger +from PySide6.QtWidgets import ( + QFrame, + QGraphicsDropShadowEffect, + QHBoxLayout, + QLabel, + QSizePolicy, +) logger = setup_logger("aareGUI") @@ -53,7 +58,9 @@ class AlertBanner(QFrame): self.update() @Slot(str, bool) - def show_message(self, msg: str, is_error: bool = True, auto_clear_ms: int | None = None): + 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() @@ -113,7 +120,9 @@ class AlertBanner(QFrame): 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) ⏳" + decorated = ( + f"⏳ {self._countdown_base_message} ({self._countdown_remaining}s) ⏳" + ) else: decorated = f"⏳ {self._countdown_base_message} ⏳" self._label.setText(decorated) @@ -141,4 +150,4 @@ class AlertBanner(QFrame): self._clear_timer.stop() self._set_alert_kind("error") self._label.clear() - self.setVisible(False) \ No newline at end of file + self.setVisible(False) diff --git a/src/aare/gui/widgets/automation_progress.py b/src/aare/gui/widgets/automation_progress.py index c700bf9c..03eaf568 100644 --- a/src/aare/gui/widgets/automation_progress.py +++ b/src/aare/gui/widgets/automation_progress.py @@ -3,11 +3,14 @@ from __future__ import annotations import time from datetime import datetime +from aarecommon.models.automation import ( + AutomationProgress, + StepStatus, + WorkflowStateKind, +) from PySide6.QtCore import QTimer, Slot from PySide6.QtWidgets import QFrame, QHBoxLayout, QLabel, QVBoxLayout, QWidget -from aare.common.automation_models import AutomationProgress, StepStatus, WorkflowStateKind - class CompactAutomationProgressStrip(QFrame): DEFAULT_SAMPLE_ESTIMATE_S = 150.0 @@ -155,7 +158,11 @@ class CompactAutomationProgressStrip(QFrame): self._timer.stop() current_sample = progress.current_sample_name or "None" - avg_time = progress.avg_time_per_sample if progress.avg_time_per_sample > 0 else self.DEFAULT_SAMPLE_ESTIMATE_S + avg_time = ( + progress.avg_time_per_sample + if progress.avg_time_per_sample > 0 + else self.DEFAULT_SAMPLE_ESTIMATE_S + ) queue_remaining = avg_time * self._queue_count eta = time.time() + queue_remaining if queue_remaining > 0 else None @@ -182,4 +189,6 @@ class CompactAutomationProgressStrip(QFrame): states = {step.step: step.status for step in progress.steps} for step, label in self._step_labels.items(): - label.setText(self._format_step_html(step, states.get(step, StepStatus.PENDING))) \ No newline at end of file + label.setText( + self._format_step_html(step, states.get(step, StepStatus.PENDING)) + ) diff --git a/src/aare/gui/widgets/busy_overlay.py b/src/aare/gui/widgets/busy_overlay.py index 0c2119ee..217a6f6d 100644 --- a/src/aare/gui/widgets/busy_overlay.py +++ b/src/aare/gui/widgets/busy_overlay.py @@ -1,10 +1,9 @@ from dataclasses import dataclass +from aarecommon.models.models import SessionsStateEnum +from aarecommon.models.tell import TellStateModel from PySide6.QtGui import QColor -from aare.common.tell_models import TellStateModel -from aare.common.models import SessionsStateEnum - @dataclass(frozen=True) class BusyOverlayStyle: @@ -18,10 +17,10 @@ class BusyOverlayStyle: def build_busy_overlay_style( - *, - is_busy: bool, - tell_state: TellStateModel | None, - session_state: SessionsStateEnum | None = None, + *, + is_busy: bool, + tell_state: TellStateModel | None, + session_state: SessionsStateEnum | None = None, ) -> BusyOverlayStyle | None: if session_state == SessionsStateEnum.Vacant: return BusyOverlayStyle( @@ -51,7 +50,9 @@ def build_busy_overlay_style( if not is_busy: return None - activity_value = str(getattr(getattr(tell_state, "activity", None), "value", "") or "").lower() + activity_value = str( + getattr(getattr(tell_state, "activity", None), "value", "") or "" + ).lower() if activity_value == "mounting": return BusyOverlayStyle( @@ -105,4 +106,4 @@ def build_busy_overlay_style( overlay_border=QColor(255, 225, 220, 235), overlay_text=QColor(255, 255, 255), accent_dot="#ffd8d1", - ) \ No newline at end of file + ) diff --git a/src/aare/gui/widgets/camera_image.py b/src/aare/gui/widgets/camera_image.py index f83fc5e6..d9be90da 100644 --- a/src/aare/gui/widgets/camera_image.py +++ b/src/aare/gui/widgets/camera_image.py @@ -2,44 +2,48 @@ import math import time from enum import Enum -from PySide6.QtCore import Qt, QTimer, QRect, QPoint, Signal, Slot, QPointF, QRectF +from aarecommon.config.logger import setup_logger +from aarecommon.math.coordinate import Coordinate, SmargonCoordinate +from aarecommon.math.sample_geometry import SampleGeometryModel +from aarecommon.models.models import ( + AutofocusSettings, + BeamlineStateEnum, + DAQStatusModel, + SampleCameraSettings, + SessionsStateEnum, +) +from PySide6.QtCore import QPoint, QPointF, QRect, QRectF, Qt, QTimer, Signal, Slot from PySide6.QtGui import ( - QPixmap, + QBrush, + QColor, + QCursor, + QFont, + QFontMetrics, + QLinearGradient, QPainter, QPen, - QColor, - QWheelEvent, + QPixmap, + QPolygonF, QTransform, - QCursor, - QLinearGradient, QFont, QFontMetrics, QPolygonF, QBrush, + QWheelEvent, ) from PySide6.QtWidgets import ( - QMenu, - QToolTip, - QGraphicsView, - QGraphicsScene, - QGraphicsPixmapItem, QFileDialog, QFrame, + QGraphicsPixmapItem, + QGraphicsScene, + QGraphicsView, + QMenu, + QToolTip, ) from aare.gui.models.bookmark import SmargonBookmarkList -from aare.gui.widgets.busy_overlay import build_busy_overlay_style, BusyOverlayStyle from aare.gui.scan_logic.raster_grid_manager import RasterGridManager - -from aare.common.models import ( - DAQStatusModel, - AutofocusSettings, - SampleCameraSettings, - BeamlineStateEnum, - SessionsStateEnum -) -from aare.common.coordinate import Coordinate, SmargonCoordinate -from aare.common.sample_geometry import SampleGeometryModel -from aare.common.logger_config import setup_logger +from aare.gui.widgets.busy_overlay import BusyOverlayStyle, build_busy_overlay_style logger = setup_logger("aareGUI") + class SampleCameraImageState(Enum): IDLE = 0 DRAWING_RASTER_GRID = 1 @@ -47,6 +51,7 @@ class SampleCameraImageState(Enum): RESIZE_RASTER_GRID = 3 BEAM_MARKING = 4 + class SampleCameraImageLabel(QGraphicsView): smargon = Signal(SmargonCoordinate) @@ -66,11 +71,11 @@ class SampleCameraImageLabel(QGraphicsView): switch_raster_grid = Signal() def __init__( - self, - geom: SampleGeometryModel, - raster: RasterGridManager, - default_image: str | None, - parent=None, + self, + geom: SampleGeometryModel, + raster: RasterGridManager, + default_image: str | None, + parent=None, ): super().__init__(parent) @@ -163,7 +168,9 @@ class SampleCameraImageLabel(QGraphicsView): self.viewport().setCursor(Qt.CursorShape.ForbiddenCursor) self.setToolTip("Sample camera feed unavailable") - def __show_camera_unavailable_tooltip(self, event, action: str = "Sample camera interaction") -> None: + def __show_camera_unavailable_tooltip( + self, event, action: str = "Sample camera interaction" + ) -> None: QToolTip.showText( self.mapToGlobal(event.pos()), f"{action} disabled: sample camera feed unavailable", @@ -182,8 +189,8 @@ class SampleCameraImageLabel(QGraphicsView): @Slot(dict) def update_detections(self, payload: dict): try: - self.__det_shape = payload.get('shape', None) - self.__detections = payload.get('boxes', []) or [] + self.__det_shape = payload.get("shape", None) + self.__detections = payload.get("boxes", []) or [] except Exception as e: logger.debug(f"Exception in update_detections: {e}") self.__detections = [] @@ -262,12 +269,7 @@ class SampleCameraImageLabel(QGraphicsView): position_x = int((viewport_width - bg_width) / 2) position_y = int(viewport_height * 0.68 - bg_height / 2) - bg_rect = QRect( - position_x, - position_y, - bg_width, - bg_height - ) + bg_rect = QRect(position_x, position_y, bg_width, bg_height) painter.setPen(QPen(style.overlay_border, 2, Qt.PenStyle.SolidLine)) painter.setBrush(style.overlay_fill) @@ -275,8 +277,7 @@ class SampleCameraImageLabel(QGraphicsView): painter.setPen(QPen(style.overlay_text, 2, Qt.PenStyle.SolidLine)) text_pos = QPoint( - position_x + padding_x, - position_y + padding_y + font_metrics.ascent() + position_x + padding_x, position_y + padding_y + font_metrics.ascent() ) painter.drawText(text_pos, style.text) @@ -286,7 +287,10 @@ class SampleCameraImageLabel(QGraphicsView): if self.__busy_overlay_style is not None: return - if self.__session_state in (SessionsStateEnum.OwnedByYou, SessionsStateEnum.PendingElseToYou): + if self.__session_state in ( + SessionsStateEnum.OwnedByYou, + SessionsStateEnum.PendingElseToYou, + ): return painter.save() @@ -326,7 +330,9 @@ class SampleCameraImageLabel(QGraphicsView): painter.drawRoundedRect(bg_rect, 10, 10) painter.setPen(QPen(QColor(255, 255, 255))) - painter.drawText(QPoint(position_x + padding, position_y + padding + fm.ascent()), text) + painter.drawText( + QPoint(position_x + padding, position_y + padding + fm.ascent()), text + ) painter.restore() @@ -384,13 +390,18 @@ class SampleCameraImageLabel(QGraphicsView): def mousePressEvent(self, event): if not self.__camera_interaction_enabled(): - if event.button() in (Qt.MouseButton.LeftButton, Qt.MouseButton.RightButton): + if event.button() in ( + Qt.MouseButton.LeftButton, + Qt.MouseButton.RightButton, + ): self.__show_camera_unavailable_tooltip(event) event.accept() return self.start_point = self.mapToScene(event.pos()) - ctrl_override_move = bool(event.modifiers() & Qt.KeyboardModifier.ControlModifier) + ctrl_override_move = bool( + event.modifiers() & Qt.KeyboardModifier.ControlModifier + ) match self.__state: case SampleCameraImageState.BEAM_MARKING: @@ -409,7 +420,6 @@ class SampleCameraImageLabel(QGraphicsView): ): self.__state = SampleCameraImageState.MOVING_RASTER_GRID - def _update_grid(self): now = time.monotonic() if (now - self.__last_grid_update_ts) < self.__grid_update_min_interval_s: @@ -541,7 +551,9 @@ class SampleCameraImageLabel(QGraphicsView): return elif action == show_coord_action: self.__show_coords = not self.__show_coords - logger.info(f"Show-coordinates hover tooltip {'enabled' if self.__show_coords else 'disabled'}") + logger.info( + f"Show-coordinates hover tooltip {'enabled' if self.__show_coords else 'disabled'}" + ) self.update() elif action == scale_action: self.__autoscale = not self.__autoscale @@ -560,10 +572,15 @@ class SampleCameraImageLabel(QGraphicsView): elif action == grab_with_overlay_action: self.__screenshot_with_dialog(overlay=True) elif action == autofocus_action: - self.autofocus.emit(AutofocusSettings(center_x_pxl=None, center_y_pxl=None, - radius_pxl=30, - z_range_um=2000, - z_steps=10)) + self.autofocus.emit( + AutofocusSettings( + center_x_pxl=None, + center_y_pxl=None, + radius_pxl=30, + z_range_um=2000, + z_steps=10, + ) + ) elif action == beam_mark_action: c = self.mapToScene(event.pos()) self.update_beam_mark.emit(c.x(), c.y()) @@ -603,11 +620,19 @@ class SampleCameraImageLabel(QGraphicsView): self.setVerticalScrollBarPolicy(Qt.ScrollBarPolicy.ScrollBarAlwaysOff) self.setHorizontalScrollBarPolicy(Qt.ScrollBarPolicy.ScrollBarAlwaysOff) - if self.pixmap_item.boundingRect().width() == 0 or self.pixmap_item.boundingRect().height() == 0: + if ( + self.pixmap_item.boundingRect().width() == 0 + or self.pixmap_item.boundingRect().height() == 0 + ): ratio = 1.0 else: - ratio_w = self.viewport().size().width() / self.pixmap_item.boundingRect().width() - ratio_h = self.viewport().size().height() / self.pixmap_item.boundingRect().height() + ratio_w = ( + self.viewport().size().width() / self.pixmap_item.boundingRect().width() + ) + ratio_h = ( + self.viewport().size().height() + / self.pixmap_item.boundingRect().height() + ) ratio = min(ratio_w, ratio_h) if ratio < 0.1: @@ -638,7 +663,7 @@ class SampleCameraImageLabel(QGraphicsView): self.__bounding_box = s.box self.__tell_state = s.tell_state - new_session_state = s.session.session if hasattr(s, 'session') else None + new_session_state = s.session.session if hasattr(s, "session") else None if new_session_state != self.__session_state: self.__session_state = new_session_state self.update() @@ -672,7 +697,10 @@ class SampleCameraImageLabel(QGraphicsView): match self.__state: case SampleCameraImageState.BEAM_MARKING: - if event.button() == Qt.MouseButton.LeftButton and event.modifiers() & Qt.KeyboardModifier.ShiftModifier: + if ( + event.button() == Qt.MouseButton.LeftButton + and event.modifiers() & Qt.KeyboardModifier.ShiftModifier + ): point = self.mapToScene(event.pos()) self.update_beam_mark.emit(point.x(), point.y()) case SampleCameraImageState.IDLE: @@ -685,7 +713,9 @@ class SampleCameraImageLabel(QGraphicsView): sc = SmargonCoordinate(sh_mm=self.__geom.smargon_z(point.y())) self.smargon.emit(sc) else: - sample_coord = self.__geom.picture_to_sample(Coordinate(x=point.x(), y=point.y())) + sample_coord = self.__geom.picture_to_sample( + Coordinate(x=point.x(), y=point.y()) + ) smargon_coord = SmargonCoordinate( sh_mm=self.__geom.beamline_to_smargon(sample_coord) ) @@ -712,41 +742,45 @@ class SampleCameraImageLabel(QGraphicsView): sy = disp_h / float(img_h) color_map = { - 'pin': QColor('red'), - 'loop_all': QColor('green'), - 'loop_face': QColor('yellow'), - 'crystal': QColor('blue'), - 'needle': QColor('magenta'), - 'ice': QColor('cyan'), + "pin": QColor("red"), + "loop_all": QColor("green"), + "loop_face": QColor("yellow"), + "crystal": QColor("blue"), + "needle": QColor("magenta"), + "ice": QColor("cyan"), } for det in self.__detections: try: - x1 = det['x1'] * sx - y1 = det['y1'] * sy - x2 = det['x2'] * sx - y2 = det['y2'] * sy - poly = det.get('poly', None) - label = str(det.get('label', '')).lower() - conf = det.get('conf', 0.0) + x1 = det["x1"] * sx + y1 = det["y1"] * sy + x2 = det["x2"] * sx + y2 = det["y2"] * sy + poly = det.get("poly", None) + label = str(det.get("label", "")).lower() + conf = det.get("conf", 0.0) except Exception as e: logger.debug(f"Error in draw detection {det}: {e}") continue - color = color_map.get(label, QColor('magenta')) + color = color_map.get(label, QColor("magenta")) pen = QPen(color, 3) painter.setPen(pen) painter.setBrush(Qt.BrushStyle.NoBrush) - detection_rect = QRect(int(x1), int(y1), int(max(1, x2 - x1)), int(max(1, y2 - y1))) + detection_rect = QRect( + int(x1), int(y1), int(max(1, x2 - x1)), int(max(1, y2 - y1)) + ) painter.drawRect(detection_rect) # Only draw polygon if the separate toggle is enabled if self.__show_detection_polygons and poly and len(poly) >= 3: - polygon = QPolygonF([ - QPointF(x1 + float(p[0]) * sx, y1 + float(p[1]) * sy) - for p in poly - ]) + polygon = QPolygonF( + [ + QPointF(x1 + float(p[0]) * sx, y1 + float(p[1]) * sy) + for p in poly + ] + ) painter.drawPolygon(polygon) painter.setPen(QPen(QColor(255, 255, 255), 1)) @@ -794,7 +828,9 @@ class SampleCameraImageLabel(QGraphicsView): img_h, img_w = int(shape[0]), int(shape[1]) tx, ty = self.__smoothed_target_point except Exception as e: - logger.debug(f"Error using smoothed target point {self.__smoothed_target_point}: {e}") + logger.debug( + f"Error using smoothed target point {self.__smoothed_target_point}: {e}" + ) return sx = pix.width() / float(img_w) @@ -838,10 +874,7 @@ class SampleCameraImageLabel(QGraphicsView): fm = QFontMetrics(font) text_rect = fm.boundingRect(label_text) bubble_rect = QRectF( - px + 16, - py - 28, - text_rect.width() + 16, - text_rect.height() + 10 + px + 16, py - 28, text_rect.width() + 16, text_rect.height() + 10 ) painter.setPen(QPen(color, 2)) @@ -851,7 +884,7 @@ class SampleCameraImageLabel(QGraphicsView): painter.setPen(QPen(QColor(255, 255, 255), 1)) painter.drawText( QPointF(bubble_rect.left() + 8, bubble_rect.top() + 7 + fm.ascent()), - label_text + label_text, ) painter.restore() @@ -875,12 +908,14 @@ class SampleCameraImageLabel(QGraphicsView): lines.append((target_label, self.__target_color())) if self.__show_detections: - lines.extend([ - ("Pin", QColor("red")), - ("Loop", QColor("green")), - ("Face", QColor("yellow")), - ("Crystal", QColor("blue")), - ]) + lines.extend( + [ + ("Pin", QColor("red")), + ("Loop", QColor("green")), + ("Face", QColor("yellow")), + ("Crystal", QColor("blue")), + ] + ) if self.__show_coords: lines.append(("Coords tooltip", QColor(230, 230, 230))) @@ -897,14 +932,16 @@ class SampleCameraImageLabel(QGraphicsView): lines.append((target_label, self.__target_color())) if self.__show_detections: - lines.extend([ - ("Prediction: Pin", QColor("red")), - ("Prediction: Loop_all", QColor("green")), - ("Prediction: Loop_face", QColor("yellow")), - ("Prediction: Crystal", QColor("blue")), - ("Prediction: Needle", QColor("magenta")), - ("Prediction: Ice", QColor("cyan")), - ]) + lines.extend( + [ + ("Prediction: Pin", QColor("red")), + ("Prediction: Loop_all", QColor("green")), + ("Prediction: Loop_face", QColor("yellow")), + ("Prediction: Crystal", QColor("blue")), + ("Prediction: Needle", QColor("magenta")), + ("Prediction: Ice", QColor("cyan")), + ] + ) if self.__show_coords: lines.append(("Cursor tooltip: pixel coordinates", QColor(230, 230, 230))) @@ -954,14 +991,24 @@ class SampleCameraImageLabel(QGraphicsView): if color is not None: painter.setPen(QPen(color, 2)) painter.setBrush(color) - painter.drawRect(QRectF(bg_rect.left() + section_padding, y + 2, swatch_size, swatch_size)) + painter.drawRect( + QRectF( + bg_rect.left() + section_padding, + y + 2, + swatch_size, + swatch_size, + ) + ) else: painter.setPen(Qt.PenStyle.NoPen) painter.setBrush(Qt.BrushStyle.NoBrush) painter.setPen(QPen(QColor(240, 240, 240), 1)) painter.drawText( - QPointF(bg_rect.left() + section_padding + swatch_size + text_padding, y + fm.ascent() + 1), + QPointF( + bg_rect.left() + section_padding + swatch_size + text_padding, + y + fm.ascent() + 1, + ), text, ) y += line_height @@ -1018,7 +1065,7 @@ class SampleCameraImageLabel(QGraphicsView): if self.__bounding_box is None: return - painter.setPen(QPen(QColor(50,205, 50), 3, Qt.PenStyle.SolidLine)) + painter.setPen(QPen(QColor(50, 205, 50), 3, Qt.PenStyle.SolidLine)) painter.setBrush(Qt.BrushStyle.NoBrush) painter.drawRect( @@ -1115,7 +1162,10 @@ class SampleCameraImageLabel(QGraphicsView): mouse_scene_pos = self.mapToScene(mouse_view_pos) if event.key() == Qt.Key.Key_Shift: - if self.__state == SampleCameraImageState.IDLE and self.__camera_interaction_enabled(): + if ( + self.__state == SampleCameraImageState.IDLE + and self.__camera_interaction_enabled() + ): if not self.raster_timer.isActive(): self.load_image.emit(mouse_scene_pos) self.raster_timer.start(self.raster_timer_interval) @@ -1128,26 +1178,46 @@ class SampleCameraImageLabel(QGraphicsView): match self.__state: case SampleCameraImageState.BEAM_MARKING: if event.modifiers() & Qt.KeyboardModifier.AltModifier: - new_exp_time = self.__sam_cam.exposure * (1.0 + math.copysign(0.1, delta_y)) - new_settings = SampleCameraSettings(exposure=new_exp_time, gain=self.__sam_cam.gain) + new_exp_time = self.__sam_cam.exposure * ( + 1.0 + math.copysign(0.1, delta_y) + ) + new_settings = SampleCameraSettings( + exposure=new_exp_time, gain=self.__sam_cam.gain + ) self.samcam_updated.emit(new_settings) else: - new_exp_time = self.__sam_cam.exposure * (1.0 + math.copysign(0.5, delta_y)) - new_settings = SampleCameraSettings(exposure=new_exp_time, gain=self.__sam_cam.gain) + new_exp_time = self.__sam_cam.exposure * ( + 1.0 + math.copysign(0.5, delta_y) + ) + new_settings = SampleCameraSettings( + exposure=new_exp_time, gain=self.__sam_cam.gain + ) self.samcam_updated.emit(new_settings) case SampleCameraImageState.IDLE: if event.modifiers() & Qt.KeyboardModifier.AltModifier: - new_exp_time = self.__sam_cam.exposure * (1.0 + math.copysign(0.1, delta_y)) - new_settings = SampleCameraSettings(exposure=new_exp_time, gain=self.__sam_cam.gain) + new_exp_time = self.__sam_cam.exposure * ( + 1.0 + math.copysign(0.1, delta_y) + ) + new_settings = SampleCameraSettings( + exposure=new_exp_time, gain=self.__sam_cam.gain + ) self.samcam_updated.emit(new_settings) elif event.modifiers() & Qt.KeyboardModifier.ControlModifier: - new_exp_time = self.__sam_cam.exposure * (1.0 + math.copysign(0.5, delta_y)) - new_settings = SampleCameraSettings(exposure=new_exp_time, gain=self.__sam_cam.gain) + new_exp_time = self.__sam_cam.exposure * ( + 1.0 + math.copysign(0.5, delta_y) + ) + new_settings = SampleCameraSettings( + exposure=new_exp_time, gain=self.__sam_cam.gain + ) self.samcam_updated.emit(new_settings) elif event.modifiers() & Qt.KeyboardModifier.ShiftModifier: - self.set_omega.emit(self.__geom.omega_deg + math.copysign(10.0, delta_y)) + self.set_omega.emit( + self.__geom.omega_deg + math.copysign(10.0, delta_y) + ) else: - self.set_omega.emit(self.__geom.omega_deg + math.copysign(90.0, delta_y)) + self.set_omega.emit( + self.__geom.omega_deg + math.copysign(90.0, delta_y) + ) # Start the timer to throttle further events self.wheel_event_timer.start(self.wheel_event_threshold) diff --git a/src/aare/gui/widgets/local_contact_status_widget.py b/src/aare/gui/widgets/local_contact_status_widget.py index 5297c968..db27f08e 100644 --- a/src/aare/gui/widgets/local_contact_status_widget.py +++ b/src/aare/gui/widgets/local_contact_status_widget.py @@ -2,6 +2,7 @@ from __future__ import annotations from collections.abc import Iterable +from aarecommon.models.models import DAQStatusModel from PySide6.QtCore import Qt, Slot from PySide6.QtWidgets import ( QFrame, @@ -11,7 +12,6 @@ from PySide6.QtWidgets import ( QVBoxLayout, ) -from aare.common.models import DAQStatusModel from aare.gui.widgets.title_label import TitleLabel @@ -98,7 +98,9 @@ class LocalContactStatusWidget(QFrame): self._summary = QLabel(summary, self) self._summary.setWordWrap(True) - self._summary.setStyleSheet("border: none; background: transparent; color: #334155;") + self._summary.setStyleSheet( + "border: none; background: transparent; color: #334155;" + ) layout.addWidget(self._summary) self._grid = QGridLayout() @@ -129,12 +131,18 @@ class LocalContactStatusWidget(QFrame): self._row_widgets.clear() for row, key in enumerate(self._visible_fields): - title = QLabel(self.FIELD_TITLES.get(key, key.replace("_", " ").title()), self) - title.setStyleSheet("font-weight: 700; border: none; background: transparent; color: #1e293b;") + title = QLabel( + self.FIELD_TITLES.get(key, key.replace("_", " ").title()), self + ) + title.setStyleSheet( + "font-weight: 700; border: none; background: transparent; color: #1e293b;" + ) value = QLabel(self._badge("WAITING", tone="neutral"), self) value.setWordWrap(True) value.setTextInteractionFlags(Qt.TextInteractionFlag.TextSelectableByMouse) - value.setStyleSheet("border: none; background: transparent; color: #334155;") + value.setStyleSheet( + "border: none; background: transparent; color: #334155;" + ) self._row_widgets[key] = (title, value) self._grid.addWidget(title, row, 0, alignment=Qt.AlignmentFlag.AlignTop) self._grid.addWidget(value, row, 1) @@ -153,7 +161,13 @@ class LocalContactStatusWidget(QFrame): f"padding:2px 6px;border-radius:8px;'>{text}" ) - def _format_bool(self, value: bool | None, *, true_text: str = "CONNECTED", false_text: str = "DISCONNECTED") -> str: + def _format_bool( + self, + value: bool | None, + *, + true_text: str = "CONNECTED", + false_text: str = "DISCONNECTED", + ) -> str: if value is True: return self._badge(true_text, tone="good") if value is False: @@ -161,7 +175,11 @@ class LocalContactStatusWidget(QFrame): return self._badge("UNKNOWN", tone="neutral") def _format_device_status(self, device: str) -> str: - payload = self._device_state_payload.get(device, {}) if isinstance(self._device_state_payload, dict) else {} + payload = ( + self._device_state_payload.get(device, {}) + if isinstance(self._device_state_payload, dict) + else {} + ) mode = str(payload.get("mode", "unknown")) error = payload.get("error") @@ -213,13 +231,17 @@ class LocalContactStatusWidget(QFrame): display_name = getattr(state, "display_name", None) if callable(display_name): return self._badge(str(display_name()).upper(), tone="info") - return self._badge(str(getattr(state, 'name', state)).upper(), tone="info") + return self._badge(str(getattr(state, "name", state)).upper(), tone="info") def _format_field_value(self, key: str, status: DAQStatusModel) -> str: if key == "beamline_state": return self._format_state(status) if key == "busy": - return self._badge("BUSY", tone="warn") if bool(status.busy) else self._badge("IDLE", tone="good") + return ( + self._badge("BUSY", tone="warn") + if bool(status.busy) + else self._badge("IDLE", tone="good") + ) if key == "sample": return self._format_sample(status) if key == "tell_connected": @@ -258,10 +280,18 @@ class LocalContactStatusWidget(QFrame): return f"({box.top_x:.1f}, {box.top_y:.1f}) → ({box.bottom_x:.1f}, {box.bottom_y:.1f})" if key == "last_best_res": value = getattr(status, "last_best_res", None) - return f"{value:.3f} Å" if value is not None else self._badge("NONE", tone="neutral") + return ( + f"{value:.3f} Å" + if value is not None + else self._badge("NONE", tone="neutral") + ) if key == "last_best_b_factor": value = getattr(status, "last_best_b_factor", None) - return f"{value:.3f}" if value is not None else self._badge("NONE", tone="neutral") + return ( + f"{value:.3f}" + if value is not None + else self._badge("NONE", tone="neutral") + ) if key == "crystal_size": crystal = getattr(status, "crystal_size", None) if crystal is None: @@ -278,18 +308,34 @@ class LocalContactStatusWidget(QFrame): if key == "cryojet": return f"{getattr(status.bl, 'cryojet_K', 0.0):.2f} K" if key == "shutter": - return self._format_bool(getattr(status.bl, "shutter_open", None), true_text="OPEN", false_text="CLOSED") + return self._format_bool( + getattr(status.bl, "shutter_open", None), + true_text="OPEN", + false_text="CLOSED", + ) if key == "exposure_shutter": - return self._format_bool(getattr(status.bl, "exp_shutter_open", None), true_text="OPEN", false_text="CLOSED") + return self._format_bool( + getattr(status.bl, "exp_shutter_open", None), + true_text="OPEN", + false_text="CLOSED", + ) if key == "flux": return f"{getattr(status.bl, 'flux_ph_s', 0.0):.3g} ph/s" if key == "transmission": transmission = getattr(status.bl, "transmission", None) - return f"{100.0 * transmission:.2f} %" if transmission is not None else self._badge("NONE", tone="neutral") + return ( + f"{100.0 * transmission:.2f} %" + if transmission is not None + else self._badge("NONE", tone="neutral") + ) if key == "zoom": return f"{getattr(status.bl, 'zoom', 0.0):.3f}" if key == "commissioning_mode": - return self._format_bool(getattr(status.bl, "commissioning_mode", None), true_text="ON", false_text="OFF") + return self._format_bool( + getattr(status.bl, "commissioning_mode", None), + true_text="ON", + false_text="OFF", + ) if key == "omega": return f"{getattr(status.geom, 'omega_deg', 0.0):.3f}°" if key == "beam_size": @@ -320,7 +366,11 @@ class LocalContactStatusWidget(QFrame): return f"{getattr(status.diffraction, 'energy_keV', 0.0):.4f} keV" if key == "wavelength": wavelength = getattr(status.diffraction, "wavelength_angstrom", None) - return f"{wavelength:.5f} Å" if wavelength is not None else self._badge("NONE", tone="neutral") + return ( + f"{wavelength:.5f} Å" + if wavelength is not None + else self._badge("NONE", tone="neutral") + ) if key == "beam_center": beam_center = getattr(status.diffraction, "beam_center_pxl", None) if beam_center is None: @@ -346,4 +396,4 @@ class LocalContactStatusWidget(QFrame): @Slot(DAQStatusModel) def set_daq_status(self, status: DAQStatusModel) -> None: self._last_status = status - self._refresh() \ No newline at end of file + self._refresh() diff --git a/src/aare/gui/widgets/login.py b/src/aare/gui/widgets/login.py index 23c9cbe0..e9c7ac2a 100644 --- a/src/aare/gui/widgets/login.py +++ b/src/aare/gui/widgets/login.py @@ -2,15 +2,19 @@ import json import os import jwt -from PySide6.QtCore import Slot, QByteArray, QUrl, QUrlQuery -from PySide6.QtNetwork import QNetworkAccessManager, QNetworkRequest, QNetworkReply +from aarecommon.models.auth import get_user +from aarecommon.models.models import TokenData +from PySide6.QtCore import QByteArray, QUrl, QUrlQuery, Slot +from PySide6.QtNetwork import QNetworkAccessManager, QNetworkReply, QNetworkRequest from PySide6.QtWidgets import ( - QDialog, QVBoxLayout, QLineEdit, QPushButton, QLabel, QHBoxLayout + QDialog, + QHBoxLayout, + QLabel, + QLineEdit, + QPushButton, + QVBoxLayout, ) -from aare.common.models import TokenData -from aare.common.auth_models import get_user - class LoginDialog(QDialog): def __init__(self, base_url: str | None): @@ -38,9 +42,13 @@ class LoginDialog(QDialog): # Buttons button_layout = QHBoxLayout() self.ok_button = QPushButton("OK") - self.ok_button.clicked.connect(self.authenticate) # Close dialog with accept status + self.ok_button.clicked.connect( + self.authenticate + ) # Close dialog with accept status self.cancel_button = QPushButton("Cancel") - self.cancel_button.clicked.connect(self.reject) # Close dialog with reject status + self.cancel_button.clicked.connect( + self.reject + ) # Close dialog with reject status button_layout.addWidget(self.ok_button) button_layout.addWidget(self.cancel_button) @@ -49,27 +57,36 @@ class LoginDialog(QDialog): @Slot() def authenticate(self): if self.__base_url is None: - token_data = TokenData(sub=self.name_entry.text(), - staff=True, - session=15, - pgroups=["p16371", "p22233"]) + token_data = TokenData( + sub=self.name_entry.text(), + staff=True, + session=15, + pgroups=["p16371", "p22233"], + ) self.token = jwt.encode(token_data.model_dump(), "ABC123") self.accept() self.__network_manager = QNetworkAccessManager(self) request = QNetworkRequest(QUrl(f"{self.__base_url}/token")) - request.setHeader(QNetworkRequest.KnownHeaders.ContentTypeHeader, "application/x-www-form-urlencoded") + request.setHeader( + QNetworkRequest.KnownHeaders.ContentTypeHeader, + "application/x-www-form-urlencoded", + ) payload = QUrlQuery() payload.addQueryItem("username", f"{self.name_entry.text()}") payload.addQueryItem("password", "") payload_string = payload.toString() - self.__reply = self.__network_manager.post(request, QByteArray.fromStdString(payload_string)) + self.__reply = self.__network_manager.post( + request, QByteArray.fromStdString(payload_string) + ) self.__reply.finished.connect(self.handle_token_response) - @Slot() def handle_token_response(self): - if self.__reply is not None and self.__reply.error() == QNetworkReply.NetworkError.NoError: + if ( + self.__reply is not None + and self.__reply.error() == QNetworkReply.NetworkError.NoError + ): response_data = self.__reply.readAll().data() response_json = json.loads(response_data.decode("utf-8")) if "access_token" in response_json: diff --git a/src/aare/gui/widgets/message_box.py b/src/aare/gui/widgets/message_box.py index 461ba1ad..ba355800 100644 --- a/src/aare/gui/widgets/message_box.py +++ b/src/aare/gui/widgets/message_box.py @@ -1,8 +1,8 @@ import time -from PySide6.QtCore import QTimer, QEventLoop -from PySide6.QtWidgets import QMessageBox, QCheckBox -from aare.common.logger_config import setup_logger +from aarecommon.config.logger import setup_logger +from PySide6.QtCore import QEventLoop, QTimer +from PySide6.QtWidgets import QCheckBox, QMessageBox logger = setup_logger("aareGUI") @@ -44,7 +44,9 @@ def precondition_problems(ring_current, shutter_open, door_prohibited) -> list[s if not shutter_open: problems.append("Experiment safety shutter is closed.") if door_prohibited is False: - problems.append("Hutch is not in the prohibited state (doors open / not searched).") + problems.append( + "Hutch is not in the prohibited state (doors open / not searched)." + ) return problems @@ -66,7 +68,9 @@ def precondition_check(parent, *, ring_current, shutter_open, door_prohibited) - box.setIcon(QMessageBox.Icon.Warning) box.setWindowTitle("Beamline not ready") box.setText("\n".join(problems) + "\n\nDo you wish to continue?") - box.setStandardButtons(QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No) + box.setStandardButtons( + QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No + ) box.setDefaultButton(QMessageBox.StandardButton.No) snooze_cb = QCheckBox("Don't ask me again for 1 hour") box.setCheckBox(snooze_cb) @@ -76,7 +80,8 @@ def precondition_check(parent, *, ring_current, shutter_open, door_prohibited) - precondition_snooze.snooze() return proceed -def reply_box(parent, title: str = "Warning", msg: str = "Warning." ): + +def reply_box(parent, title: str = "Warning", msg: str = "Warning."): return QMessageBox.question( parent, title, @@ -85,11 +90,16 @@ def reply_box(parent, title: str = "Warning", msg: str = "Warning." ): QMessageBox.StandardButton.No, ) -def timer_box(parent, title: str = "Warning", msg: str = "Warning.", condition_func=None) -> QMessageBox: + +def timer_box( + parent, title: str = "Warning", msg: str = "Warning.", condition_func=None +) -> QMessageBox: box = QMessageBox(parent) box.setWindowTitle(title) box.setText(f"{msg} Do you wish to continue?") - box.setStandardButtons(QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No) + box.setStandardButtons( + QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No + ) box.show() timer = QTimer(box) @@ -119,6 +129,7 @@ def timer_box(parent, title: str = "Warning", msg: str = "Warning.", condition_f timer.start() return box + def ring_current_low_check(parent, ring_current) -> bool: if ring_current is None: logger.debug("Ring current: unknown") @@ -138,6 +149,7 @@ def ring_current_low_check(parent, ring_current) -> bool: reply = reply_box(parent, title="Ring current too low", msg=msg) return reply == QMessageBox.StandardButton.Yes + def experiment_hutch_shutter_check(parent, shutter_state) -> bool: if shutter_state: logger.debug("Experiment shutter open") @@ -146,16 +158,23 @@ def experiment_hutch_shutter_check(parent, shutter_state) -> bool: reply = reply_box( parent, title="Experiment shutter open", - msg=f"Experiment shutter is Closed." + msg=f"Experiment shutter is Closed.", ) return reply == QMessageBox.StandardButton.Yes + def ring_current_auto_check(parent, ring_current, check_func) -> bool: - msg = "Ring current: unknown." if ring_current is None else f"Ring current is low {round(ring_current, 2)} mA." + msg = ( + "Ring current: unknown." + if ring_current is None + else f"Ring current is low {round(ring_current, 2)} mA." + ) return conditions_auto_check(parent, msg, check_func, title="Ring current too low") -def conditions_auto_check(parent, msg: str, check_func, title: str = "Beamline not ready") -> bool: +def conditions_auto_check( + parent, msg: str, check_func, title: str = "Beamline not ready" +) -> bool: """Pause-and-wait dialog that auto-resumes when ``check_func()`` becomes True. Used to hold automation between samples until the beamline recovers. The @@ -165,13 +184,15 @@ def conditions_auto_check(parent, msg: str, check_func, title: str = "Beamline n """ box = timer_box(parent, title=title, msg=msg, condition_func=check_func) loop = QEventLoop() + def finish(_=None): if loop.isRunning(): loop.quit() + box.finished.connect(finish) box.show() loop.exec() # returns as soon as box finishes without blocking UI processing of its signals result = box.result() - return result == QMessageBox.StandardButton.Yes \ No newline at end of file + return result == QMessageBox.StandardButton.Yes diff --git a/src/aare/gui/widgets/status_bar.py b/src/aare/gui/widgets/status_bar.py index d6a1890f..b9a538d7 100644 --- a/src/aare/gui/widgets/status_bar.py +++ b/src/aare/gui/widgets/status_bar.py @@ -1,20 +1,32 @@ import math -from PySide6.QtCore import Signal, Slot, QPoint, QTimer +from aarecommon.config.logger import setup_logger +from aarecommon.models.auth import BatonRequestStatus, BatonStatus +from aarecommon.models.models import ( + BeamlineStateEnum, + DAQStatusModel, + SessionsStateEnum, + TokenData, +) +from PySide6.QtCore import QPoint, QTimer, Signal, Slot from PySide6.QtGui import QFont -from PySide6.QtWidgets import QStatusBar, QDialog, QMenu, QMessageBox, QLabel, QSizePolicy +from PySide6.QtWidgets import ( + QDialog, + QLabel, + QMenu, + QMessageBox, + QSizePolicy, + QStatusBar, +) -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 - logger = setup_logger("aareGUI") + class StatusBar(QStatusBar): set_pgroup = Signal(str) dewar_exchange = Signal() @@ -56,7 +68,9 @@ class StatusBar(QStatusBar): self.message_label = QLabel("", self) self.message_label.setVisible(False) - self.message_label.setSizePolicy(QSizePolicy.Policy.Maximum, QSizePolicy.Policy.Preferred) + self.message_label.setSizePolicy( + QSizePolicy.Policy.Maximum, QSizePolicy.Policy.Preferred + ) self.sharpness = ValueLabel("Samcam image sharpness", "", self) self.samcam_fps = ValueLabel("Samcam FPS", "fps", self) @@ -147,7 +161,9 @@ class StatusBar(QStatusBar): if status.bl.ring_current_mA < 5.0: self.ring_current.set_value(f"{status.bl.ring_current_mA:.2f}", "red") elif status.bl.ring_current_mA < 390.0: - self.ring_current.set_value(f"{status.bl.ring_current_mA:.2f}", "orange") + self.ring_current.set_value( + f"{status.bl.ring_current_mA:.2f}", "orange" + ) else: self.ring_current.set_value(f"{status.bl.ring_current_mA:.2f}") @@ -162,17 +178,27 @@ class StatusBar(QStatusBar): self.cryo_label.set_value(f"{status.bl.cryojet_K:.1f}", "red") if status.bl.shutter_open: - self.shutter_label.setText(f"""Fast Shutter: Open ☢️ """) + self.shutter_label.setText( + f"""Fast Shutter: Open ☢️ """ + ) else: - self.shutter_label.setText(f"""Fast Shutter: Closed 🚪 """) + self.shutter_label.setText( + f"""Fast Shutter: Closed 🚪 """ + ) if status.bl.exp_shutter_open: - self.exp_shutter_label.setText("""ExpHutch Shutter: Open """) + self.exp_shutter_label.setText( + """ExpHutch Shutter: Open """ + ) else: - self.exp_shutter_label.setText("""ExpHutch Shutter: Closed 🚪 """) + self.exp_shutter_label.setText( + """ExpHutch Shutter: Closed 🚪 """ + ) if status.session.current_pgroup is not None: - self.pgroup_label.setText(f"""p-group: {status.session.current_pgroup} """) + self.pgroup_label.setText( + f"""p-group: {status.session.current_pgroup} """ + ) else: self.pgroup_label.setText(f"Inactive p-group ") @@ -185,7 +211,12 @@ class StatusBar(QStatusBar): if status.tell_state.activity.value == "error": tell_color = "red" - elif status.tell_state.activity.value in {"mounting", "unmounting", "drying", "cooling"}: + elif status.tell_state.activity.value in { + "mounting", + "unmounting", + "drying", + "cooling", + }: tell_color = "orange" else: tell_color = "green" @@ -210,7 +241,9 @@ class StatusBar(QStatusBar): elif status.session.session == SessionsStateEnum.OwnedByElse: session_flag = """ Other 🔒 """ elif status.session.session == SessionsStateEnum.PendingYouToElse: - session_flag = """ Waiting... ⏳ """ + session_flag = ( + """ Waiting... ⏳ """ + ) elif status.session.session == SessionsStateEnum.PendingElseToYou: session_flag = """ Request! ⚡ """ @@ -246,13 +279,14 @@ class StatusBar(QStatusBar): logger.info(f"Incoming baton request detected: {status.pending_request}") self._emit_incoming_baton_request(status) elif not incoming and self._baton_request_dialog is not None: - logger.info("Baton request no longer incoming, closing local dialog reference") + logger.info( + "Baton request no longer incoming, closing local dialog reference" + ) try: self._baton_request_dialog.close() except Exception as e: logger.error(f"Error closing baton request dialog: {e}") - self._baton_request_dialog = None\ - + self._baton_request_dialog = None def _show_pgroup_after_baton_grant(self) -> None: """ @@ -273,10 +307,12 @@ class StatusBar(QStatusBar): 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_received.emit( + { + "requester": requester, + "timeout": timeout, + } + ) @Slot() def _on_baton_dialog_accepted(self): @@ -318,22 +354,32 @@ class StatusBar(QStatusBar): def show_session_menu(self): menu = QMenu(self) is_busy = self.__status and self.__status.busy - session_state = self.__status.session.session if self.__status else SessionsStateEnum.Vacant + session_state = ( + self.__status.session.session if self.__status else SessionsStateEnum.Vacant + ) # Determine if we are the holder or waiting for baton - is_yours = session_state in (SessionsStateEnum.OwnedByYou, SessionsStateEnum.PendingElseToYou) + is_yours = session_state in ( + SessionsStateEnum.OwnedByYou, + SessionsStateEnum.PendingElseToYou, + ) # Check baton status for fallback if status.session is not yet updated if not is_yours and self._baton_status: - is_yours = self._baton_status.you_are_holder or self._baton_status.incoming_request + is_yours = ( + self._baton_status.you_are_holder or self._baton_status.incoming_request + ) is_vacant = session_state == SessionsStateEnum.Vacant - is_other = session_state in (SessionsStateEnum.OwnedByElse, SessionsStateEnum.PendingYouToElse) + is_other = session_state in ( + SessionsStateEnum.OwnedByElse, + SessionsStateEnum.PendingYouToElse, + ) has_pending = session_state == SessionsStateEnum.PendingYouToElse # Determine holder info from baton status holder_is_staff = ( - self._baton_status and - self._baton_status.holder and - self._baton_status.holder.is_staff + self._baton_status + and self._baton_status.holder + and self._baton_status.holder.is_staff ) # --- GRAB / REQUEST --- @@ -355,7 +401,9 @@ class StatusBar(QStatusBar): action_grab.triggered.connect(self._on_grab_clicked) elif holder_is_staff: # Non-staff cannot request from staff - allowed = self._baton_status and getattr(self._baton_status, "allow_non_staff_request", False) + 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: @@ -368,7 +416,7 @@ class StatusBar(QStatusBar): elif is_yours: # You have it - show release option action_release = menu.addAction("Release") - action_release.setEnabled(True) # Always allow release + action_release.setEnabled(True) # Always allow release action_release.triggered.connect(self._on_release_clicked) if session_state == SessionsStateEnum.PendingElseToYou: @@ -390,7 +438,10 @@ class StatusBar(QStatusBar): label_geometry = self.session_label.geometry() menu_width = max(label_geometry.width(), menu.sizeHint().width()) - menu.move(self.mapToGlobal(label_geometry.topLeft()) - QPoint(0, menu.sizeHint().height())) + menu.move( + self.mapToGlobal(label_geometry.topLeft()) + - QPoint(0, menu.sizeHint().height()) + ) menu.setFixedWidth(menu_width) menu.exec() @@ -404,7 +455,10 @@ class StatusBar(QStatusBar): allow_pgroup_menu = self.__is_staff or in_curr - owned_by_you = self.__status and self.__status.session.session == SessionsStateEnum.OwnedByYou + owned_by_you = ( + self.__status + and self.__status.session.session == SessionsStateEnum.OwnedByYou + ) if allow_pgroup_menu: logger.info("showing pgroup menu") @@ -415,7 +469,10 @@ class StatusBar(QStatusBar): action_1.triggered.connect(self.show_change_dialog) label_geometry = self.pgroup_label.geometry() - menu.move(self.mapToGlobal(label_geometry.topLeft()) - QPoint(0, menu.sizeHint().height())) + menu.move( + self.mapToGlobal(label_geometry.topLeft()) + - QPoint(0, menu.sizeHint().height()) + ) menu.setFixedWidth(label_geometry.width()) menu.exec() @@ -429,7 +486,10 @@ class StatusBar(QStatusBar): action_2.triggered.connect(self.open_shutter_clicked) label_geometry = self.shutter_label.geometry() - menu.move(self.mapToGlobal(label_geometry.topLeft()) - QPoint(0, menu.sizeHint().height())) + menu.move( + self.mapToGlobal(label_geometry.topLeft()) + - QPoint(0, menu.sizeHint().height()) + ) menu.setFixedWidth(label_geometry.width()) menu.exec() @@ -452,28 +512,37 @@ class StatusBar(QStatusBar): action_0.setEnabled(False) menu.addSeparator() - if self.__status.state in [BeamlineStateEnum.RobotSampleExchange, - BeamlineStateEnum.SampleExchange, - BeamlineStateEnum.DewarTransfer, - BeamlineStateEnum.BeamLocation, - BeamlineStateEnum.DataCollection, - BeamlineStateEnum.XrayFluorescence, - BeamlineStateEnum.XtalSnapshot]: + if self.__status.state in [ + BeamlineStateEnum.RobotSampleExchange, + BeamlineStateEnum.SampleExchange, + BeamlineStateEnum.DewarTransfer, + BeamlineStateEnum.BeamLocation, + BeamlineStateEnum.DataCollection, + BeamlineStateEnum.XrayFluorescence, + BeamlineStateEnum.XtalSnapshot, + ]: action_1 = menu.addAction("Sample alignment") action_1.triggered.connect(self.sa) - elif self.__status.state in [BeamlineStateEnum.SampleAlignment, ]: + elif self.__status.state in [ + BeamlineStateEnum.SampleAlignment, + ]: action_2 = menu.addAction("Manual sample exchange") action_2.triggered.connect(self.se) action_3 = menu.addAction("Dewar transfer") action_3.triggered.connect(self.dl) action_4 = menu.addAction("Beam location") action_4.triggered.connect(self.beam_location) - elif self.__status.state in [BeamlineStateEnum.Maintenance, ]: + elif self.__status.state in [ + BeamlineStateEnum.Maintenance, + ]: action_2 = menu.addAction("Manual sample exchange") action_2.triggered.connect(self.se) label_geometry = self.state_label.geometry() - menu.move(self.mapToGlobal(label_geometry.topLeft()) - QPoint(0, menu.sizeHint().height())) + menu.move( + self.mapToGlobal(label_geometry.topLeft()) + - QPoint(0, menu.sizeHint().height()) + ) menu.setFixedWidth(label_geometry.width()) menu.exec() @@ -500,11 +569,20 @@ class StatusBar(QStatusBar): def show_change_dialog(self): logger.debug(self.__decoded_token.pgroups) curr = self.__status.session.current_pgroup - pgroups = [str(p) for p in (self.__allowed_pgroups or []) if p is not None and str(p).strip()] + pgroups = [ + str(p) + for p in (self.__allowed_pgroups or []) + if p is not None and str(p).strip() + ] if self.__is_staff: + def _on_loaded(lst: list): try: - merged = {str(p).strip() for p in (lst or []) if p is not None and str(p).strip()} + merged = { + str(p).strip() + for p in (lst or []) + if p is not None and str(p).strip() + } if not merged: merged = set(pgroups) self._generate_pgroup_dialogue(curr=curr, pgroups=sorted(merged)) @@ -563,7 +641,9 @@ class StatusBar(QStatusBar): self.get_all_pgroups.emit() return - def _generate_pgroup_dialogue(self, curr: str | None = None, pgroups: list | None = None): + 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: @@ -574,7 +654,7 @@ class StatusBar(QStatusBar): self, "Invalid P-Group", f"P-group '{entered_text}' is not in your allowed list.\n" - f"Please select from: {', '.join(pgroups)}" + f"Please select from: {', '.join(pgroups)}", ) self._generate_pgroup_dialogue(curr=curr, pgroups=pgroups) return diff --git a/tests/conftest.py b/tests/conftest.py index 3d91fa42..638a9cfb 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -5,10 +5,9 @@ from unittest.mock import MagicMock, patch import numpy as np import pytest -from PySide6.QtWidgets import QApplication +from aarecommon.models.models import DewarAddress, SampleShortInfo from fastapi.testclient import TestClient - -from aare.common.models import SampleShortInfo, DewarAddress +from PySide6.QtWidgets import QApplication # Ensure safe defaults during tests os.environ.setdefault("BEAMLINE", "SIMULATED") @@ -52,15 +51,17 @@ def sample_info(): @pytest.fixture(scope="session") def server_module(): import aare.daq.server as server + return server @pytest.fixture def mock_backend(server_module): - with patch.object(server_module, "daq", create=True) as m_daq, \ - patch.object(server_module, "bl", create=True) as m_bl, \ - patch.object(server_module, "cfg", create=True) as m_cfg: - + with ( + patch.object(server_module, "daq", create=True) as m_daq, + patch.object(server_module, "bl", create=True) as m_bl, + patch.object(server_module, "cfg", create=True) as m_cfg, + ): m_daq.busy = False m_daq.camera_image = np.zeros((10, 10, 3), dtype=np.uint8) @@ -92,11 +93,13 @@ def auth_token_data(): @pytest.fixture def client(server_module, mock_backend, auth_token_data): - with patch("aare.daq.auth.parse_token", return_value=auth_token_data), \ - patch("cv2.imencode", return_value=(True, np.array([1, 2, 3], dtype=np.uint8))), \ - patch.object(server_module, "mx_beamline", return_value=mock_backend["bl"]), \ - patch.object(server_module, "BeamlineConfig", return_value=mock_backend["cfg"]), \ - patch.object(server_module, "AareDAQ", return_value=mock_backend["daq"]): + with ( + patch("aare.daq.auth.parse_token", return_value=auth_token_data), + patch("cv2.imencode", return_value=(True, np.array([1, 2, 3], dtype=np.uint8))), + patch.object(server_module, "mx_beamline", return_value=mock_backend["bl"]), + patch.object(server_module, "BeamlineConfig", return_value=mock_backend["cfg"]), + patch.object(server_module, "AareDAQ", return_value=mock_backend["daq"]), + ): with TestClient(server_module.app) as c: yield c @@ -113,17 +116,17 @@ def api(client, mock_backend): @pytest.fixture def daq_status_factory(): - from aare.common.coordinate import Coordinate, SmargonCoordinate - from aare.common.diffraction_geometry import DiffractionGeometry - from aare.common.models import ( + from aarecommon.math.coordinate import Coordinate, SmargonCoordinate + from aarecommon.math.diffraction_geometry import DiffractionGeometry + from aarecommon.models.models import ( BeamlineStateEnum, BeamlineStatus, CrystalSize, DAQStatusModel, SampleCameraSettings, SampleGeometryModel, - SessionStatus, SessionsStateEnum, + SessionStatus, ) def _build( @@ -202,4 +205,4 @@ def daq_status_factory(): return status - return _build \ No newline at end of file + return _build diff --git a/tests/unit/common/test_aare_exception.py b/tests/unit/common/test_aare_exception.py index 3379d555..82ce9be8 100644 --- a/tests/unit/common/test_aare_exception.py +++ b/tests/unit/common/test_aare_exception.py @@ -10,53 +10,52 @@ These cover: from __future__ import annotations import pytest - -from aare.common.exception_handler import ( - AareException, - AutomationError, - AareUserError, +from aarecommon.errors.exception_handler import ( AareAuthError, - TellException, - SmargonException, - AerotechException, - JFJochException, - BECException, - BeamlineStateException, - DataCollectionException, - RasterScanException, - AutoRasterSampleSkipped, - LoopCenteringFailed, - MountingFailed, - UnmountingFailed, - TellCommunicationError, - TellConnectionException, - TellCommandWhileBusyException, - WarningTellException, - CriticalTellException, - SmargonCommunicationError, - AerotechCommunicationError, - JFJochCommunicationError, - BECCommunicationError, AareDBCommunicationError, - MagnetPositionSensorErorr, - SmartMagnetFaultException, - TransformationInvalidException, - StateTransitionFailed, - MaintenanceStateException, + AareException, + AareUserError, + AerotechCommunicationError, + AerotechException, + AuthenticationException, + AutomationError, + AutoRasterSampleSkipped, + AXCFailed, BeamlineBusyException, BeamlineBusyTimeoutException, - AXCFailed, - SampleException, + BeamlineStateException, + BECCommunicationError, + BECException, + CriticalTellException, + DataCollectionException, + JFJochCommunicationError, + JFJochException, + LoopCenteringFailed, + MagnetPositionSensorErorr, + MaintenanceStateException, ManualMountException, - AuthenticationException, + MountingFailed, + RasterScanException, + SampleException, + SmargonCommunicationError, + SmargonException, + SmartMagnetFaultException, + StateTransitionFailed, + TellCommandWhileBusyException, + TellCommunicationError, + TellConnectionException, + TellException, + TransformationInvalidException, + UnmountingFailed, UserRightsException, + WarningTellException, ) - # --------------------------------------------------------------------------- # Class-level criticality defaults # --------------------------------------------------------------------------- + def test_aare_exception_class_critical_default_false(): assert AareException.critical is False @@ -78,7 +77,9 @@ def test_phase3_classes_have_critical_default_true(): SmartMagnetFaultException, TransformationInvalidException, ): - assert cls.critical is True, f"{cls.__name__} should default to critical=True in Phase 3" + assert cls.critical is True, ( + f"{cls.__name__} should default to critical=True in Phase 3" + ) def test_other_core_classes_keep_critical_default_false(): @@ -106,6 +107,7 @@ def test_other_core_classes_keep_critical_default_false(): # Instance-level critical override # --------------------------------------------------------------------------- + def test_instance_critical_override_true(): exc = MountingFailed("mount fail", critical=True) assert exc.critical is True @@ -134,16 +136,30 @@ def test_aaredb_communication_error_instance_critical_override(): # Routing-root membership # --------------------------------------------------------------------------- + def test_automation_errors_are_automation_error(): - for cls in (LoopCenteringFailed, MountingFailed, UnmountingFailed, - AXCFailed, TellCommunicationError, SmargonCommunicationError, - AerotechCommunicationError, JFJochCommunicationError, - BECCommunicationError, AareDBCommunicationError, - BeamlineBusyException, RasterScanException, - MagnetPositionSensorErorr, SmartMagnetFaultException, - TransformationInvalidException, AutoRasterSampleSkipped): + for cls in ( + LoopCenteringFailed, + MountingFailed, + UnmountingFailed, + AXCFailed, + TellCommunicationError, + SmargonCommunicationError, + AerotechCommunicationError, + JFJochCommunicationError, + BECCommunicationError, + AareDBCommunicationError, + BeamlineBusyException, + RasterScanException, + MagnetPositionSensorErorr, + SmartMagnetFaultException, + TransformationInvalidException, + AutoRasterSampleSkipped, + ): exc = cls() if cls is not AutoRasterSampleSkipped else cls("skipped") - assert isinstance(exc, AutomationError), f"{cls.__name__} should be AutomationError" + assert isinstance(exc, AutomationError), ( + f"{cls.__name__} should be AutomationError" + ) assert isinstance(exc, AareException), f"{cls.__name__} should be AareException" @@ -165,12 +181,21 @@ def test_user_errors_are_aare_user_error(): # Family-base membership (drives watcher matching) # --------------------------------------------------------------------------- + def test_tell_family_membership(): - for cls in (TellCommunicationError, TellConnectionException, - CriticalTellException, TellCommandWhileBusyException, - WarningTellException, MountingFailed, UnmountingFailed): + for cls in ( + TellCommunicationError, + TellConnectionException, + CriticalTellException, + TellCommandWhileBusyException, + WarningTellException, + MountingFailed, + UnmountingFailed, + ): exc = cls() - assert isinstance(exc, TellException), f"{cls.__name__} should be in Tell family" + assert isinstance(exc, TellException), ( + f"{cls.__name__} should be in Tell family" + ) assert isinstance(exc, AutomationError) @@ -199,8 +224,12 @@ def test_bec_family_membership(): def test_beamline_state_family_membership(): - for cls in (StateTransitionFailed, MaintenanceStateException, - BeamlineBusyException, BeamlineBusyTimeoutException): + for cls in ( + StateTransitionFailed, + MaintenanceStateException, + BeamlineBusyException, + BeamlineBusyTimeoutException, + ): exc = cls() assert isinstance(exc, BeamlineStateException) assert isinstance(exc, AutomationError) @@ -216,6 +245,7 @@ def test_data_collection_family_membership(): # Backward compat: existing __str__/.message contracts preserved # --------------------------------------------------------------------------- + def test_message_attribute_preserved(): exc = LoopCenteringFailed("custom message") assert exc.message == "custom message" diff --git a/tests/unit/common/test_aerotech_models.py b/tests/unit/common/test_aerotech_models.py index ada2e1be..094e927b 100644 --- a/tests/unit/common/test_aerotech_models.py +++ b/tests/unit/common/test_aerotech_models.py @@ -1,15 +1,23 @@ import pytest -from aare.common.aerotech_models import ( - TaskEnum, AxisEnum, AerotechRunEnum, VariableTypeEnum, - AerotechAxisStatus, AerotechStatus, AerotechTarget, AerotechRotationScanRequest +from aarecommon.models.aerotech import ( + AerotechAxisStatus, + AerotechRotationScanRequest, + AerotechRunEnum, + AerotechStatus, + AerotechTarget, + AxisEnum, + TaskEnum, + VariableTypeEnum, ) + def test_enums(): assert TaskEnum.TASK_0.value == 0 assert AxisEnum.X.value == "x" assert AerotechRunEnum.START.value == 1 assert VariableTypeEnum.REAL.value == 1 + def test_aerotech_axis_status(): status = AerotechAxisStatus( enabled=True, @@ -19,11 +27,12 @@ def test_aerotech_axis_status(): moving=False, position=10.0, status=1, - velocity=0.0 + velocity=0.0, ) assert status.position == 10.0 assert status.enabled is True + def test_aerotech_status_strings(): axis_status = AerotechAxisStatus( enabled=True, @@ -33,10 +42,10 @@ def test_aerotech_status_strings(): moving=False, position=1.234567, status=1, - velocity=0.1 + velocity=0.1, ) status = AerotechStatus(state="READY", x=axis_status) - + pretty = status.to_pretty_string() assert "STATE: READY" in pretty assert " X | pos= 1.234567" in pretty @@ -53,12 +62,14 @@ def test_aerotech_status_strings(): assert str(status) == pretty + def test_aerotech_target(): target = AerotechTarget(x=10.0, y=20.0) payload = target.to_payload() assert payload == {"x": 10.0, "y": 20.0} assert "z" not in payload + def test_aerotech_rotation_scan_request(): request = AerotechRotationScanRequest( rotation_deg=360.0, @@ -67,7 +78,7 @@ def test_aerotech_rotation_scan_request(): async_move=True, exp_time_s=0.1, incr_omega_deg=1.0, - steps=360 + steps=360, ) payload = request.to_payload() assert payload["rotation_deg"] == 360.0 diff --git a/tests/unit/common/test_autofocus_tools.py b/tests/unit/common/test_autofocus_tools.py index 2ba32019..222834f5 100644 --- a/tests/unit/common/test_autofocus_tools.py +++ b/tests/unit/common/test_autofocus_tools.py @@ -1,74 +1,81 @@ import numpy as np import pytest -from aare.common.autofocus_tools import focus_measure_edges, focus_measure_blob_size +from aarecommon.math.autofocus import focus_measure_blob_size, focus_measure_edges + def test_focus_measure_edges_all_zeros(): gray = np.zeros((100, 100), dtype=np.uint8) assert focus_measure_edges(gray) == 0.0 + def test_focus_measure_edges_sharp_vs_blurry(): # Sharp image (large blocks to survive GaussianBlur) sharp = np.zeros((100, 100), dtype=np.uint8) sharp[:, 0:50] = 255 - + # Blurry image (flat) blurry = np.full((100, 100), 128, dtype=np.uint8) - + fm_sharp = focus_measure_edges(sharp, verbose=True) fm_blurry = focus_measure_edges(blurry) - + assert fm_sharp > fm_blurry assert fm_blurry == 0.0 + def test_focus_measure_edges_with_mask(): gray = np.zeros((100, 100), dtype=np.uint8) gray[40:60, 40:60] = 255 - + mask = np.zeros((100, 100), dtype=bool) mask[40:60, 40:60] = True - + fm_with_mask = focus_measure_edges(gray, mask=mask) assert fm_with_mask > 0 - + empty_mask = np.zeros((100, 100), dtype=bool) assert focus_measure_edges(gray, mask=empty_mask) == 0.0 + def test_focus_measure_edges_verbose(capsys): gray = np.zeros((100, 100), dtype=np.uint8) gray[40:60, 40:60] = 255 mask = np.ones((100, 100), dtype=bool) - + focus_measure_edges(gray, mask=mask, verbose=True) captured = capsys.readouterr() assert "focus=" in captured.out + def test_focus_measure_blob_size_all_zeros(): gray = np.zeros((100, 100), dtype=np.uint8) assert focus_measure_blob_size(gray) == 0.0 + def test_focus_measure_blob_size_sharp_vs_blurry(): # Small sharp blob sharp = np.zeros((100, 100), dtype=np.uint8) sharp[50, 50] = 255 - + # Larger blurry blob blurry = np.zeros((100, 100), dtype=np.uint8) blurry[45:55, 45:55] = 255 - + fm_sharp = focus_measure_blob_size(sharp) fm_blurry = focus_measure_blob_size(blurry) - + assert fm_sharp > fm_blurry + def test_focus_measure_blob_size_with_mask(): gray = np.zeros((100, 100), dtype=np.uint8) gray[50, 50] = 255 - + mask = np.zeros((100, 100), dtype=bool) mask[50, 50] = True - + fm_with_mask = focus_measure_blob_size(gray, mask=mask) assert fm_with_mask > 0 - + empty_mask = np.zeros((100, 100), dtype=bool) assert focus_measure_blob_size(gray, mask=empty_mask) == 0.0 diff --git a/tests/unit/common/test_beamline.py b/tests/unit/common/test_beamline.py index a9fec7a9..00de4a70 100644 --- a/tests/unit/common/test_beamline.py +++ b/tests/unit/common/test_beamline.py @@ -1,21 +1,26 @@ import os -import pytest from unittest.mock import patch -from aare.common.beamline import mx_beamline, MXBeamline + +from aarecommon.config.beamline import mx_beamline +from aarecommon.models.beamline import MXBeamline + def test_mx_beamline_default(): with patch.dict(os.environ, {}, clear=True): # If BEAMLINE is not set, it should return SIMULATED assert mx_beamline() == MXBeamline.SIMULATED + def test_mx_beamline_x10sa(): with patch.dict(os.environ, {"BEAMLINE": "x10sa"}): assert mx_beamline() == MXBeamline.X10SA + def test_mx_beamline_x06sa(): with patch.dict(os.environ, {"BEAMLINE": "X06SA "}): assert mx_beamline() == MXBeamline.X06SA + def test_mx_beamline_invalid(): with patch.dict(os.environ, {"BEAMLINE": "INVALID"}): assert mx_beamline() == MXBeamline.SIMULATED diff --git a/tests/unit/common/test_coordinate.py b/tests/unit/common/test_coordinate.py index 4f8062aa..c43f073a 100644 --- a/tests/unit/common/test_coordinate.py +++ b/tests/unit/common/test_coordinate.py @@ -1,6 +1,12 @@ -import pytest import numpy as np -from aare.common.coordinate import Coordinate, SmargonCoordinate, AerotechCoordinate, positive_coords +import pytest +from aarecommon.math.coordinate import ( + AerotechCoordinate, + Coordinate, + SmargonCoordinate, + positive_coords, +) + def test_coordinate_addition(): c1 = Coordinate(x=1, y=2, z=3) @@ -10,6 +16,7 @@ def test_coordinate_addition(): assert c3.y == 7 assert c3.z == 9 + def test_coordinate_subtraction(): c1 = Coordinate(x=5, y=7, z=9) c2 = Coordinate(x=1, y=2, z=3) @@ -18,6 +25,7 @@ def test_coordinate_subtraction(): assert c3.y == 5 assert c3.z == 6 + def test_coordinate_multiplication_scalar(): c1 = Coordinate(x=1, y=2, z=3) c2 = c1 * 2.0 @@ -25,11 +33,13 @@ def test_coordinate_multiplication_scalar(): assert c2.y == 4.0 assert c2.z == 6.0 + def test_coordinate_dot_product(): c1 = Coordinate(x=1, y=2, z=3) c2 = Coordinate(x=4, y=5, z=6) dot = c1 * c2 - assert dot == (1*4 + 2*5 + 3*6) + assert dot == (1 * 4 + 2 * 5 + 3 * 6) + def test_coordinate_division(): c1 = Coordinate(x=2, y=4, z=6) @@ -38,11 +48,13 @@ def test_coordinate_division(): assert c2.y == 2.0 assert c2.z == 3.0 + def test_coordinate_division_by_zero(): c1 = Coordinate(x=2, y=4, z=6) with pytest.raises(ValueError, match="Cannot divide by zero"): _ = c1 / 0.0 + def test_coordinate_normalize(): c1 = Coordinate(x=3, y=0, z=4) c2 = c1.normalize() @@ -50,66 +62,79 @@ def test_coordinate_normalize(): assert c2.y == 0.0 assert c2.z == 0.8 + def test_coordinate_normalize_zero(): c1 = Coordinate(x=0, y=0, z=0) with pytest.raises(ValueError, match="Cannot normalize a zero-magnitude vector."): c1.normalize() + def test_coordinate_rotate_x(): c1 = Coordinate(x=1, y=1, z=0) # Rotate 90 degrees around X. (1, 1, 0) -> (1, 0, 1) - c2 = c1.rotate(90, 'x') + c2 = c1.rotate(90, "x") assert np.isclose(c2.x, 1) assert np.isclose(c2.y, 0) assert np.isclose(c2.z, 1) + def test_coordinate_rotate_y(): c1 = Coordinate(x=1, y=0, z=1) # Rotate 90 degrees around Y. (1, 0, 1) -> (1, 0, -1) - c2 = c1.rotate(90, 'y') + c2 = c1.rotate(90, "y") assert np.isclose(c2.x, 1) assert np.isclose(c2.y, 0) assert np.isclose(c2.z, -1) + def test_coordinate_rotate_z(): c1 = Coordinate(x=1, y=0, z=0) # Rotate 90 degrees around Z. (1, 0, 0) -> (0, 1, 0) - c2 = c1.rotate(90, 'z') + c2 = c1.rotate(90, "z") assert np.isclose(c2.x, 0) assert np.isclose(c2.y, 1) assert np.isclose(c2.z, 0) + def test_coordinate_rotate_invalid_axis(): c1 = Coordinate(x=1, y=1, z=1) with pytest.raises(ValueError, match="Invalid axis"): - c1.rotate(90, 'w') + c1.rotate(90, "w") + def test_smargon_coordinate_equality(): s1 = SmargonCoordinate(sh_mm=Coordinate(x=1, y=1, z=1), phi_deg=10, chi_deg=20) - s2 = SmargonCoordinate(sh_mm=Coordinate(x=1.05, y=0.95, z=1.01), phi_deg=10.05, chi_deg=19.95) + s2 = SmargonCoordinate( + sh_mm=Coordinate(x=1.05, y=0.95, z=1.01), phi_deg=10.05, chi_deg=19.95 + ) assert s1 == s2 - + s3 = SmargonCoordinate(sh_mm=Coordinate(x=2, y=1, z=1), phi_deg=10, chi_deg=20) assert s1 != s3 + def test_aerotech_coordinate_equality(): a1 = AerotechCoordinate(at_mm=Coordinate(x=1, y=1, z=1), omega_deg=10) - a2 = AerotechCoordinate(at_mm=Coordinate(x=1.005, y=0.995, z=1.001), omega_deg=10.005) + a2 = AerotechCoordinate( + at_mm=Coordinate(x=1.005, y=0.995, z=1.001), omega_deg=10.005 + ) assert a1 == a2 - + a3 = AerotechCoordinate(at_mm=Coordinate(x=1.1, y=1, z=1), omega_deg=10) assert a1 != a3 + def test_positive_coords(): c1 = Coordinate(x=1, y=1, z=1) assert positive_coords(c1) == c1 - + with pytest.raises(ValueError, match="Coordinates must be positive"): positive_coords(Coordinate(x=-1, y=1, z=1)) - + with pytest.raises(ValueError, match="Coordinates must be positive"): positive_coords(Coordinate(x=1, y=-1, z=1)) + def test_coordinate_unsupported_ops(): c1 = Coordinate(x=1, y=1, z=1) assert c1.__add__(1) == NotImplemented @@ -117,10 +142,12 @@ def test_coordinate_unsupported_ops(): assert c1.__mul__("string") == NotImplemented assert c1.__truediv__("string") == NotImplemented + def test_smargon_coordinate_eq_not_implemented(): s1 = SmargonCoordinate(sh_mm=Coordinate(x=1, y=1, z=1), phi_deg=10, chi_deg=20) assert s1.__eq__(1) == NotImplemented + def test_aerotech_coordinate_eq_not_implemented(): a1 = AerotechCoordinate(at_mm=Coordinate(x=1, y=1, z=1), omega_deg=10) assert a1.__eq__(1) == NotImplemented diff --git a/tests/unit/common/test_data_collection_parameters.py b/tests/unit/common/test_data_collection_parameters.py index 41e5f887..15362a4e 100644 --- a/tests/unit/common/test_data_collection_parameters.py +++ b/tests/unit/common/test_data_collection_parameters.py @@ -1,8 +1,7 @@ import pytest +from aarecommon.models.models import DataCollectionParameters from pydantic import ValidationError -from aare.common.models import DataCollectionParameters - def test_directory_defaults_when_missing(): params = DataCollectionParameters() @@ -72,4 +71,4 @@ def test_datacollectionparameters_accepts_new_fields(): assert params.unitcell == "11,22,33,90,90,120" assert params.processingresolution == 1.2 assert params.pdbmodel == "model.pdb" - assert params.cloud is False \ No newline at end of file + assert params.cloud is False diff --git a/tests/unit/common/test_diffraction_geometry.py b/tests/unit/common/test_diffraction_geometry.py index 0051d67d..2cf3c378 100644 --- a/tests/unit/common/test_diffraction_geometry.py +++ b/tests/unit/common/test_diffraction_geometry.py @@ -1,5 +1,6 @@ import pytest -from aare.common.diffraction_geometry import DiffractionGeometry +from aarecommon.math.diffraction_geometry import DiffractionGeometry + @pytest.fixture def sample_dg(): @@ -12,9 +13,10 @@ def sample_dg(): detector_description="Eiger 16M", detector_serial_number="E-123", poni_rot1_rad=0.0, - poni_rot2_rad=0.0 + poni_rot2_rad=0.0, ) + def test_detector_max_radius_pxl(sample_dg): # center (1000, 1000), size (2000, 2000) # x0 = 2000-1000 = 1000 @@ -24,14 +26,17 @@ def test_detector_max_radius_pxl(sample_dg): # max = 1000 assert sample_dg.detector_max_radius_pxl == 1000.0 + def test_detector_radius_mm(sample_dg): # 1000 * 0.172 = 172.0 assert sample_dg.detector_radius_mm == 172.0 + def test_wavelength_angstrom(sample_dg): # 12.398 / 12.4 = 0.9998387... assert sample_dg.wavelength_angstrom == pytest.approx(0.9998387) + def test_resolution_angstrom(sample_dg): # dtz = 100 # radius = 172 @@ -40,20 +45,24 @@ def test_resolution_angstrom(sample_dg): res = sample_dg.resolution_angstrom(100.0) assert res > 0 assert res == pytest.approx(1.002469, abs=1e-5) - + with pytest.raises(ValueError): sample_dg.resolution_angstrom(0) + def test_max_resolution_angstrom(sample_dg): - assert sample_dg.max_resolution_angstrom == sample_dg.resolution_angstrom(sample_dg.dtz_mm) + assert sample_dg.max_resolution_angstrom == sample_dg.resolution_angstrom( + sample_dg.dtz_mm + ) + def test_calc_dtz_mm(sample_dg): res = sample_dg.resolution_angstrom(100.0) dtz = sample_dg.calc_dtz_mm(res) assert dtz == pytest.approx(100.0) - + with pytest.raises(ValueError): sample_dg.calc_dtz_mm(-1) - + # test x >= 1.0 case: wavelength / (2*res) >= 1.0 -> res <= wavelength / 2 assert sample_dg.calc_dtz_mm(0.0001) == 0.0 diff --git a/tests/unit/common/test_error_codes.py b/tests/unit/common/test_error_codes.py index 70545adb..80e41ca2 100644 --- a/tests/unit/common/test_error_codes.py +++ b/tests/unit/common/test_error_codes.py @@ -1,70 +1,112 @@ -from aare.common.error_codes import ( - AuthErrorCode, +from aarecommon.errors.codes import ( AareErrorCode, + AuthErrorCode, code_for_exception_class, error_code_help, export_error_codes, export_error_codes_grouped, ) + def test_error_code_help_returns_string(): help_text = error_code_help(AuthErrorCode.AUTHENTICATION_FAILED) assert isinstance(help_text, str) assert len(help_text) > 0 + def test_error_code_help_unknown_code(): help_text = error_code_help("UNKNOWN_CODE") assert help_text is None + def test_export_error_codes_contains_known_codes(): exported = export_error_codes() assert AuthErrorCode.AUTHENTICATION_FAILED.name in exported def test_code_for_exception_class_basic(): - assert code_for_exception_class("TellCommunicationError") == "TELL_COMMUNICATION_ERROR" + assert ( + code_for_exception_class("TellCommunicationError") == "TELL_COMMUNICATION_ERROR" + ) assert code_for_exception_class("MountingFailed") == "MOUNTING_FAILED" assert code_for_exception_class("LoopCenteringFailed") == "LOOP_CENTERING_FAILED" def test_code_for_exception_class_acronyms(): assert code_for_exception_class("AXCFailed") == "AXC_FAILED" - assert code_for_exception_class("AareDBCommunicationError") == "AARE_DB_COMMUNICATION_ERROR" - assert code_for_exception_class("JFJochCommunicationError") == "JF_JOCH_COMMUNICATION_ERROR" + assert ( + code_for_exception_class("AareDBCommunicationError") + == "AARE_DB_COMMUNICATION_ERROR" + ) + assert ( + code_for_exception_class("JFJochCommunicationError") + == "JF_JOCH_COMMUNICATION_ERROR" + ) def test_code_for_each_concrete_exception_is_in_aare_error_code_enum(): """Every code we'd produce from the rebuilt hierarchy must exist in the AareErrorCode enum -- otherwise clients have no symbol to branch on.""" - from aare.common.exception_handler import ( - TellCommunicationError, TellConnectionException, CriticalTellException, - WarningTellException, TellCommandWhileBusyException, - MountingFailed, UnmountingFailed, - SmargonCommunicationError, AerotechCommunicationError, - JFJochCommunicationError, BECCommunicationError, - AareDBCommunicationError, StateTransitionFailed, - MaintenanceStateException, BeamlineBusyException, - BeamlineBusyTimeoutException, DataCollectionException, - RasterScanException, LoopCenteringFailed, AXCFailed, - AutoRasterSampleSkipped, TransformationInvalidException, - MagnetPositionSensorErorr, SmartMagnetFaultException, - ManualMountException, SampleException, AuthenticationException, + from aarecommon.errors.exception_handler import ( + AareDBCommunicationError, + AerotechCommunicationError, + AuthenticationException, + AutoRasterSampleSkipped, + AXCFailed, + BeamlineBusyException, + BeamlineBusyTimeoutException, + BECCommunicationError, + CriticalTellException, + DataCollectionException, + JFJochCommunicationError, + LoopCenteringFailed, + MagnetPositionSensorErorr, + MaintenanceStateException, + ManualMountException, + MountingFailed, + RasterScanException, + SampleException, + SmargonCommunicationError, + SmartMagnetFaultException, + StateTransitionFailed, + TellCommandWhileBusyException, + TellCommunicationError, + TellConnectionException, + TransformationInvalidException, + UnmountingFailed, UserRightsException, + WarningTellException, ) + valid = {c.value for c in AareErrorCode} classes = [ - TellCommunicationError, TellConnectionException, CriticalTellException, - WarningTellException, TellCommandWhileBusyException, - MountingFailed, UnmountingFailed, - SmargonCommunicationError, AerotechCommunicationError, - JFJochCommunicationError, BECCommunicationError, - AareDBCommunicationError, StateTransitionFailed, - MaintenanceStateException, BeamlineBusyException, - BeamlineBusyTimeoutException, DataCollectionException, - RasterScanException, LoopCenteringFailed, AXCFailed, - AutoRasterSampleSkipped, TransformationInvalidException, - MagnetPositionSensorErorr, SmartMagnetFaultException, - ManualMountException, SampleException, AuthenticationException, + TellCommunicationError, + TellConnectionException, + CriticalTellException, + WarningTellException, + TellCommandWhileBusyException, + MountingFailed, + UnmountingFailed, + SmargonCommunicationError, + AerotechCommunicationError, + JFJochCommunicationError, + BECCommunicationError, + AareDBCommunicationError, + StateTransitionFailed, + MaintenanceStateException, + BeamlineBusyException, + BeamlineBusyTimeoutException, + DataCollectionException, + RasterScanException, + LoopCenteringFailed, + AXCFailed, + AutoRasterSampleSkipped, + TransformationInvalidException, + MagnetPositionSensorErorr, + SmartMagnetFaultException, + ManualMountException, + SampleException, + AuthenticationException, UserRightsException, ] missing = [] diff --git a/tests/unit/common/test_exception_handler.py b/tests/unit/common/test_exception_handler.py index bb23da5e..bbbac751 100644 --- a/tests/unit/common/test_exception_handler.py +++ b/tests/unit/common/test_exception_handler.py @@ -1,132 +1,158 @@ import pytest -from aare.common.exception_handler import ( - DataCollectionException, - AuthenticationException, - AuthErrorCode, - TellCommunicationError, - TransformationInvalidException, - RasterScanException, - LoopCenteringFailed, - UnmountingFailed, - MountingFailed, - ManualMountException, - SmartMagnetFaultException, - TellCommandWhileBusyException, - TellConnectionException, - WarningTellException, - CriticalTellException, - AXCFailed, - BeamlineBusyException, - SampleException, - UserRightsException, - SmargonCommunicationError, - JFJochCommunicationError, +from aarecommon.errors.exception_handler import ( AareDBCommunicationError, AerotechCommunicationError, - MagnetPositionSensorErorr + AuthenticationException, + AuthErrorCode, + AXCFailed, + BeamlineBusyException, + CriticalTellException, + DataCollectionException, + JFJochCommunicationError, + LoopCenteringFailed, + MagnetPositionSensorErorr, + ManualMountException, + MountingFailed, + RasterScanException, + SampleException, + SmargonCommunicationError, + SmartMagnetFaultException, + TellCommandWhileBusyException, + TellCommunicationError, + TellConnectionException, + TransformationInvalidException, + UnmountingFailed, + UserRightsException, + WarningTellException, ) + def test_data_collection_exception_message(): exc = DataCollectionException("Custom error") assert str(exc) == "Custom error" + def test_data_collection_exception_default_message(): exc = DataCollectionException() assert str(exc) == "Data collection failed" + def test_authentication_exception_properties(): - exc = AuthenticationException("Failed", status_code=403, code=AuthErrorCode.FORBIDDEN) + exc = AuthenticationException( + "Failed", status_code=403, code=AuthErrorCode.FORBIDDEN + ) assert exc.status_code == 403 assert exc.code == AuthErrorCode.FORBIDDEN assert "Failed" in str(exc) + def test_tell_communication_error_str(): exc = TellCommunicationError("Timeout", endpoint="/state", operation="GET") assert str(exc) == "Timeout" assert exc.endpoint == "/state" assert exc.operation == "GET" + def test_transformation_invalid_exception(): exc = TransformationInvalidException() assert "Transformation is not implemented" in str(exc) + def test_raster_scan_exception(): exc = RasterScanException("Raster failed") assert "Raster failed" in str(exc) + def test_loop_centering_failed(): exc = LoopCenteringFailed() assert "Loop Centering did not detect a sample" in str(exc) + def test_unmounting_failed(): exc = UnmountingFailed() assert "A sample was not unmounted" in str(exc) + def test_mounting_failed(): exc = MountingFailed() assert "A sample was not mounted" in str(exc) + def test_manual_mount_exception(): exc = ManualMountException() assert "Manual mounting failed" in str(exc) + def test_smart_magnet_fault_exception(): exc = SmartMagnetFaultException() assert "Smart magnet fault" in str(exc) + def test_tell_command_while_busy_exception(): exc = TellCommandWhileBusyException() assert "Tell is busy" in str(exc) + def test_tell_connection_exception(): exc = TellConnectionException() assert "Lost connection to Tell" in str(exc) + def test_warning_tell_exception(): exc = WarningTellException() assert "Warning error in TELL" in str(exc) + def test_critical_tell_exception(): exc = CriticalTellException() assert "Critical error in TELL" in str(exc) + def test_axc_failed(): exc = AXCFailed() assert "Auto X-ray centering failed" in str(exc) + def test_beamline_busy_exception(): exc = BeamlineBusyException() assert "Beamline is in busy state" in str(exc) + def test_sample_exception(): exc = SampleException() assert "Sample not found" in str(exc) + def test_user_rights_exception(): exc = UserRightsException(code=AuthErrorCode.NOT_STAFF) assert exc.status_code == 403 assert exc.code == AuthErrorCode.NOT_STAFF assert "User does not have rights" in str(exc) + def test_smargon_communication_error(): exc = SmargonCommunicationError("Conn error", endpoint="/move", status_code=500) assert str(exc) == "Conn error" assert exc.endpoint == "/move" assert exc.status_code == 500 + def test_jfjoch_communication_error(): exc = JFJochCommunicationError("JFJoch error", operation="POST") assert str(exc) == "JFJoch error" assert exc.operation == "POST" + def test_aaredb_communication_error(): exc = AareDBCommunicationError("DB error") assert "DB error" in str(exc) + def test_aerotech_communication_error(): exc = AerotechCommunicationError("Aerotech error") assert "Aerotech error" in str(exc) + def test_magnet_position_sensor_error(): exc = MagnetPositionSensorErorr() assert "Magnet position sensor error" in str(exc) diff --git a/tests/unit/common/test_find_xtal.py b/tests/unit/common/test_find_xtal.py index d7cd264c..fab48d35 100644 --- a/tests/unit/common/test_find_xtal.py +++ b/tests/unit/common/test_find_xtal.py @@ -1,14 +1,22 @@ -import pytest -import numpy as np from unittest.mock import MagicMock -from aare.common.find_xtal import ( - identify_crystal_raster, rebuild_array_from_scan_results, - create_quality_filtered_array, get_xtal_size, get_best_b_factor, - get_best_res, raster_centre_of_mass, has_sufficient_low_res_spots, - compute_crystal_score_array, raster_highest_score + +import numpy as np +import pytest +from aarecommon.math.find_xtal import ( + compute_crystal_score_array, + create_quality_filtered_array, + get_best_b_factor, + get_best_res, + get_xtal_size, + has_sufficient_low_res_spots, + identify_crystal_raster, + raster_centre_of_mass, + raster_highest_score, + rebuild_array_from_scan_results, ) -from aare.common.raster_grid import RasterGridRequest, CenterOfMassModel -from aare.common.models import CrystalSize, Coordinate +from aarecommon.models.models import Coordinate, CrystalSize +from aarecommon.models.raster_grid import CenterOfMassModel, RasterGridRequest + @pytest.fixture def mock_raster_results(): @@ -29,6 +37,7 @@ def mock_raster_results(): results.append(res) return results + @pytest.fixture def raster_request(): return RasterGridRequest( @@ -36,13 +45,14 @@ def raster_request(): n_x=3, n_y=3, grid_size_mm=Coordinate(x=0.02, y=0.02), - smargon_top_left=None + smargon_top_left=None, ) + def test_identify_crystal_raster(mock_raster_results, raster_request): mock_result = MagicMock() mock_result.images = mock_raster_results - + com = identify_crystal_raster(mock_result, raster_request) assert isinstance(com, CenterOfMassModel) # The max spots_low_res is 8.0 at (2, 2) @@ -50,40 +60,54 @@ def test_identify_crystal_raster(mock_raster_results, raster_request): assert com.n_y == 2 assert com.max_image == 8 + def test_rebuild_array_from_scan_results(mock_raster_results): - arr = rebuild_array_from_scan_results(mock_raster_results, "spots_low_res", array_shape=(3, 3)) + arr = rebuild_array_from_scan_results( + mock_raster_results, "spots_low_res", array_shape=(3, 3) + ) assert arr.shape == (3, 3) assert arr[0, 0] == 0.0 assert arr[2, 2] == 8.0 + def test_create_quality_filtered_array(mock_raster_results): # Testing create_quality_filtered_array with a more lenient filter - arr = create_quality_filtered_array(mock_raster_results, "spots", min_low_res_spots=0.0, min_background=0.0, array_shape=(3, 3)) + arr = create_quality_filtered_array( + mock_raster_results, + "spots", + min_low_res_spots=0.0, + min_background=0.0, + array_shape=(3, 3), + ) # i=8 has nx=2, ny=2 -> row=2, col=2 assert arr[2, 2] == 18.0 + def test_get_xtal_size(raster_request): # 3x3 array where only center is 1 arr = np.zeros((3, 3)) arr[1, 1] = 1 - + size = CrystalSize(x=0, y=0, z=0) new_size = get_xtal_size(size, arr, raster_request) # row_max-row_min = 0, so size 0? Wait, it should probably be at least 1 grid unit if present. # The code does (row_max - row_min) * grid_size_mm.x * 1000 - # If row_min=1, row_max=1, then size is 0. + # If row_min=1, row_max=1, then size is 0. # This might be a bug in find_xtal.py if it doesn't account for the pixel itself. - assert new_size.x == 0 + assert new_size.x == 0 + def test_get_best_b_factor(mock_raster_results): assert get_best_b_factor(mock_raster_results) == 20.0 assert get_best_b_factor([]) is None + def test_get_best_res(mock_raster_results): # min res. i=8 -> res = 2.0 - 0.8 = 1.2 assert pytest.approx(get_best_res(mock_raster_results)) == 1.2 assert get_best_res([]) is None + def test_raster_centre_of_mass(mock_raster_results): arr = np.zeros((3, 3)) arr[1, 1] = 10.0 @@ -91,12 +115,14 @@ def test_raster_centre_of_mass(mock_raster_results): assert com.n_x == 1.0 assert com.n_y == 1.0 + def test_has_sufficient_low_res_spots(): arr = np.array([[1, 2], [3, 4]]) assert has_sufficient_low_res_spots(arr, 3.0) is True assert has_sufficient_low_res_spots(arr, 5.0) is False assert has_sufficient_low_res_spots(None, 1.0) is False + @pytest.fixture def mock_score_results(): # 2x2 grid; cell (1, 1) is the strongest crystal across bkg, spots_low_res @@ -116,6 +142,7 @@ def mock_score_results(): results.append(res) return results + def test_compute_crystal_score_array(mock_score_results): score = compute_crystal_score_array(mock_score_results) assert score.shape == (2, 2) @@ -126,6 +153,7 @@ def test_compute_crystal_score_array(mock_score_results): assert score[0, 0] == pytest.approx(0.0) assert np.unravel_index(np.argmax(score), score.shape) == (1, 1) + def test_compute_crystal_score_array_weights(): # Cell A is the sole max in spots_low_res (weight 0.55); cell B is the sole # max in spots_indexed (weight 0.20). A must outscore B. @@ -134,11 +162,15 @@ def test_compute_crystal_score_array_weights(): r.nx, r.ny, r.number = nx, 0, n r.bkg, r.spots_low_res, r.spots_indexed = 0.0, low, idx return r - score = compute_crystal_score_array([_cell(0, 0, 10.0, 0.0), _cell(1, 1, 0.0, 10.0)]) + + score = compute_crystal_score_array( + [_cell(0, 0, 10.0, 0.0), _cell(1, 1, 0.0, 10.0)] + ) assert score[0, 0] > score[1, 0] assert score[0, 0] == pytest.approx(55.0) assert score[1, 0] == pytest.approx(20.0) + def test_raster_highest_score(mock_score_results): com = raster_highest_score(mock_score_results) assert com.n_x == 1.0 diff --git a/tests/unit/common/test_logger_events.py b/tests/unit/common/test_logger_events.py index e5e22b4e..c575b3ff 100644 --- a/tests/unit/common/test_logger_events.py +++ b/tests/unit/common/test_logger_events.py @@ -1,19 +1,21 @@ import logging import time -import pytest from unittest.mock import MagicMock -from aare.common.logger_events import log_timing, merge_log_context + +import pytest +from aarecommon.config.logger_events import log_timing, merge_log_context + def test_log_timing_success(): logger = MagicMock(spec=logging.Logger) - + @log_timing(logger, message_prefix="Test", level=logging.INFO) def sample_func(x): time.sleep(0.01) return x * 2 result = sample_func(21) - + assert result == 42 assert logger.log.call_count == 2 # First call: Starting @@ -25,9 +27,10 @@ def test_log_timing_success(): assert "duration_s" in kwargs["extra"] assert kwargs["extra"]["duration_s"] >= 0.01 + def test_log_timing_failure(): logger = MagicMock(spec=logging.Logger) - + @log_timing(logger, level=logging.ERROR) def failing_func(): time.sleep(0.01) @@ -35,7 +38,7 @@ def test_log_timing_failure(): with pytest.raises(ValueError, match="Something went wrong"): failing_func() - + assert logger.log.call_count == 2 # First call: Starting logger.log.assert_any_call(logging.ERROR, "Starting failing_func") @@ -46,11 +49,12 @@ def test_log_timing_failure(): assert "Something went wrong" in args[1] assert "duration_s" in kwargs["extra"] + def test_merge_log_context(): ctx1 = {"a": 1, "b": 2} ctx2 = {"b": 3, "c": 4} merged = merge_log_context(ctx1, ctx2, d=5) assert merged == {"a": 1, "b": 3, "c": 4, "d": 5} - + merged_none = merge_log_context(ctx1, None) assert merged_none == {"a": 1, "b": 2} diff --git a/tests/unit/common/test_mlbox_model.py b/tests/unit/common/test_mlbox_model.py index 0ae92607..5d720700 100644 --- a/tests/unit/common/test_mlbox_model.py +++ b/tests/unit/common/test_mlbox_model.py @@ -1,4 +1,4 @@ -from aare.common.models import MLOutputModel, MLBoxType +from aarecommon.models.models import MLBoxType, MLOutputModel def test_add_box_generates_unique_keys(): @@ -20,6 +20,7 @@ def test_get_best_for_class_returns_highest_confidence(): assert best is not None assert best.conf == 0.7 + def test_get_best_for_class_returns_none_when_missing(): model = MLOutputModel() - assert model.get_best_for_class(MLBoxType.CRYSTAL) is None \ No newline at end of file + assert model.get_best_for_class(MLBoxType.CRYSTAL) is None diff --git a/tests/unit/common/test_models_extra.py b/tests/unit/common/test_models_extra.py index 7c9557b8..2d683ac8 100644 --- a/tests/unit/common/test_models_extra.py +++ b/tests/unit/common/test_models_extra.py @@ -1,12 +1,12 @@ import pytest -from aare.common.models import ( - SampleShortInfo, - DewarAddress, - BeamMarkCoeffModel, - MLOutputModel, - MLBoxType, +from aarecommon.models.models import ( BeamlineStateEnum, + BeamMarkCoeffModel, DataCollectionParameters, + DewarAddress, + MLBoxType, + MLOutputModel, + SampleShortInfo, ) @@ -55,7 +55,7 @@ def test_sample_short_info_methods(): }, "user": "group2", "pin": 4, - "location": {"segment": "B", "pos": 5} + "location": {"segment": "B", "pos": 5}, } info2 = SampleShortInfo.from_dict(data) assert info2.db_id == 2 @@ -67,10 +67,7 @@ def test_sample_short_info_methods(): def test_beam_mark_coeff_model_apply(): - model = BeamMarkCoeffModel( - coeff_x=(1.0, 2.0, 5.0), - coeff_y=(3.0, 4.0, 6.0) - ) + model = BeamMarkCoeffModel(coeff_x=(1.0, 2.0, 5.0), coeff_y=(3.0, 4.0, 6.0)) res = model.apply(10.0) assert res.x == 125.0 assert res.y == 346.0 @@ -109,4 +106,4 @@ def test_ml_output_model_extra_methods(): def test_beamline_state_enum_display_name(): assert BeamlineStateEnum.SampleExchange.display_name() == "Sample exchange" assert BeamlineStateEnum.Moving.display_name() == "Moving" - assert BeamlineStateEnum.display_name(None) == "-" \ No newline at end of file + assert BeamlineStateEnum.display_name(None) == "-" diff --git a/tests/unit/common/test_raster_grid_common.py b/tests/unit/common/test_raster_grid_common.py index 9244122c..38d7e5a1 100644 --- a/tests/unit/common/test_raster_grid_common.py +++ b/tests/unit/common/test_raster_grid_common.py @@ -1,6 +1,5 @@ import pytest - -from aare.common.raster_grid import grid_to_image_id, image_id_to_grid +from aarecommon.math.raster_grid import grid_to_image_id, image_id_to_grid def test_grid_to_image_id_single_cell(): @@ -84,7 +83,9 @@ def test_grid_to_image_id_rejects_zero_grid_y(): def test_grid_to_image_id_rejects_grid_x_larger_than_number_of_cols(): - with pytest.raises(ValueError, match="grid_x cannot be greater than number_of_cols"): + with pytest.raises( + ValueError, match="grid_x cannot be greater than number_of_cols" + ): grid_to_image_id(grid_x=6, grid_y=1, number_of_cols=5) diff --git a/tests/unit/common/test_recurrence_watcher.py b/tests/unit/common/test_recurrence_watcher.py index 6711c1c8..9aa360c5 100644 --- a/tests/unit/common/test_recurrence_watcher.py +++ b/tests/unit/common/test_recurrence_watcher.py @@ -1,10 +1,10 @@ -from aare.common.exception_handler import ( +from aarecommon.errors.exception_handler import ( LoopCenteringFailed, MountingFailed, SampleException, TellException, ) -from aare.common.recurrence_watcher import ( +from aarecommon.recurrence_watcher import ( RecurrenceWatcher, create_default_watchers, load_watcher_threshold_overrides, @@ -47,7 +47,9 @@ def test_load_watcher_threshold_overrides_reads_valid_values_only(): def test_create_default_watchers_applies_overrides(): - watchers = {watcher.name: watcher for watcher in create_default_watchers({"alc": 4})} + watchers = { + watcher.name: watcher for watcher in create_default_watchers({"alc": 4}) + } assert watchers["alc"].threshold == 4 assert watchers["tell"].threshold > 0 @@ -79,4 +81,4 @@ def test_load_watcher_threshold_overrides_through_env_var_adapter(monkeypatch): get_value=lambda k: os.getenv(redis_key_to_env_var(k)), ) - assert overrides == {"tell": 9} \ No newline at end of file + assert overrides == {"tell": 9} diff --git a/tests/unit/common/test_zoom_model_camera_settings.py b/tests/unit/common/test_zoom_model_camera_settings.py index b02f9fdd..f5c794ab 100644 --- a/tests/unit/common/test_zoom_model_camera_settings.py +++ b/tests/unit/common/test_zoom_model_camera_settings.py @@ -1,6 +1,5 @@ import pytest - -from aare.common.models import ZoomModel, SampleCameraSettings +from aarecommon.models.models import SampleCameraSettings, ZoomModel def test_get_camera_settings_exact_match(): @@ -34,6 +33,7 @@ def test_get_camera_settings_raises_when_empty(): with pytest.raises(ValueError): model.get_camera_settings(100) + def test_get_camera_settings_below_minimum_returns_first(): model = ZoomModel( z={ @@ -44,4 +44,4 @@ def test_get_camera_settings_below_minimum_returns_first(): settings = model.get_camera_settings(50) assert settings.gain == 1.0 - assert settings.exposure == 0.1 \ No newline at end of file + assert settings.exposure == 0.1 diff --git a/tests/unit/daq/operations/face_detection/test_face_detection_service.py b/tests/unit/daq/operations/face_detection/test_face_detection_service.py index b06c012c..d16e30a7 100644 --- a/tests/unit/daq/operations/face_detection/test_face_detection_service.py +++ b/tests/unit/daq/operations/face_detection/test_face_detection_service.py @@ -1,9 +1,19 @@ import types import pytest +from aarecommon.math.coordinate import Coordinate, SmargonCoordinate +from aarecommon.models.models import ( + BoundingBoxModel, + MLBoxModel, + MLBoxType, + ZoomModeEnum, +) -from aare.common.coordinate import Coordinate, SmargonCoordinate -from aare.common.models import BoundingBoxModel, MLBoxModel, MLBoxType, ZoomModeEnum +from aare.daq.operations.common.runtime import DAQRuntimeState +from aare.daq.operations.common.services import ( + FaceDetectionProgressEmitter, + OperationServices, +) from aare.daq.operations.face_detection.models import ( FaceDetectionContext, FaceDetectionDependencies, @@ -11,11 +21,7 @@ from aare.daq.operations.face_detection.models import ( FaceDetectionSettings, ) from aare.daq.operations.face_detection.service import FaceDetectionService -from aare.daq.operations.common.runtime import DAQRuntimeState -from aare.daq.operations.common.services import ( - FaceDetectionProgressEmitter, - OperationServices, -) + class DummyGeometry: def __init__(self): @@ -63,9 +69,7 @@ def context(): ) services = OperationServices( - screenshots=types.SimpleNamespace( - save_to_db=lambda *args, **kwargs: None - ), + screenshots=types.SimpleNamespace(save_to_db=lambda *args, **kwargs: None), face_detection_progress=FaceDetectionProgressEmitter( reporter=types.SimpleNamespace(emit_progress=emit_progress) ), @@ -117,19 +121,39 @@ def test_service_returns_empty_payload_when_no_boxes(monkeypatch, context, mock_ assert context._progress_events[-1]["running"] is False -def test_service_prefers_face_boxes_when_ratio_is_high(monkeypatch, context, mock_logger): +def test_service_prefers_face_boxes_when_ratio_is_high( + monkeypatch, context, mock_logger +): service = FaceDetectionService(context=context, logger=mock_logger) predictions = iter( [ - types.SimpleNamespace(box=_box(MLBoxType.LOOP_FACE, 10, 20, 30, 40), target_point=None, focus=None), - types.SimpleNamespace(box=_box(MLBoxType.LOOP_FACE, 12, 20, 32, 40), target_point=None, focus=None), - types.SimpleNamespace(box=_box(MLBoxType.LOOP_ALL, 14, 20, 34, 40), target_point=None, focus=None), - types.SimpleNamespace(box=_box(MLBoxType.LOOP_FACE, 16, 20, 36, 40), target_point=None, focus=None), + types.SimpleNamespace( + box=_box(MLBoxType.LOOP_FACE, 10, 20, 30, 40), + target_point=None, + focus=None, + ), + types.SimpleNamespace( + box=_box(MLBoxType.LOOP_FACE, 12, 20, 32, 40), + target_point=None, + focus=None, + ), + types.SimpleNamespace( + box=_box(MLBoxType.LOOP_ALL, 14, 20, 34, 40), + target_point=None, + focus=None, + ), + types.SimpleNamespace( + box=_box(MLBoxType.LOOP_FACE, 16, 20, 36, 40), + target_point=None, + focus=None, + ), ] ) - monkeypatch.setattr(context.deps.mlbox, "predict", lambda **kwargs: next(predictions)) + monkeypatch.setattr( + context.deps.mlbox, "predict", lambda **kwargs: next(predictions) + ) monkeypatch.setattr( "aare.daq.operations.face_detection.service.fd.get_flat_face", @@ -144,7 +168,9 @@ def test_service_prefers_face_boxes_when_ratio_is_high(monkeypatch, context, moc ) monkeypatch.setattr( "aare.daq.operations.face_detection.service.fd.get_samples_out", - lambda boxes: [{"angle": angle, "box": box} for angle, box in sorted(boxes.items())], + lambda boxes: [ + {"angle": angle, "box": box} for angle, box in sorted(boxes.items()) + ], ) result = service.run(steps=3, step_size=30, face_min_ratio=0.5) @@ -158,20 +184,44 @@ def test_service_prefers_face_boxes_when_ratio_is_high(monkeypatch, context, moc assert context._progress_events[-1]["running"] is False -def test_service_falls_back_to_loop_all_when_face_ratio_is_low(monkeypatch, context, mock_logger): +def test_service_falls_back_to_loop_all_when_face_ratio_is_low( + monkeypatch, context, mock_logger +): service = FaceDetectionService(context=context, logger=mock_logger) predictions = iter( [ - types.SimpleNamespace(box=_box(MLBoxType.LOOP_FACE, 10, 20, 30, 40), target_point=None, focus=None), - types.SimpleNamespace(box=_box(MLBoxType.LOOP_ALL, 11, 20, 31, 40), target_point=None, focus=None), - types.SimpleNamespace(box=_box(MLBoxType.LOOP_ALL, 12, 20, 32, 40), target_point=None, focus=None), - types.SimpleNamespace(box=_box(MLBoxType.LOOP_ALL, 13, 20, 33, 40), target_point=None, focus=None), - types.SimpleNamespace(box=_box(MLBoxType.LOOP_ALL, 14, 20, 34, 40), target_point=None, focus=None), + types.SimpleNamespace( + box=_box(MLBoxType.LOOP_FACE, 10, 20, 30, 40), + target_point=None, + focus=None, + ), + types.SimpleNamespace( + box=_box(MLBoxType.LOOP_ALL, 11, 20, 31, 40), + target_point=None, + focus=None, + ), + types.SimpleNamespace( + box=_box(MLBoxType.LOOP_ALL, 12, 20, 32, 40), + target_point=None, + focus=None, + ), + types.SimpleNamespace( + box=_box(MLBoxType.LOOP_ALL, 13, 20, 33, 40), + target_point=None, + focus=None, + ), + types.SimpleNamespace( + box=_box(MLBoxType.LOOP_ALL, 14, 20, 34, 40), + target_point=None, + focus=None, + ), ] ) - monkeypatch.setattr(context.deps.mlbox, "predict", lambda **kwargs: next(predictions)) + monkeypatch.setattr( + context.deps.mlbox, "predict", lambda **kwargs: next(predictions) + ) captured = {} @@ -179,7 +229,10 @@ def test_service_falls_back_to_loop_all_when_face_ratio_is_low(monkeypatch, cont captured["boxes_used"] = dict(boxes) return 60, {"A": 1.0, "B": 2.0, "phi_rad": 0.2, "C": 4.0} - monkeypatch.setattr("aare.daq.operations.face_detection.service.fd.get_flat_face", fake_get_flat_face) + monkeypatch.setattr( + "aare.daq.operations.face_detection.service.fd.get_flat_face", + fake_get_flat_face, + ) monkeypatch.setattr( "aare.daq.operations.face_detection.service.fd.choose_best_fit", lambda fit_results: (60, {"A": 1.0}, "Height"), @@ -198,7 +251,9 @@ def test_service_falls_back_to_loop_all_when_face_ratio_is_low(monkeypatch, cont assert context._progress_events[-1]["running"] is False -def test_service_applies_centre_correction_when_target_is_far_from_beam(monkeypatch, context, mock_logger): +def test_service_applies_centre_correction_when_target_is_far_from_beam( + monkeypatch, context, mock_logger +): service = FaceDetectionService(context=context, logger=mock_logger) model = _box(MLBoxType.LOOP_FACE, 40, 160, 80, 200) @@ -210,7 +265,9 @@ def test_service_applies_centre_correction_when_target_is_far_from_beam(monkeypa assert context.deps.devs.smargon_pos.sh_mm.y == pytest.approx(1.8) -def test_service_returns_failed_result_when_prediction_raises(monkeypatch, context, mock_logger): +def test_service_returns_failed_result_when_prediction_raises( + monkeypatch, context, mock_logger +): service = FaceDetectionService(context=context, logger=mock_logger) monkeypatch.setattr( @@ -226,4 +283,4 @@ def test_service_returns_failed_result_when_prediction_raises(monkeypatch, conte assert result.comment == "Face detection sequence failed" assert result.payload["running"] is False assert context.deps.cfg.zoom_mode == ZoomModeEnum.User - assert context._progress_events[-1]["running"] is False \ No newline at end of file + assert context._progress_events[-1]["running"] is False diff --git a/tests/unit/daq/operations/loop_centering/test_loop_centering_analyzer.py b/tests/unit/daq/operations/loop_centering/test_loop_centering_analyzer.py index 76ac5684..49cd3aee 100644 --- a/tests/unit/daq/operations/loop_centering/test_loop_centering_analyzer.py +++ b/tests/unit/daq/operations/loop_centering/test_loop_centering_analyzer.py @@ -1,15 +1,15 @@ import types import pytest +from aarecommon.math.coordinate import Coordinate, SmargonCoordinate +from aarecommon.models.models import MLBoxType, MLOutputModel -from aare.common.coordinate import Coordinate, SmargonCoordinate -from aare.common.models import MLBoxType, MLOutputModel from aare.daq.mlbox import MLBoxPredictionsResult from aare.daq.operations.common.runtime import DAQRuntimeState from aare.daq.operations.common.services import ( - TraceWriter, - PredictionProvider, OperationServices, + PredictionProvider, + TraceWriter, ) from aare.daq.operations.loop_centering.analyzer import LoopCenteringAnalyzer from aare.daq.operations.loop_centering.models import ( @@ -56,11 +56,11 @@ def analyzer(mock_logger): ), runtime=runtime_state, services=OperationServices( - screenshots=types.SimpleNamespace( - save_to_db=lambda *args, **kwargs: None - ), + screenshots=types.SimpleNamespace(save_to_db=lambda *args, **kwargs: None), traces=TraceWriter( - appender=types.SimpleNamespace(append_smargon_trace=lambda *args, **kwargs: None) + appender=types.SimpleNamespace( + append_smargon_trace=lambda *args, **kwargs: None + ) ), predictions=PredictionProvider( getter=types.SimpleNamespace( @@ -79,32 +79,40 @@ def analyzer(mock_logger): def test_is_ignore_only_classes_true(analyzer): - assert analyzer.is_ignore_only_classes([ - MLBoxType.PIN.value, - MLBoxType.ICE.value, - MLBoxType.NEEDLE.value, - ]) + assert analyzer.is_ignore_only_classes( + [ + MLBoxType.PIN.value, + MLBoxType.ICE.value, + MLBoxType.NEEDLE.value, + ] + ) def test_is_ignore_only_classes_false_when_loop_present(analyzer): - assert not analyzer.is_ignore_only_classes([ - MLBoxType.PIN.value, - MLBoxType.LOOP_FACE.value, - ]) + assert not analyzer.is_ignore_only_classes( + [ + MLBoxType.PIN.value, + MLBoxType.LOOP_FACE.value, + ] + ) def test_has_valid_target_classes_true(analyzer): - assert analyzer.has_valid_target_classes([ - MLBoxType.PIN.value, - MLBoxType.CRYSTAL.value, - ]) + assert analyzer.has_valid_target_classes( + [ + MLBoxType.PIN.value, + MLBoxType.CRYSTAL.value, + ] + ) def test_has_valid_target_classes_false(analyzer): - assert not analyzer.has_valid_target_classes([ - MLBoxType.PIN.value, - MLBoxType.ICE.value, - ]) + assert not analyzer.has_valid_target_classes( + [ + MLBoxType.PIN.value, + MLBoxType.ICE.value, + ] + ) def test_select_smargon_target_prefers_prediction_when_within_tolerance(analyzer): @@ -150,11 +158,13 @@ def test_analyze_angle_with_crystal_prediction_returns_valid_target(analyzer): boxes = MLOutputModel() boxes.add_box(MLBoxType.CRYSTAL, (10, 20, 30, 40), conf=0.9) - analyzer.ctx.services.predictions.getter.get_predictions = lambda: MLBoxPredictionsResult( - predictions=boxes, - image=None, - target_point=(25.0, 30.0), - focus=1.0, + analyzer.ctx.services.predictions.getter.get_predictions = lambda: ( + MLBoxPredictionsResult( + predictions=boxes, + image=None, + target_point=(25.0, 30.0), + focus=1.0, + ) ) analysis = analyzer.analyze_angle(angle_deg=90, zoom_value=200.0, sample_id=1) @@ -165,4 +175,4 @@ def test_analyze_angle_with_crystal_prediction_returns_valid_target(analyzer): assert analysis.ignore_only is False assert analysis.calculated_target is not None assert analysis.predicted_target is not None - assert analysis.final_target is not None \ No newline at end of file + assert analysis.final_target is not None diff --git a/tests/unit/daq/operations/loop_centering/test_loop_centering_service.py b/tests/unit/daq/operations/loop_centering/test_loop_centering_service.py index 41588c87..61b61385 100644 --- a/tests/unit/daq/operations/loop_centering/test_loop_centering_service.py +++ b/tests/unit/daq/operations/loop_centering/test_loop_centering_service.py @@ -1,13 +1,13 @@ import types import pytest +from aarecommon.models.models import LoopCenteringResult -from aare.common.models import LoopCenteringResult from aare.daq.operations.common.runtime import DAQRuntimeState from aare.daq.operations.common.services import ( OperationServices, - TraceWriter, PredictionProvider, + TraceWriter, ) from aare.daq.operations.loop_centering.models import ( AngleAnalysis, @@ -32,7 +32,9 @@ def context(): runtime_state = DAQRuntimeState( sample_provider=types.SimpleNamespace(sample=None), - sample_geometry_provider=types.SimpleNamespace(sample_geometry=types.SimpleNamespace()), + sample_geometry_provider=types.SimpleNamespace( + sample_geometry=types.SimpleNamespace() + ), status_provider=types.SimpleNamespace(status=None), ) @@ -45,7 +47,9 @@ def context(): runtime=runtime_state, services=OperationServices( screenshots=types.SimpleNamespace( - save_to_db=lambda sample_id, filename, wait=0.0: screenshots.append((sample_id, filename, wait)) + save_to_db=lambda sample_id, filename, wait=0.0: screenshots.append( + (sample_id, filename, wait) + ) ), traces=TraceWriter( appender=types.SimpleNamespace( @@ -63,7 +67,9 @@ def context(): return ctx -def test_service_fails_when_attempt_one_has_no_valid_targets(monkeypatch, context, mock_logger): +def test_service_fails_when_attempt_one_has_no_valid_targets( + monkeypatch, context, mock_logger +): service = LoopCenteringService(context=context, logger=mock_logger) monkeypatch.setattr( @@ -84,7 +90,9 @@ def test_service_fails_when_attempt_one_has_no_valid_targets(monkeypatch, contex assert "attempt 1" in result.comment.lower() -def test_service_succeeds_when_correction_pass_has_valid_target(monkeypatch, context, mock_logger): +def test_service_succeeds_when_correction_pass_has_valid_target( + monkeypatch, context, mock_logger +): service = LoopCenteringService(context=context, logger=mock_logger) call_count = {"count": 0} @@ -106,7 +114,8 @@ def test_service_succeeds_when_correction_pass_has_valid_target(monkeypatch, con moved=True, ) - from aare.common.models import MLBoxType + from aarecommon.models.models import MLBoxType + monkeypatch.setattr(service, "_run_angle_pass", fake_run_angle_pass) result = service.run(sample_id=2) @@ -135,4 +144,4 @@ def test_service_uses_settings_values(monkeypatch, context, mock_logger): result = service.run(sample_id=3) assert context.deps.devs.zoom == 250.0 - assert result.success is False \ No newline at end of file + assert result.success is False diff --git a/tests/unit/daq/operations/mounting/test_mounting_service.py b/tests/unit/daq/operations/mounting/test_mounting_service.py index ca0d417d..0d8fdb9f 100644 --- a/tests/unit/daq/operations/mounting/test_mounting_service.py +++ b/tests/unit/daq/operations/mounting/test_mounting_service.py @@ -1,13 +1,22 @@ -import types import sys +import types -from aare.common.coordinate import AerotechCoordinate, Coordinate -from aare.common.exception_handler import CriticalTellException, DoorSafetyError, MountingFailed -from aare.common.models import DewarAddress, SampleShortInfo -from aare.devices.tell_client import TellEventValueEnum -from aare.daq.operations.mounting.models import MountingContext, MountingResult, MountingDependencies, MountingSettings +from aarecommon.errors.exception_handler import ( + CriticalTellException, + DoorSafetyError, + MountingFailed, +) +from aarecommon.math.coordinate import AerotechCoordinate, Coordinate +from aarecommon.models.models import DewarAddress, SampleShortInfo + +from aare.daq.operations.mounting.models import ( + MountingContext, + MountingDependencies, + MountingResult, + MountingSettings, +) from aare.daq.operations.mounting.service import MountingService - +from aare.devices.tell_client import TellEventValueEnum if "jfjoch_client.models.scan_result" not in sys.modules: jfjoch_client_mod = types.ModuleType("jfjoch_client") @@ -62,9 +71,13 @@ def _make_context(previous_sample=None, *, prohibited=True, alarm=False): cfg = types.SimpleNamespace( current_sample=previous_sample, get_mount_failure_streak=lambda: streak["count"], - increment_mount_failure_streak=lambda: streak.__setitem__("count", streak["count"] + 1) or streak["count"], + increment_mount_failure_streak=lambda: ( + streak.__setitem__("count", streak["count"] + 1) or streak["count"] + ), reset_mount_failure_streak=lambda: streak.__setitem__("count", 0), - record_mount_failure=lambda: streak.__setitem__("count", streak["count"] + 1) or streak["count"], + record_mount_failure=lambda: ( + streak.__setitem__("count", streak["count"] + 1) or streak["count"] + ), record_mount_success=lambda: streak.__setitem__("count", 0), ) @@ -74,7 +87,9 @@ def _make_context(previous_sample=None, *, prohibited=True, alarm=False): devs=devs, ), settings=MountingSettings( - mount_position=AerotechCoordinate(at_mm=Coordinate(x=0, y=0, z=0), omega_deg=0), + mount_position=AerotechCoordinate( + at_mm=Coordinate(x=0, y=0, z=0), omega_deg=0 + ), ), ) @@ -214,7 +229,9 @@ def test_execute_mount_success_resets_failure_streak(mock_logger): assert ctx.deps.cfg.get_mount_failure_streak() == 0 -def test_execute_mount_critical_tell_error_does_not_increment_failure_streak(mock_logger): +def test_execute_mount_critical_tell_error_does_not_increment_failure_streak( + mock_logger, +): target_sample = _make_sample(2, "new") ctx = _make_context() @@ -231,6 +248,7 @@ def test_execute_mount_critical_tell_error_does_not_increment_failure_streak(moc assert isinstance(result.error, CriticalTellException) assert ctx.deps.cfg.get_mount_failure_streak() == 0 + def test_mount_blocked_when_not_prohibited(mock_logger): """Doors open / hutch not in prohibited state -> critical DoorSafetyError, and it is not counted as a mount-failure streak.""" diff --git a/tests/unit/daq/operations/screenshot/test_screenshot_service.py b/tests/unit/daq/operations/screenshot/test_screenshot_service.py index 5680da8e..4d9f2018 100644 --- a/tests/unit/daq/operations/screenshot/test_screenshot_service.py +++ b/tests/unit/daq/operations/screenshot/test_screenshot_service.py @@ -2,8 +2,8 @@ import types import numpy as np import pytest +from aarecommon.models.models import DewarAddress, SampleShortInfo -from aare.common.models import DewarAddress, SampleShortInfo from aare.daq.operations.screenshot.service import ScreenshotService @@ -27,8 +27,8 @@ def test_save_to_db_uploads_image_via_shared_aare(mock_logger): mlbox = types.SimpleNamespace(get_latest_image=lambda: image) aare = types.SimpleNamespace( - upload_image=lambda sample_id, filename, bgr_image, **kwargs: upload_calls.append( - (sample_id, filename, bgr_image, kwargs) + upload_image=lambda sample_id, filename, bgr_image, **kwargs: ( + upload_calls.append((sample_id, filename, bgr_image, kwargs)) ) ) @@ -55,14 +55,16 @@ def test_save_to_db_uploads_image_via_shared_aare(mock_logger): assert noncritical_calls == [("screenshot upload 'mounted'", sample)] -def test_send_to_db_writes_photo_and_uploads_with_default_message(mock_logger, tmp_path): +def test_send_to_db_writes_photo_and_uploads_with_default_message( + mock_logger, tmp_path +): image = np.zeros((10, 10, 3), dtype=np.uint8) upload_calls = [] mlbox = types.SimpleNamespace(get_latest_image=lambda: image) aare = types.SimpleNamespace( - upload_image=lambda sample_id, filename, bgr_image, **kwargs: upload_calls.append( - (sample_id, filename, kwargs) + upload_image=lambda sample_id, filename, bgr_image, **kwargs: ( + upload_calls.append((sample_id, filename, kwargs)) ) ) sample = _make_sample(15) @@ -93,7 +95,9 @@ def test_send_to_db_writes_photo_and_uploads_with_default_message(mock_logger, t def test_send_to_db_requires_mounted_sample(mock_logger): service = ScreenshotService( - mlbox=types.SimpleNamespace(get_latest_image=lambda: np.zeros((4, 4, 3), dtype=np.uint8)), + mlbox=types.SimpleNamespace( + get_latest_image=lambda: np.zeros((4, 4, 3), dtype=np.uint8) + ), aare=types.SimpleNamespace(upload_image=lambda *args, **kwargs: None), logger=mock_logger, run_noncritical=lambda action, **kwargs: action(), @@ -102,4 +106,4 @@ def test_send_to_db_requires_mounted_sample(mock_logger): ) with pytest.raises(ValueError, match="valid sample_id"): - service.send_to_db(default_message="x") \ No newline at end of file + service.send_to_db(default_message="x") diff --git a/tests/unit/daq/operations/test_ml_raster_plan.py b/tests/unit/daq/operations/test_ml_raster_plan.py index 9a93c749..7e5737e8 100644 --- a/tests/unit/daq/operations/test_ml_raster_plan.py +++ b/tests/unit/daq/operations/test_ml_raster_plan.py @@ -3,10 +3,10 @@ import types import numpy as np import pytest +from aarecommon.math.coordinate import Coordinate, SmargonCoordinate +from aarecommon.math.sample_geometry import SampleGeometryModel +from aarecommon.models.models import MLBoxType, MLOutputModel -from aare.common.coordinate import Coordinate, SmargonCoordinate -from aare.common.models import MLBoxType, MLOutputModel -from aare.common.sample_geometry import SampleGeometryModel from aare.daq.mlbox import MLBoxPredictionsResult from aare.daq.operations.common import ml_bounding_box as mlb from aare.daq.operations.common.ml_bounding_box import ( @@ -25,7 +25,9 @@ def _geom() -> SampleGeometryModel: pixel_in_mm=0.001, aerotech=Coordinate(x=0, y=0), aerotech_meas=Coordinate(x=0, y=0), - smargon=SmargonCoordinate(sh_mm=Coordinate(x=0, y=0, z=0), phi_deg=0, chi_deg=0), + smargon=SmargonCoordinate( + sh_mm=Coordinate(x=0, y=0, z=0), phi_deg=0, chi_deg=0 + ), omega_deg=0, beam_size_mm=Coordinate(x=0.01, y=0.01), ) @@ -69,7 +71,9 @@ def test_box_helpers(): def test_plan_returns_loop_boxes_and_prefers_loop_face(): - plan = _plan(_fake_mlbox(loop_all=(100, 100, 400, 400), loop_face=(150, 150, 300, 300))) + plan = _plan( + _fake_mlbox(loop_all=(100, 100, 400, 400), loop_face=(150, 150, 300, 300)) + ) assert plan is not None assert plan.loop_all_box == (100, 100, 400, 400) assert plan.loop_face_box == (150, 150, 300, 300) @@ -85,16 +89,27 @@ def test_plan_none_when_no_loop(): def _grid(box, *, grid_padding): x1, y1, x2, y2 = box return _box_to_raster_request( - x1=x1, y1=y1, x2=x2, y2=y2, - sample=None, sample_geometry=_geom(), logger=logger, filename=None, - sample_id=None, max_images=100000, min_cell_size_mm=0.0001, - skip_if_exceed_max_image_threshold=False, grid_padding=grid_padding, + x1=x1, + y1=y1, + x2=x2, + y2=y2, + sample=None, + sample_geometry=_geom(), + logger=logger, + filename=None, + sample_id=None, + max_images=100000, + min_cell_size_mm=0.0001, + skip_if_exceed_max_image_threshold=False, + grid_padding=grid_padding, ) def test_grid_padding_grows_and_shifts_top_left(monkeypatch): # fraction 0 -> minimum one cell of padding per side - monkeypatch.setattr(mlb, "cfg_get", lambda k, d=None: 0.0 if "grid_padding_fraction" in k else d) + monkeypatch.setattr( + mlb, "cfg_get", lambda k, d=None: 0.0 if "grid_padding_fraction" in k else d + ) box = (150, 150, 400, 300) nopad = _grid(box, grid_padding=False) pad = _grid(box, grid_padding=True) @@ -110,37 +125,45 @@ def test_grid_padding_y_bottom_asymmetric(monkeypatch): def cfg(y_bottom): return lambda k, d=None: ( - y_bottom if "grid_padding_fraction_y_bottom" in k + y_bottom + if "grid_padding_fraction_y_bottom" in k else (0.0 if "grid_padding_fraction" in k else d) ) - monkeypatch.setattr(mlb, "cfg_get", cfg(0.0)) # bottom == top (min 1 cell each) + monkeypatch.setattr(mlb, "cfg_get", cfg(0.0)) # bottom == top (min 1 cell each) sym = _grid(box, grid_padding=True) - monkeypatch.setattr(mlb, "cfg_get", cfg(0.6)) # much more padding at the bottom + monkeypatch.setattr(mlb, "cfg_get", cfg(0.6)) # much more padding at the bottom bottom = _grid(box, grid_padding=True) - assert bottom.n_y > sym.n_y # extra cells added at the bottom - assert bottom.n_x == sym.n_x # x unaffected + assert bottom.n_y > sym.n_y # extra cells added at the bottom + assert bottom.n_x == sym.n_x # x unaffected # top padding identical -> smargon_top_left (cell 0) unchanged assert bottom.smargon_top_left == sym.smargon_top_left def test_grid_padding_fraction_scales(monkeypatch): box = (150, 150, 520, 420) - monkeypatch.setattr(mlb, "cfg_get", lambda k, d=None: 0.0 if "grid_padding_fraction" in k else d) + monkeypatch.setattr( + mlb, "cfg_get", lambda k, d=None: 0.0 if "grid_padding_fraction" in k else d + ) small = _grid(box, grid_padding=True) - monkeypatch.setattr(mlb, "cfg_get", lambda k, d=None: 0.5 if "grid_padding_fraction" in k else d) + monkeypatch.setattr( + mlb, "cfg_get", lambda k, d=None: 0.5 if "grid_padding_fraction" in k else d + ) big = _grid(box, grid_padding=True) assert big.n_x > small.n_x and big.n_y > small.n_y def test_crystal_union_extends_grid_only_when_enabled(monkeypatch): # crystal extends well beyond the loop_face box on +x - mlbox = lambda: _fake_mlbox(loop_face=(150, 150, 300, 300), crystals=[(350, 150, 520, 300)]) + mlbox = lambda: _fake_mlbox( + loop_face=(150, 150, 300, 300), crystals=[(350, 150, 520, 300)] + ) def cfg(enabled): return lambda k, d=None: ( - enabled if "include_crystal" in k + enabled + if "include_crystal" in k else (0.0 if "grid_padding_fraction" in k else d) ) @@ -160,21 +183,30 @@ def test_zoom_box_uses_loop_all_unless_clipped(): # loop_all fully inside the frame -> used for zoom ok = mlb.MLRasterPlan( - grid_request=None, loop_all_box=(100, 100, 400, 400), - loop_face_box=(150, 150, 300, 300), image_width=1000, image_height=1000, + grid_request=None, + loop_all_box=(100, 100, 400, 400), + loop_face_box=(150, 150, 300, 300), + image_width=1000, + image_height=1000, ) assert svc._zoom_to_fit_box(ok) == (100, 100, 400, 400) # loop_all touches the left edge (clipped) -> fall back to loop_face clipped = mlb.MLRasterPlan( - grid_request=None, loop_all_box=(0, 100, 400, 400), - loop_face_box=(150, 150, 300, 300), image_width=1000, image_height=1000, + grid_request=None, + loop_all_box=(0, 100, 400, 400), + loop_face_box=(150, 150, 300, 300), + image_width=1000, + image_height=1000, ) assert svc._zoom_to_fit_box(clipped) == (150, 150, 300, 300) # no loop_all -> loop_face only_face = mlb.MLRasterPlan( - grid_request=None, loop_all_box=None, - loop_face_box=(150, 150, 300, 300), image_width=1000, image_height=1000, + grid_request=None, + loop_all_box=None, + loop_face_box=(150, 150, 300, 300), + image_width=1000, + image_height=1000, ) assert svc._zoom_to_fit_box(only_face) == (150, 150, 300, 300) diff --git a/tests/unit/daq/test_aare_daq_loop_centering.py b/tests/unit/daq/test_aare_daq_loop_centering.py index 463bbf38..b47564f3 100644 --- a/tests/unit/daq/test_aare_daq_loop_centering.py +++ b/tests/unit/daq/test_aare_daq_loop_centering.py @@ -1,9 +1,15 @@ import types from unittest.mock import MagicMock, patch -from aare.common.automation_models import AutomationProgress, StepState, StepStatus, WorkflowStateKind -from aare.common.models import SampleShortInfo -from aare.common.exception_handler import LoopCenteringFailed +from aarecommon.errors.exception_handler import LoopCenteringFailed +from aarecommon.models.automation import ( + AutomationProgress, + StepState, + StepStatus, + WorkflowStateKind, +) +from aarecommon.models.models import SampleShortInfo + from aare.daq.daq import AareDAQ @@ -94,7 +100,13 @@ def test_record_best_effort_step_failure_marks_progress_and_logs_warning(mock_lo progress = AutomationProgress( current_step="Center", - steps=[StepState(step=WorkflowStateKind.LOOP_CENTRE, status=StepStatus.RUNNING, message="Centering")], + steps=[ + StepState( + step=WorkflowStateKind.LOOP_CENTRE, + status=StepStatus.RUNNING, + message="Centering", + ) + ], finished=False, success=None, ) @@ -127,4 +139,4 @@ def test_record_best_effort_step_failure_marks_progress_and_logs_warning(mock_lo assert progress.events[0].exception_class == "LoopCenteringFailed" assert progress.events[0].message == "Loop centering failed" assert progress.events[0].sample_id == 1 - mock_logger.warning.assert_called_once() \ No newline at end of file + mock_logger.warning.assert_called_once() diff --git a/tests/unit/daq/test_aaredb.py b/tests/unit/daq/test_aaredb.py index 39d69c4a..8c85782e 100644 --- a/tests/unit/daq/test_aaredb.py +++ b/tests/unit/daq/test_aaredb.py @@ -1,16 +1,17 @@ -import pytest -import numpy as np from unittest.mock import MagicMock, patch -from aare.daq.aaredb import AareWrapper -from aare.common.models import ( - SampleShortInfo, PuckLoadedInfo, DewarAddress -) -from aare.common.raster_grid import RasterGridRequest, CenterOfMassModel -from aare.common.rotation_scan import RotationScanRequest -from aare.common.sample_geometry import SampleGeometryModel -from aare.common.coordinate import Coordinate, SmargonCoordinate + +import numpy as np +import pytest +from aarecommon.math.coordinate import Coordinate, SmargonCoordinate +from aarecommon.math.sample_geometry import SampleGeometryModel +from aarecommon.models.models import DewarAddress, PuckLoadedInfo, SampleShortInfo +from aarecommon.models.raster_grid import CenterOfMassModel, RasterGridRequest +from aarecommon.models.rotation_scan import RotationScanRequest from jfjoch_client.models import ScanResult +from aare.daq.aaredb import AareWrapper + + @pytest.fixture def mock_bl(): bl = MagicMock() @@ -18,6 +19,7 @@ def mock_bl(): bl.name = "X10SA" return bl + @pytest.fixture def sample_info(): return SampleShortInfo( @@ -28,9 +30,10 @@ def sample_info(): run_number=1, user="testuser", pin=5, - location=DewarAddress(segment="A", pos=1) + location=DewarAddress(segment="A", pos=1), ) + @pytest.fixture def daq_status(mock_bl): status = MagicMock() @@ -48,9 +51,12 @@ def daq_status(mock_bl): status.bl.flux_ph_s = 1e12 status.bl.cryojet_K = 100.0 status.geom.beam_size_mm = Coordinate(x=0.01, y=0.01) - status.geom.smargon = SmargonCoordinate(sh_mm=Coordinate(x=0,y=0,z=0), phi_deg=0, chi_deg=0) + status.geom.smargon = SmargonCoordinate( + sh_mm=Coordinate(x=0, y=0, z=0), phi_deg=0, chi_deg=0 + ) return status + @pytest.fixture def geom_model(): return SampleGeometryModel( @@ -58,32 +64,41 @@ def geom_model(): pixel_in_mm=0.001, aerotech=Coordinate(x=0, y=0), aerotech_meas=Coordinate(x=0, y=0), - smargon=SmargonCoordinate(sh_mm=Coordinate(x=0,y=0,z=0), phi_deg=0, chi_deg=0), + smargon=SmargonCoordinate( + sh_mm=Coordinate(x=0, y=0, z=0), phi_deg=0, chi_deg=0 + ), omega_deg=0, - beam_size_mm=Coordinate(x=0.01, y=0.01) + beam_size_mm=Coordinate(x=0.01, y=0.01), ) + @patch("aareDB.ApiClient") @patch("aareDB.TellsRunnerApi") @patch("aareDB.SamplesRunnerApi") @patch("aareDB.ProcessingsRunnerApi") @patch("aareDB.GridscanRunnerApi") -def test_aare_wrapper_init(mock_grid, mock_proc, mock_sample, mock_tell, mock_api, mock_bl): +def test_aare_wrapper_init( + mock_grid, mock_proc, mock_sample, mock_tell, mock_api, mock_bl +): wrapper = AareWrapper(bl=mock_bl) mock_api.assert_called_once() assert wrapper._AareWrapper__bl == mock_bl + @patch("aareDB.ApiClient") @patch("aareDB.TellsRunnerApi") @patch("aareDB.SamplesRunnerApi") @patch("aareDB.ProcessingsRunnerApi") @patch("aareDB.GridscanRunnerApi") -def test_set_pucks_beamline(mock_grid, mock_proc, mock_sample, mock_tell, mock_api, mock_bl): +def test_set_pucks_beamline( + mock_grid, mock_proc, mock_sample, mock_tell, mock_api, mock_bl +): wrapper = AareWrapper(bl=mock_bl) pucks = [PuckLoadedInfo(puck_name="P1", location=DewarAddress(segment="A", pos=1))] wrapper.set_pucks_beamline(pucks) mock_tell.return_value.set_tell_positions.assert_called_once() + @patch("aareDB.ApiClient") @patch("aareDB.SamplesRunnerApi") def test_create_manual_sample(mock_sample, mock_api, mock_bl, sample_info): @@ -93,12 +108,16 @@ def test_create_manual_sample(mock_sample, mock_api, mock_bl, sample_info): assert sample_info.db_id == 456 mock_sample.return_value.insert_sample.assert_called_once() + @patch("aareDB.ApiClient") @patch("aareDB.SamplesRunnerApi") def test_send_sample_event(mock_sample, mock_api, mock_bl, sample_info): from aareDB import SampleEventType + wrapper = AareWrapper(bl=mock_bl) - wrapper.send_sample_event(sample_info.db_id, SampleEventType.MOUNTED, "Test comment") + wrapper.send_sample_event( + sample_info.db_id, SampleEventType.MOUNTED, "Test comment" + ) mock_sample.return_value.create_sample_event.assert_called_once() # Test None sample @@ -111,21 +130,28 @@ def test_send_sample_event(mock_sample, mock_api, mock_bl, sample_info): wrapper.send_sample_event(sample_info.db_id, SampleEventType.MOUNTED) mock_sample.return_value.create_sample_event.assert_not_called() + @patch("aareDB.ApiClient") @patch("aareDB.SamplesRunnerApi") -def test_create_manual_sample_error(mock_sample, mock_api, mock_bl, sample_info, caplog): +def test_create_manual_sample_error( + mock_sample, mock_api, mock_bl, sample_info, caplog +): import logging + caplog.set_level(logging.ERROR) wrapper = AareWrapper(bl=mock_bl) mock_sample.return_value.insert_sample.side_effect = Exception("DB Error") wrapper.create_manual_sample(sample_info) assert "Error inserting sample: DB Error" in caplog.text + @patch("aareDB.ApiClient") @patch("aareDB.SamplesRunnerApi") def test_send_sample_event_error(mock_sample, mock_api, mock_bl, sample_info, caplog): - from aareDB import SampleEventType import logging + + from aareDB import SampleEventType + caplog.set_level(logging.ERROR) wrapper = AareWrapper(bl=mock_bl) mock_sample.return_value.create_sample_event.side_effect = Exception("Event Error") @@ -134,6 +160,7 @@ def test_send_sample_event_error(mock_sample, mock_api, mock_bl, sample_info, ca assert "MOUNTED" in caplog.text assert "Event Error" in caplog.text + @patch("aareDB.ApiClient") @patch("requests.post") def test_upload_image(mock_post, mock_api, mock_bl): @@ -144,6 +171,7 @@ def test_upload_image(mock_post, mock_api, mock_bl): mock_post.assert_called_once() assert "test_img.jpg" in str(mock_post.call_args) + @patch("aareDB.ApiClient") @patch("requests.post") def test_upload_jpg(mock_post, mock_api, mock_bl): @@ -152,111 +180,197 @@ def test_upload_jpg(mock_post, mock_api, mock_bl): wrapper.upload_jpg(123, "test_img", b"fake jpeg bytes") mock_post.assert_called_once() + @patch("aareDB.ApiClient") @patch("aareDB.SamplesRunnerApi") def test_create_rotation_run(mock_sample, mock_api, mock_bl, sample_info, daq_status): wrapper = AareWrapper(bl=mock_bl) # Standard rotation - req = RotationScanRequest(exp_time_s=0.1, incr_omega_deg=0.1, steps=100, file_prefix="test_prefix", dtz=200.0) + req = RotationScanRequest( + exp_time_s=0.1, + incr_omega_deg=0.1, + steps=100, + file_prefix="test_prefix", + dtz=200.0, + ) wrapper.create_rotation_run(sample_info, req, daq_status) mock_sample.return_value.create_experiment_parameters_for_sample.assert_called_once() - + # None sample - should return early mock_sample.return_value.create_experiment_parameters_for_sample.reset_mock() wrapper.create_rotation_run(None, req, daq_status) mock_sample.return_value.create_experiment_parameters_for_sample.assert_not_called() -@patch("aareDB.ApiClient") -@patch("aareDB.SamplesRunnerApi") -def test_create_rotation_run_screening(mock_sample, mock_api, mock_bl, sample_info, daq_status): - wrapper = AareWrapper(bl=mock_bl) - req = RotationScanRequest(exp_time_s=0.1, incr_omega_deg=0.1, steps=100, file_prefix="test_prefix", screening=True, wedge_omega_deg=5.0, dtz=200.0) - wrapper.create_rotation_run(sample_info, req, daq_status) - mock_sample.return_value.create_experiment_parameters_for_sample.assert_called_once() @patch("aareDB.ApiClient") @patch("aareDB.SamplesRunnerApi") -def test_create_rotation_run_error(mock_sample, mock_api, mock_bl, sample_info, daq_status, caplog): +def test_create_rotation_run_screening( + mock_sample, mock_api, mock_bl, sample_info, daq_status +): wrapper = AareWrapper(bl=mock_bl) - req = RotationScanRequest(exp_time_s=0.1, incr_omega_deg=0.1, steps=100, file_prefix="test_prefix", dtz=200.0) - mock_sample.return_value.create_experiment_parameters_for_sample.side_effect = Exception("Rotation error") + req = RotationScanRequest( + exp_time_s=0.1, + incr_omega_deg=0.1, + steps=100, + file_prefix="test_prefix", + screening=True, + wedge_omega_deg=5.0, + dtz=200.0, + ) + wrapper.create_rotation_run(sample_info, req, daq_status) + mock_sample.return_value.create_experiment_parameters_for_sample.assert_called_once() + + +@patch("aareDB.ApiClient") +@patch("aareDB.SamplesRunnerApi") +def test_create_rotation_run_error( + mock_sample, mock_api, mock_bl, sample_info, daq_status, caplog +): + wrapper = AareWrapper(bl=mock_bl) + req = RotationScanRequest( + exp_time_s=0.1, + incr_omega_deg=0.1, + steps=100, + file_prefix="test_prefix", + dtz=200.0, + ) + mock_sample.return_value.create_experiment_parameters_for_sample.side_effect = ( + Exception("Rotation error") + ) wrapper.create_rotation_run(sample_info, req, daq_status) assert "Rotation error" in caplog.text + @patch("aareDB.ApiClient") @patch("aareDB.SamplesRunnerApi") def test_create_gridscan_run(mock_sample, mock_api, mock_bl, sample_info, daq_status): wrapper = AareWrapper(bl=mock_bl) - req = RasterGridRequest(exp_time_s=0.1, n_x=10, n_y=10, grid_size_mm=Coordinate(x=0.01, y=0.01), smargon_top_left=None, file_prefix="test_grid_prefix", dtz=200.0) + req = RasterGridRequest( + exp_time_s=0.1, + n_x=10, + n_y=10, + grid_size_mm=Coordinate(x=0.01, y=0.01), + smargon_top_left=None, + file_prefix="test_grid_prefix", + dtz=200.0, + ) wrapper.create_gridscan_run(sample_info, req, daq_status) mock_sample.return_value.create_experiment_parameters_for_sample.assert_called_once() - + # None sample mock_sample.return_value.create_experiment_parameters_for_sample.reset_mock() wrapper.create_gridscan_run(None, req, daq_status) mock_sample.return_value.create_experiment_parameters_for_sample.assert_not_called() + @patch("aareDB.ApiClient") @patch("aareDB.SamplesRunnerApi") -def test_create_gridscan_run_error(mock_sample, mock_api, mock_bl, sample_info, daq_status, caplog): +def test_create_gridscan_run_error( + mock_sample, mock_api, mock_bl, sample_info, daq_status, caplog +): wrapper = AareWrapper(bl=mock_bl) - req = RasterGridRequest(exp_time_s=0.1, n_x=10, n_y=10, grid_size_mm=Coordinate(x=0.01, y=0.01), smargon_top_left=None, file_prefix="test_grid_prefix", dtz=200.0) - mock_sample.return_value.create_experiment_parameters_for_sample.side_effect = Exception("Grid error") + req = RasterGridRequest( + exp_time_s=0.1, + n_x=10, + n_y=10, + grid_size_mm=Coordinate(x=0.01, y=0.01), + smargon_top_left=None, + file_prefix="test_grid_prefix", + dtz=200.0, + ) + mock_sample.return_value.create_experiment_parameters_for_sample.side_effect = ( + Exception("Grid error") + ) wrapper.create_gridscan_run(sample_info, req, daq_status) assert "Grid error" in caplog.text + @patch("aareDB.ApiClient") @patch("requests.post") def test_ingest_gridscan(mock_post, mock_api, mock_bl, sample_info, geom_model): wrapper = AareWrapper(bl=mock_bl) mock_post.return_value.status_code = 200 raster_result = MagicMock(spec=ScanResult) - raster_request = RasterGridRequest(exp_time_s=0.1, n_x=10, n_y=10, grid_size_mm=Coordinate(x=0.01, y=0.01), smargon_top_left=None) + raster_request = RasterGridRequest( + exp_time_s=0.1, + n_x=10, + n_y=10, + grid_size_mm=Coordinate(x=0.01, y=0.01), + smargon_top_left=None, + ) com = CenterOfMassModel(n_x=5.0, n_y=5.0) - - wrapper.ingest_gridscan(sample_info, raster_result, raster_request, geom_model, com, (500.0, 500.0)) + + wrapper.ingest_gridscan( + sample_info, raster_result, raster_request, geom_model, com, (500.0, 500.0) + ) mock_post.assert_called_once() - + # None sample mock_post.reset_mock() - wrapper.ingest_gridscan(None, raster_result, raster_request, geom_model, com, (500.0, 500.0)) + wrapper.ingest_gridscan( + None, raster_result, raster_request, geom_model, com, (500.0, 500.0) + ) mock_post.assert_not_called() + @patch("aareDB.ApiClient") @patch("requests.post") def test_ingest_scan(mock_post, mock_api, mock_bl, sample_info, geom_model): wrapper = AareWrapper(bl=mock_bl) mock_post.return_value.status_code = 200 result = MagicMock(spec=ScanResult) - + wrapper.ingest_scan(sample_info, result, geom_model, (500.0, 500.0)) mock_post.assert_called_once() - + # None sample mock_post.reset_mock() wrapper.ingest_scan(None, result, geom_model, (500.0, 500.0)) mock_post.assert_not_called() + @patch("aareDB.ApiClient") def test_format_gridscan_payload_no_com(mock_api, mock_bl, sample_info, geom_model): wrapper = AareWrapper(bl=mock_bl) raster_result = MagicMock(spec=ScanResult) - raster_request = RasterGridRequest(exp_time_s=0.1, n_x=10, n_y=10, grid_size_mm=Coordinate(x=0.01, y=0.01), smargon_top_left=None) - - payload = wrapper.format_gridscan_payload(sample_info, raster_result, raster_request, geom_model, None, (500.0, 500.0)) + raster_request = RasterGridRequest( + exp_time_s=0.1, + n_x=10, + n_y=10, + grid_size_mm=Coordinate(x=0.01, y=0.01), + smargon_top_left=None, + ) + + payload = wrapper.format_gridscan_payload( + sample_info, raster_result, raster_request, geom_model, None, (500.0, 500.0) + ) assert payload.center_pxl is None assert payload.sample_id == sample_info.db_id + @patch("aareDB.ApiClient") -def test_format_gridscan_payload_with_top_left(mock_api, mock_bl, sample_info, geom_model): +def test_format_gridscan_payload_with_top_left( + mock_api, mock_bl, sample_info, geom_model +): wrapper = AareWrapper(bl=mock_bl) raster_result = MagicMock(spec=ScanResult) - top_left = SmargonCoordinate(sh_mm=Coordinate(x=0.1, y=0.1, z=0.1), phi_deg=0, chi_deg=0) - raster_request = RasterGridRequest(exp_time_s=0.1, n_x=10, n_y=10, grid_size_mm=Coordinate(x=0.01, y=0.01), smargon_top_left=top_left) - - payload = wrapper.format_gridscan_payload(sample_info, raster_result, raster_request, geom_model, None, (500.0, 500.0)) + top_left = SmargonCoordinate( + sh_mm=Coordinate(x=0.1, y=0.1, z=0.1), phi_deg=0, chi_deg=0 + ) + raster_request = RasterGridRequest( + exp_time_s=0.1, + n_x=10, + n_y=10, + grid_size_mm=Coordinate(x=0.01, y=0.01), + smargon_top_left=top_left, + ) + + payload = wrapper.format_gridscan_payload( + sample_info, raster_result, raster_request, geom_model, None, (500.0, 500.0) + ) assert payload.sample_id == sample_info.db_id + @patch("aareDB.ApiClient") def test_ingest_gridscan_payload_none(mock_api, mock_bl, sample_info, geom_model): wrapper = AareWrapper(bl=mock_bl) @@ -265,24 +379,26 @@ def test_ingest_gridscan_payload_none(mock_api, mock_bl, sample_info, geom_model # But wait, it raises 'e' after logging. Oh, I see. # Lines 351-352 in aaredb.py: if payload is None: return # This happens if format_gridscan_payload returns None. - with patch.object(AareWrapper, 'format_gridscan_payload', return_value=None): - wrapper.ingest_gridscan(sample_info, None, None, geom_model, None, (0,0)) - - with patch.object(AareWrapper, 'format_scan_payload', return_value=None): - wrapper.ingest_scan(sample_info, None, geom_model, (0,0)) + with patch.object(AareWrapper, "format_gridscan_payload", return_value=None): + wrapper.ingest_gridscan(sample_info, None, None, geom_model, None, (0, 0)) + + with patch.object(AareWrapper, "format_scan_payload", return_value=None): + wrapper.ingest_scan(sample_info, None, geom_model, (0, 0)) + @patch("aareDB.ApiClient") def test_format_gridscan_payload_error(mock_api, mock_bl, sample_info, caplog): wrapper = AareWrapper(bl=mock_bl) # Passing None for geom_model should trigger an error in smargon_to_picture with pytest.raises(Exception): - wrapper.format_gridscan_payload(sample_info, None, None, None, None, (0,0)) + wrapper.format_gridscan_payload(sample_info, None, None, None, None, (0, 0)) assert "NoneType" in caplog.text + @patch("aareDB.ApiClient") def test_format_scan_payload_error(mock_api, mock_bl, sample_info, caplog): wrapper = AareWrapper(bl=mock_bl) with pytest.raises(Exception): # Passing None for geom should trigger error when accessing geom.beam_size_mm - wrapper.format_scan_payload(sample_info, None, None, (0,0)) + wrapper.format_scan_payload(sample_info, None, None, (0, 0)) assert "NoneType" in caplog.text diff --git a/tests/unit/daq/test_auth.py b/tests/unit/daq/test_auth.py index bbfc2f2b..ef28cfc4 100644 --- a/tests/unit/daq/test_auth.py +++ b/tests/unit/daq/test_auth.py @@ -1,23 +1,44 @@ -import pytest -import jwt import time import uuid +from datetime import UTC, datetime from unittest.mock import MagicMock, patch -from datetime import datetime, UTC + +import jwt +import pytest from fastapi import HTTPException # Mock environment variable before importing auth -with patch.dict('os.environ', {'JWT_AAREDAQ_KEY': 'test_secret'}): +with patch.dict("os.environ", {"JWT_AAREDAQ_KEY": "test_secret"}): from aare.daq.auth import ( - TokenData, create_access_token, authenticate_user, parse_token, - check_jwt_ro, check_jwt_rw, check_jwt_staff_only, check_jwt_staff, - force_current_sesion, get_baton_status, request_baton, - respond_to_baton_request, release_baton, cancel_baton_request, - resolve_baton_timeout_if_needed + TokenData, + authenticate_user, + cancel_baton_request, + check_jwt_ro, + check_jwt_rw, + check_jwt_staff, + check_jwt_staff_only, + create_access_token, + force_current_sesion, + get_baton_status, + parse_token, + release_baton, + request_baton, + resolve_baton_timeout_if_needed, + respond_to_baton_request, ) -from aare.common.exception_handler import AuthenticationException, UserRightsException -from aare.common.auth_models import BatonStatus, BatonRequest, BatonRequestStatus, BatonHolderInfo, BatonTransferQueue -from aare.common.models import SessionsStateEnum +from aarecommon.errors.exception_handler import ( + AuthenticationException, + UserRightsException, +) +from aarecommon.models.auth import ( + BatonHolderInfo, + BatonRequest, + BatonRequestStatus, + BatonStatus, + BatonTransferQueue, +) +from aarecommon.models.models import SessionsStateEnum + @pytest.fixture def mock_cfg(): @@ -29,28 +50,35 @@ def mock_cfg(): cfg.pending_baton_request = None return cfg + @pytest.fixture def token_data(): - return TokenData(sub="testuser", pgroups=["p12345", "p67890"], session=100, staff=False) + return TokenData( + sub="testuser", pgroups=["p12345", "p67890"], session=100, staff=False + ) + @pytest.fixture def staff_token_data(): return TokenData(sub="staffuser", pgroups=["p12345"], session=101, staff=True) + def test_create_access_token(token_data): - with patch('aare.daq.auth.SECRET_KEY', 'test_secret'): + with patch("aare.daq.auth.SECRET_KEY", "test_secret"): token = create_access_token(token_data) assert isinstance(token, str) - payload = jwt.decode(token, 'test_secret', algorithms=["HS256"]) + payload = jwt.decode(token, "test_secret", algorithms=["HS256"]) assert payload["sub"] == "testuser" assert payload["session"] == 100 -def test_authenticate_user(mock_cfg): - with patch('pwd.getpwnam') as mock_pwd, \ - patch('os.getgrouplist') as mock_groups, \ - patch('grp.getgrgid') as mock_grp, \ - patch('aare.daq.auth.SECRET_KEY', 'test_secret'): +def test_authenticate_user(mock_cfg): + with ( + patch("pwd.getpwnam") as mock_pwd, + patch("os.getgrouplist") as mock_groups, + patch("grp.getgrgid") as mock_grp, + patch("aare.daq.auth.SECRET_KEY", "test_secret"), + ): mock_pwd.return_value.pw_name = "testuser" mock_pwd.return_value.pw_gid = 1000 mock_groups.return_value = [1000, 1001] @@ -68,180 +96,257 @@ def test_authenticate_user(mock_cfg): token = authenticate_user(mock_cfg, "testuser") assert isinstance(token, str) - payload = jwt.decode(token, 'test_secret', algorithms=["HS256"]) + payload = jwt.decode(token, "test_secret", algorithms=["HS256"]) assert payload["sub"] == "testuser" assert "p12345" in payload["pgroups"] assert payload["staff"] is True + def test_parse_token(): - with patch('aare.daq.auth.SECRET_KEY', 'test_secret'): - token = jwt.encode({"sub": "user", "pgroups": [], "session": 1, "staff": False}, "test_secret") + with patch("aare.daq.auth.SECRET_KEY", "test_secret"): + token = jwt.encode( + {"sub": "user", "pgroups": [], "session": 1, "staff": False}, "test_secret" + ) data = parse_token(token) assert data.sub == "user" + def test_parse_token_invalid(): with pytest.raises(AuthenticationException): parse_token("invalid.token.here") + def test_check_jwt_ro(mock_cfg, token_data): # Success check_jwt_ro(mock_cfg, token_data) - + # Fail mock_cfg.pgroup = "p99999" with pytest.raises(UserRightsException): check_jwt_ro(mock_cfg, token_data) - + # Staff success even if not in pgroup staff_data = TokenData(sub="staff", pgroups=[], session=2, staff=True) check_jwt_ro(mock_cfg, staff_data) + def test_check_jwt_rw(mock_cfg, token_data): - mock_cfg.baton_holder = BatonHolderInfo(session=100, username="testuser", is_staff=False) + mock_cfg.baton_holder = BatonHolderInfo( + session=100, username="testuser", is_staff=False + ) check_jwt_rw(mock_cfg, token_data) - + # Not holder - mock_cfg.baton_holder = BatonHolderInfo(session=200, username="other", is_staff=False) + mock_cfg.baton_holder = BatonHolderInfo( + session=200, username="other", is_staff=False + ) with pytest.raises(UserRightsException): check_jwt_rw(mock_cfg, token_data) + def test_check_jwt_staff_only(token_data, staff_token_data): check_jwt_staff_only(staff_token_data) with pytest.raises(UserRightsException): check_jwt_staff_only(token_data) + def test_force_current_sesion(mock_cfg, token_data): force_current_sesion(mock_cfg, token_data) mock_cfg.execute_baton_transfer.assert_called_once() + def test_get_baton_status(mock_cfg, token_data): mock_cfg.baton_holder = None mock_cfg.pending_baton_request = None mock_cfg.queued_baton_transfer = None - + status = get_baton_status(mock_cfg, token_data) assert isinstance(status, BatonStatus) assert status.you_are_holder is False + def test_request_baton_vacant(mock_cfg, token_data): mock_cfg.session_state.return_value = SessionsStateEnum.Vacant res = request_baton(mock_cfg, token_data) assert res["granted"] is True mock_cfg.execute_baton_transfer.assert_called_once() + def test_request_baton_owned_by_you(mock_cfg, token_data): mock_cfg.session_state.return_value = SessionsStateEnum.OwnedByYou res = request_baton(mock_cfg, token_data) assert res["already_holder"] is True + def test_request_baton_pending(mock_cfg, token_data): mock_cfg.session_state.return_value = SessionsStateEnum.OwnedByElse - mock_cfg.baton_holder = BatonHolderInfo(session=200, username="other", is_staff=False) + mock_cfg.baton_holder = BatonHolderInfo( + session=200, username="other", is_staff=False + ) mock_cfg.pending_baton_request = None mock_cfg.can_transfer_baton_now.return_value = True mock_cfg.allow_non_staff_request_from_staff = True - + res = request_baton(mock_cfg, token_data) assert res["pending"] is True mock_cfg.set_pending_baton_request.assert_called_once() + def test_respond_to_baton_request_accept(mock_cfg, token_data): - mock_cfg.baton_holder = BatonHolderInfo(session=100, username="testuser", is_staff=False) + mock_cfg.baton_holder = BatonHolderInfo( + session=100, username="testuser", is_staff=False + ) mock_cfg.pending_baton_request = BatonRequest( - request_id="1", requester_username="other", requester_session=200, - requester_is_staff=False, holder_username="testuser", holder_session=100, - created_at=time.time(), timeout_seconds=30, status=BatonRequestStatus.PENDING + request_id="1", + requester_username="other", + requester_session=200, + requester_is_staff=False, + holder_username="testuser", + holder_session=100, + created_at=time.time(), + timeout_seconds=30, + status=BatonRequestStatus.PENDING, ) mock_cfg.can_transfer_baton_now.return_value = True - + res = respond_to_baton_request(mock_cfg, token_data, accept=True) assert res["transferred"] is True mock_cfg.execute_baton_transfer.assert_called_once() + def test_release_baton(mock_cfg, token_data): - mock_cfg.baton_holder = BatonHolderInfo(session=100, username="testuser", is_staff=False) + mock_cfg.baton_holder = BatonHolderInfo( + session=100, username="testuser", is_staff=False + ) res = release_baton(mock_cfg, token_data) assert res["released"] is True mock_cfg.end_active_session.assert_called_with(100) + def test_request_baton_staff_override(mock_cfg, staff_token_data): mock_cfg.session_state.return_value = SessionsStateEnum.OwnedByElse - mock_cfg.baton_holder = BatonHolderInfo(session=200, username="other", is_staff=False) + mock_cfg.baton_holder = BatonHolderInfo( + session=200, username="other", is_staff=False + ) mock_cfg.can_transfer_baton_now.return_value = True - + res = request_baton(mock_cfg, staff_token_data) assert res["granted"] is True assert res["override"] is True + def test_request_baton_staff_override_busy(mock_cfg, staff_token_data): mock_cfg.session_state.return_value = SessionsStateEnum.OwnedByElse - mock_cfg.baton_holder = BatonHolderInfo(session=200, username="other", is_staff=False) + mock_cfg.baton_holder = BatonHolderInfo( + session=200, username="other", is_staff=False + ) mock_cfg.can_transfer_baton_now.return_value = False - + res = request_baton(mock_cfg, staff_token_data) assert res["queued"] is True + def test_resolve_baton_timeout_if_needed(mock_cfg): # Case: No pending request mock_cfg.pending_baton_request = None assert resolve_baton_timeout_if_needed(mock_cfg) is None - + # Case: Request not expired mock_cfg.pending_baton_request = BatonRequest( - request_id="1", requester_username="other", requester_session=200, - requester_is_staff=False, holder_username="testuser", holder_session=100, - created_at=time.time(), timeout_seconds=30, status=BatonRequestStatus.PENDING + request_id="1", + requester_username="other", + requester_session=200, + requester_is_staff=False, + holder_username="testuser", + holder_session=100, + created_at=time.time(), + timeout_seconds=30, + status=BatonRequestStatus.PENDING, ) assert resolve_baton_timeout_if_needed(mock_cfg) is None - + # Case: Request expired, beamline not busy mock_cfg.pending_baton_request.created_at = time.time() - 40 mock_cfg.can_transfer_baton_now.return_value = True - mock_cfg.baton_holder = BatonHolderInfo(session=100, username="testuser", is_staff=False) - + mock_cfg.baton_holder = BatonHolderInfo( + session=100, username="testuser", is_staff=False + ) + resolve_baton_timeout_if_needed(mock_cfg) mock_cfg.execute_baton_transfer.assert_called_once() mock_cfg.clear_pending_baton_request.assert_called_once() + def test_respond_to_baton_request_refuse(mock_cfg, token_data): - mock_cfg.baton_holder = BatonHolderInfo(session=100, username="testuser", is_staff=False) + mock_cfg.baton_holder = BatonHolderInfo( + session=100, username="testuser", is_staff=False + ) mock_cfg.pending_baton_request = BatonRequest( - request_id="1", requester_username="other", requester_session=200, - requester_is_staff=False, holder_username="testuser", holder_session=100, - created_at=time.time(), timeout_seconds=30, status=BatonRequestStatus.PENDING + request_id="1", + requester_username="other", + requester_session=200, + requester_is_staff=False, + holder_username="testuser", + holder_session=100, + created_at=time.time(), + timeout_seconds=30, + status=BatonRequestStatus.PENDING, ) res = respond_to_baton_request(mock_cfg, token_data, accept=False) assert res["refused"] is True assert mock_cfg.pending_baton_request.status == BatonRequestStatus.REFUSED + def test_respond_to_baton_request_accept_busy(mock_cfg, token_data): - mock_cfg.baton_holder = BatonHolderInfo(session=100, username="testuser", is_staff=False) + mock_cfg.baton_holder = BatonHolderInfo( + session=100, username="testuser", is_staff=False + ) mock_cfg.pending_baton_request = BatonRequest( - request_id="1", requester_username="other", requester_session=200, - requester_is_staff=False, holder_username="testuser", holder_session=100, - created_at=time.time(), timeout_seconds=30, status=BatonRequestStatus.PENDING + request_id="1", + requester_username="other", + requester_session=200, + requester_is_staff=False, + holder_username="testuser", + holder_session=100, + created_at=time.time(), + timeout_seconds=30, + status=BatonRequestStatus.PENDING, ) mock_cfg.can_transfer_baton_now.return_value = False - + res = respond_to_baton_request(mock_cfg, token_data, accept=True) assert res["queued"] is True assert mock_cfg.queued_baton_transfer is not None + def test_cancel_baton_request(mock_cfg, token_data): mock_cfg.pending_baton_request = BatonRequest( - request_id="1", requester_username="testuser", requester_session=100, - requester_is_staff=False, holder_username="other", holder_session=200, - created_at=time.time(), timeout_seconds=30, status=BatonRequestStatus.PENDING + request_id="1", + requester_username="testuser", + requester_session=100, + requester_is_staff=False, + holder_username="other", + holder_session=200, + created_at=time.time(), + timeout_seconds=30, + status=BatonRequestStatus.PENDING, ) res = cancel_baton_request(mock_cfg, token_data) assert res["cancelled"] is True mock_cfg.clear_pending_baton_request.assert_called_once() + def test_resolve_baton_timeout_busy(mock_cfg): mock_cfg.pending_baton_request = BatonRequest( - request_id="1", requester_username="other", requester_session=200, - requester_is_staff=False, holder_username="testuser", holder_session=100, - created_at=time.time() - 40, timeout_seconds=30, status=BatonRequestStatus.PENDING + request_id="1", + requester_username="other", + requester_session=200, + requester_is_staff=False, + holder_username="testuser", + holder_session=100, + created_at=time.time() - 40, + timeout_seconds=30, + status=BatonRequestStatus.PENDING, ) mock_cfg.can_transfer_baton_now.return_value = False resolve_baton_timeout_if_needed(mock_cfg) diff --git a/tests/unit/daq/test_automation_progress_state_manager.py b/tests/unit/daq/test_automation_progress_state_manager.py index 46796e52..31403e52 100644 --- a/tests/unit/daq/test_automation_progress_state_manager.py +++ b/tests/unit/daq/test_automation_progress_state_manager.py @@ -1,11 +1,12 @@ from dataclasses import asdict -from aare.common.automation_models import ( +from aarecommon.models.automation import ( AutomationProgress, StepState, StepStatus, WorkflowStateKind, ) + from aare.daq.config import BeamlineConfig @@ -118,4 +119,4 @@ def test_automation_progress_state_seq_increments(): final_state = cfg.get_automation_progress_state() assert final_state["seq"] == 2 - assert final_state["progress"] == asdict(second) \ No newline at end of file + assert final_state["progress"] == asdict(second) diff --git a/tests/unit/daq/test_face_detection.py b/tests/unit/daq/test_face_detection.py index 2ec9c9c7..529a2cea 100644 --- a/tests/unit/daq/test_face_detection.py +++ b/tests/unit/daq/test_face_detection.py @@ -1,6 +1,7 @@ import types -from aare.common.models import DAQOperation +from aarecommon.models.models import DAQOperation + from aare.daq.operations.face_detection.models import FaceDetectionResult @@ -22,7 +23,12 @@ def test_execute_face_detection_reports_failure(monkeypatch): lambda: types.SimpleNamespace( run=lambda **kwargs: FaceDetectionResult( success=False, - payload={"running": False, "samples": [], "height_fit": {}, "area_fit": {}}, + payload={ + "running": False, + "samples": [], + "height_fit": {}, + "area_fit": {}, + }, error=RuntimeError("fd failed"), comment="face detection sequence failed", ) @@ -61,7 +67,12 @@ def test_execute_face_detection_can_skip_error_reporting(monkeypatch): lambda: types.SimpleNamespace( run=lambda **kwargs: FaceDetectionResult( success=False, - payload={"running": False, "samples": [], "height_fit": {}, "area_fit": {}}, + payload={ + "running": False, + "samples": [], + "height_fit": {}, + "area_fit": {}, + }, error=RuntimeError("fd failed"), comment="face detection sequence failed", ) @@ -94,7 +105,12 @@ def test_public_face_detection_uses_execute_face_detection(monkeypatch): "_execute_face_detection", lambda **kwargs: FaceDetectionResult( success=True, - payload={"running": False, "samples": [{"angle": 45}], "height_fit": {}, "area_fit": {}}, + payload={ + "running": False, + "samples": [{"angle": 45}], + "height_fit": {}, + "area_fit": {}, + }, ), ) @@ -102,4 +118,4 @@ def test_public_face_detection_uses_execute_face_detection(monkeypatch): assert result["running"] is False assert result["samples"] == [{"angle": 45}] - assert cfg.state_busy is False \ No newline at end of file + assert cfg.state_busy is False diff --git a/tests/unit/daq/test_gui_timeout.py b/tests/unit/daq/test_gui_timeout.py index d16405a6..5d89f73b 100644 --- a/tests/unit/daq/test_gui_timeout.py +++ b/tests/unit/daq/test_gui_timeout.py @@ -2,11 +2,13 @@ from types import SimpleNamespace from unittest.mock import patch -def test_status_renews_gui_session_with_gui_timeout(client, mock_backend, daq_status_factory): +def test_status_renews_gui_session_with_gui_timeout( + client, mock_backend, daq_status_factory +): mock_cfg = mock_backend["cfg"] mock_daq = mock_backend["daq"] - from aare.common.models import SessionsStateEnum + from aarecommon.models.models import SessionsStateEnum mock_daq.status = daq_status_factory( current_pgroup="p12345", @@ -37,8 +39,9 @@ def test_status_renews_gui_session_with_gui_timeout(client, mock_backend, daq_st def test_status_hides_open_guis_for_non_staff(client, mock_backend, daq_status_factory): with patch("aare.daq.auth.parse_token") as mock_parse: + from aarecommon.models.models import OpenGuiSessionInfo, SessionsStateEnum + from aare.daq.auth import TokenData - from aare.common.models import OpenGuiSessionInfo, SessionsStateEnum mock_parse.return_value = TokenData( sub="user1", @@ -59,7 +62,9 @@ def test_status_hides_open_guis_for_non_staff(client, mock_backend, daq_status_f mock_backend["cfg"].pgroup = "p12345" mock_backend["cfg"].session_state.return_value = SessionsStateEnum.OwnedByElse mock_backend["cfg"].get_open_gui_sessions.return_value = [ - OpenGuiSessionInfo(session=1, username="staff1", last_seen_ts=1.0, staff=True) + OpenGuiSessionInfo( + session=1, username="staff1", last_seen_ts=1.0, staff=True + ) ] mock_backend["cfg"].get_gui_session.return_value = OpenGuiSessionInfo( session=123, @@ -79,7 +84,8 @@ def test_status_hides_open_guis_for_non_staff(client, mock_backend, daq_status_f def test_admin_gui_sessions_includes_baton_holder_flag(client, mock_backend): from types import SimpleNamespace - from aare.common.models import OpenGuiSessionInfo + + from aarecommon.models.models import OpenGuiSessionInfo mock_cfg = mock_backend["cfg"] mock_cfg.get_open_gui_sessions.return_value = [ @@ -88,9 +94,11 @@ def test_admin_gui_sessions_includes_baton_holder_flag(client, mock_backend): ] mock_cfg.baton_holder = SimpleNamespace(session=22) - response = client.get("/admin/gui_sessions", headers={"Authorization": "Bearer fake-token"}) + response = client.get( + "/admin/gui_sessions", headers={"Authorization": "Bearer fake-token"} + ) assert response.status_code == 200 payload = response.json() assert payload[0]["session"] == 11 - assert payload[1]["session"] == 22 \ No newline at end of file + assert payload[1]["session"] == 22 diff --git a/tests/unit/daq/test_mlbox.py b/tests/unit/daq/test_mlbox.py index f977ac1f..5459204c 100644 --- a/tests/unit/daq/test_mlbox.py +++ b/tests/unit/daq/test_mlbox.py @@ -1,9 +1,11 @@ -import pytest -import numpy as np from unittest.mock import MagicMock, patch + +import numpy as np +import pytest +from aarecommon.models.beamline import MXBeamline +from aarecommon.models.models import MLBoxModel, MLBoxType, MLOutputModel + from aare.daq.mlbox import MlBox -from aare.common.beamline import MXBeamline -from aare.common.models import MLOutputModel, MLBoxModel, MLBoxType @pytest.fixture(autouse=True) @@ -126,4 +128,4 @@ def test_best_by_class_keeps_highest_confidence_per_class(): assert out is not None assert out.get_best_for_class(MLBoxType.PIN).conf == 0.8 - assert out.get_best_for_class(MLBoxType.CRYSTAL).conf == 0.6 \ No newline at end of file + assert out.get_best_for_class(MLBoxType.CRYSTAL).conf == 0.6 diff --git a/tests/unit/daq/test_mount.py b/tests/unit/daq/test_mount.py index 47f88dfe..4e5a1b46 100644 --- a/tests/unit/daq/test_mount.py +++ b/tests/unit/daq/test_mount.py @@ -3,14 +3,17 @@ from types import SimpleNamespace from unittest.mock import MagicMock import pytest +from aarecommon.errors.exception_handler import ( + BECCommunicationError, + JFJochCommunicationError, +) +from aarecommon.models.models import DAQOperation, DewarAddress, SampleShortInfo +from aarecommon.models.tell import TellActivityEnum, TellPhaseEnum, TellStateModel from aareDB import SampleEventType -from aare.common.exception_handler import BECCommunicationError, JFJochCommunicationError -from aare.common.models import DewarAddress, SampleShortInfo, DAQOperation -from aare.common.tell_models import TellActivityEnum, TellPhaseEnum, TellStateModel from aare.daq.config import ABR_POS_MOUNT, BeamlineStateEnum from aare.daq.daq import AareDAQ -from aare.daq.operations.mounting.models import MountingResult, MountingContext +from aare.daq.operations.mounting.models import MountingContext, MountingResult from aare.daq.operations.mounting.service import MountingService from aare.daq.operations.screenshot.service import ScreenshotService @@ -106,9 +109,17 @@ def test_execute_mount_and_prepare_success_uses_mounting_result_fields(): assert send_calls[3].args[0] == target_sample.db_id assert send_calls[3].args[1] == SampleEventType.MOUNTED - daq.save_screenshot_db.assert_called_once_with(target_sample.db_id, f"{target_sample.db_id}_mounted") - assert daq._AareDAQ__set_state.call_args_list[0].args[0] == BeamlineStateEnum.RobotSampleExchange - assert daq._AareDAQ__set_state.call_args_list[-1].args[0] == BeamlineStateEnum.SampleAlignment + daq.save_screenshot_db.assert_called_once_with( + target_sample.db_id, f"{target_sample.db_id}_mounted" + ) + assert ( + daq._AareDAQ__set_state.call_args_list[0].args[0] + == BeamlineStateEnum.RobotSampleExchange + ) + assert ( + daq._AareDAQ__set_state.call_args_list[-1].args[0] + == BeamlineStateEnum.SampleAlignment + ) def test_execute_mount_and_prepare_marks_previous_sample_unmounted_when_mount_fails_after_auto_unmount(): @@ -140,14 +151,24 @@ def test_execute_mount_and_prepare_marks_previous_sample_unmounted_when_mount_fa assert send_calls[2].args[0] == previous_sample.db_id assert send_calls[2].args[1] == SampleEventType.UNMOUNTED - assert send_calls[2].kwargs["comment"] == "Auto-unmount succeeded before mount failed" + assert ( + send_calls[2].kwargs["comment"] == "Auto-unmount succeeded before mount failed" + ) daq._handle_operation_error.assert_called_once() - assert daq._handle_operation_error.call_args.kwargs["operation"] == DAQOperation.MOUNT + assert ( + daq._handle_operation_error.call_args.kwargs["operation"] == DAQOperation.MOUNT + ) assert daq._handle_operation_error.call_args.kwargs["sample"] == target_sample - assert daq._handle_operation_error.call_args.kwargs["event_type"] == SampleEventType.MOUNTFAILED + assert ( + daq._handle_operation_error.call_args.kwargs["event_type"] + == SampleEventType.MOUNTFAILED + ) - assert daq._AareDAQ__set_state.call_args_list[-1].args[0] == BeamlineStateEnum.SampleAlignment + assert ( + daq._AareDAQ__set_state.call_args_list[-1].args[0] + == BeamlineStateEnum.SampleAlignment + ) def test_execute_mount_and_prepare_does_not_mark_previous_sample_unmounted_when_not_confirmed(): @@ -178,7 +199,10 @@ def test_execute_mount_and_prepare_does_not_mark_previous_sample_unmounted_when_ assert send_calls[1].args[1] == SampleEventType.MOUNTING daq._handle_operation_error.assert_called_once() - assert daq._handle_operation_error.call_args.kwargs["event_type"] == SampleEventType.MOUNTFAILED + assert ( + daq._handle_operation_error.call_args.kwargs["event_type"] + == SampleEventType.MOUNTFAILED + ) def test_execute_mount_and_prepare_unmount_success_uses_unmount_operation(): @@ -206,8 +230,14 @@ def test_execute_mount_and_prepare_unmount_success_uses_unmount_operation(): assert send_calls[1].args[1] == SampleEventType.UNMOUNTED daq.save_screenshot_db.assert_not_called() - assert daq._AareDAQ__set_state.call_args_list[0].args[0] == BeamlineStateEnum.RobotSampleExchange - assert daq._AareDAQ__set_state.call_args_list[-1].args[0] == BeamlineStateEnum.SampleAlignment + assert ( + daq._AareDAQ__set_state.call_args_list[0].args[0] + == BeamlineStateEnum.RobotSampleExchange + ) + assert ( + daq._AareDAQ__set_state.call_args_list[-1].args[0] + == BeamlineStateEnum.SampleAlignment + ) def test_raise_if_critical_jfjoch_detector_error_preserves_jfjoch_exception_family(): @@ -225,7 +255,10 @@ def test_raise_if_critical_jfjoch_detector_error_preserves_jfjoch_exception_fami daq._raise_if_critical_jfjoch_detector_error(original, command="wait_till_done") exc = exc_info.value - assert "Critical detector error while running JFJoch command 'wait_till_done'" in str(exc) + assert ( + "Critical detector error while running JFJoch command 'wait_till_done'" + in str(exc) + ) assert exc.operation == "measure" assert exc.endpoint == "/measurement/start" assert exc.base_url == "http://detector" @@ -247,10 +280,14 @@ def test_raise_if_critical_bec_error_preserves_bec_exception_family(): ) with pytest.raises(BECCommunicationError) as exc_info: - daq._raise_if_critical_bec_error(original, command="planner.move_to:data_collection") + daq._raise_if_critical_bec_error( + original, command="planner.move_to:data_collection" + ) exc = exc_info.value - assert "Critical BEC error while running 'planner.move_to:data_collection'" in str(exc) + assert "Critical BEC error while running 'planner.move_to:data_collection'" in str( + exc + ) assert exc.operation == "planner.move_to:data_collection" assert exc.endpoint == "/bec" assert exc.base_url == "redis://bec" @@ -269,7 +306,10 @@ def test_raise_if_critical_jfjoch_detector_error_ignores_non_critical_jfjoch_err status_code=503, ) - assert daq._raise_if_critical_jfjoch_detector_error(original, command="wait_till_done") is None + assert ( + daq._raise_if_critical_jfjoch_detector_error(original, command="wait_till_done") + is None + ) def test_create_loop_centering_service_uses_shared_screenshot_service(): @@ -305,4 +345,4 @@ def test_create_raster_service_uses_shared_screenshot_service(): service = daq._create_raster_service() - assert service.ctx.services.screenshots is daq._screenshot_service \ No newline at end of file + assert service.ctx.services.screenshots is daq._screenshot_service diff --git a/tests/unit/daq/test_raster_logic.py b/tests/unit/daq/test_raster_logic.py index 1dfb66b6..6099849f 100644 --- a/tests/unit/daq/test_raster_logic.py +++ b/tests/unit/daq/test_raster_logic.py @@ -1,13 +1,12 @@ import types - -import pytest from types import SimpleNamespace from unittest.mock import MagicMock +import pytest +from aarecommon.math.coordinate import Coordinate, SmargonCoordinate +from aarecommon.models.raster_grid import RasterGridRequest from jfjoch_client.exceptions import NotFoundException -from aare.common.coordinate import Coordinate, SmargonCoordinate -from aare.common.raster_grid import RasterGridRequest from aare.daq.operations.common.runtime import DAQRuntimeState from aare.daq.operations.common.services import OperationServices from aare.daq.operations.raster.models import ( @@ -18,7 +17,9 @@ from aare.daq.operations.raster.models import ( from aare.daq.operations.raster.service import RasterService -def make_request(n_x: int, n_y: int, cell_x: float = 0.01, cell_y: float = 0.02) -> RasterGridRequest: +def make_request( + n_x: int, n_y: int, cell_x: float = 0.01, cell_y: float = 0.02 +) -> RasterGridRequest: return RasterGridRequest( exp_time_s=0.02, transmission=1.0, @@ -223,4 +224,4 @@ def test_upload_raster_diffraction_preview_uploads_when_present(): ) service.ctx.deps.aare.upload_jpg.assert_called_once_with( 123, "preview", b"jpeg-bytes", message="Raster diffraction (479 spots, 2.23 Å)" - ) \ No newline at end of file + ) diff --git a/tests/unit/daq/test_server.py b/tests/unit/daq/test_server.py index 48ff1407..5b7b11f1 100644 --- a/tests/unit/daq/test_server.py +++ b/tests/unit/daq/test_server.py @@ -6,6 +6,7 @@ import numpy as np os.environ["JWT_AAREDAQ_KEY"] = "test_key_for_unit_testing" + def test_meta_error_codes(client): response = client.get("/meta/error-codes") assert response.status_code == 200 @@ -14,11 +15,16 @@ def test_meta_error_codes(client): def test_status(api, daq_status_factory, monkeypatch): - from aare.common.models import BeamlineStateEnum, SessionsStateEnum + from aarecommon.models.models import BeamlineStateEnum, SessionsStateEnum + from aare.daq import server - monkeypatch.setattr(server.auth, "resolve_baton_timeout_if_needed", lambda cfg: None) - monkeypatch.setattr(server.auth, "get_baton_status", lambda cfg, data: {"dummy": "status"}) + monkeypatch.setattr( + server.auth, "resolve_baton_timeout_if_needed", lambda cfg: None + ) + monkeypatch.setattr( + server.auth, "get_baton_status", lambda cfg, data: {"dummy": "status"} + ) api.cfg.pending_baton_request = None api.cfg.queued_baton_transfer = None @@ -49,38 +55,51 @@ def test_status(api, daq_status_factory, monkeypatch): assert data["session"]["current_pgroup"] == "p12345" assert data["session"]["session"] == SessionsStateEnum.OwnedByYou.value + def test_omega_put(client): with patch("aare.daq.auth.check_jwt_rw"), patch("aare.daq.server.daq") as mock_daq: - response = client.put("/beamline/omega?val=10.5", headers={"Authorization": "Bearer fake-token"}) + response = client.put( + "/beamline/omega?val=10.5", headers={"Authorization": "Bearer fake-token"} + ) assert response.status_code == 200 assert response.json() == "OK" assert mock_daq.omega == 10.5 def test_login_success(client): - with patch("aare.daq.auth.authenticate_from_proxy_header", return_value="user"), \ - patch("aare.daq.auth.authenticate_user", return_value="fake-access-token"): + with ( + patch("aare.daq.auth.authenticate_from_proxy_header", return_value="user"), + patch("aare.daq.auth.authenticate_user", return_value="fake-access-token"), + ): response = client.post( "/token", data={"username": "user", "password": "pwd"}, headers={"X-Remote-User": "user"}, ) assert response.status_code == 200 - assert response.json() == {"access_token": "fake-access-token", "token_type": "bearer"} + assert response.json() == { + "access_token": "fake-access-token", + "token_type": "bearer", + } def test_get_image(client, mock_backend): mock_daq = mock_backend["daq"] mock_daq.camera_image = np.zeros((100, 100, 3), dtype=np.uint8) - response = client.get("/beamline/image", headers={"Authorization": "Bearer fake-token"}) + response = client.get( + "/beamline/image", headers={"Authorization": "Bearer fake-token"} + ) assert response.status_code == 200 assert response.headers["content-type"] == "image/jpeg" assert len(response.content) > 0 -def test_mount_returns_tell_exception_when_mount_precheck_fails(client, mock_backend, monkeypatch): - from aare.common.exception_handler import TellCommunicationError +def test_mount_returns_tell_exception_when_mount_precheck_fails( + client, mock_backend, monkeypatch +): + from aarecommon.errors.exception_handler import TellCommunicationError + from aare.daq import server monkeypatch.setattr(server.auth, "check_jwt_rw", lambda *_args, **_kwargs: None) @@ -102,7 +121,9 @@ def test_mount_returns_tell_exception_when_mount_precheck_fails(client, mock_bac assert "Mount can't start:" in payload["message"] -def test_mount_calls_tell_mount_precheck_before_mount(client, mock_backend, monkeypatch): +def test_mount_calls_tell_mount_precheck_before_mount( + client, mock_backend, monkeypatch +): from aare.daq import server monkeypatch.setattr(server.auth, "check_jwt_rw", lambda *_args, **_kwargs: None) @@ -123,8 +144,11 @@ def test_mount_calls_tell_mount_precheck_before_mount(client, mock_backend, monk assert mock_daq.sample == sample -def test_auto_scan_returns_tell_exception_when_mount_precheck_fails(client, mock_backend, monkeypatch): - from aare.common.exception_handler import TellCommunicationError +def test_auto_scan_returns_tell_exception_when_mount_precheck_fails( + client, mock_backend, monkeypatch +): + from aarecommon.errors.exception_handler import TellCommunicationError + from aare.daq import server monkeypatch.setattr(server.auth, "check_jwt_rw", lambda *_args, **_kwargs: None) @@ -158,7 +182,9 @@ def test_auto_scan_returns_tell_exception_when_mount_precheck_fails(client, mock assert "Mount can't start:" in payload["message"] -def test_auto_scan_calls_tell_mount_precheck_before_measure(client, mock_backend, monkeypatch): +def test_auto_scan_calls_tell_mount_precheck_before_measure( + client, mock_backend, monkeypatch +): from aare.daq import server monkeypatch.setattr(server.auth, "check_jwt_rw", lambda *_args, **_kwargs: None) @@ -190,14 +216,18 @@ def test_auto_scan_calls_tell_mount_precheck_before_measure(client, mock_backend def test_get_pgroup(api): api.cfg.pgroup = "p12345" - response = api.client.get("/access/pgroup", headers={"Authorization": "Bearer fake-token"}) + response = api.client.get( + "/access/pgroup", headers={"Authorization": "Bearer fake-token"} + ) assert response.status_code == 200 assert response.json() == "p12345" def test_set_pgroup(api): api.cfg.baton_holder = None - response = api.client.put("/access/pgroup?val=p54321", headers={"Authorization": "Bearer fake-token"}) + response = api.client.put( + "/access/pgroup?val=p54321", headers={"Authorization": "Bearer fake-token"} + ) assert response.status_code == 200 assert response.json() == "OK" assert api.cfg.pgroup == "p54321" @@ -209,69 +239,101 @@ def test_delete_pgroup(api): def test_set_commissioning_mode(api): - response = api.client.put("/beamline/commissioning_mode?val=true", headers={"Authorization": "Bearer fake-token"}) + response = api.client.put( + "/beamline/commissioning_mode?val=true", + headers={"Authorization": "Bearer fake-token"}, + ) assert response.status_code == 200 assert response.json() == "OK" assert api.cfg.commissioning_mode is True def test_get_settings(api): - from aare.common.models import BeamlineSettingsModel + from aarecommon.models.models import BeamlineSettingsModel + mock_settings = BeamlineSettingsModel() api.cfg.settings = mock_settings - response = api.client.get("/beamline/settings", headers={"Authorization": "Bearer fake-token"}) + response = api.client.get( + "/beamline/settings", headers={"Authorization": "Bearer fake-token"} + ) assert response.status_code == 200 assert response.json() == mock_settings.model_dump() def test_put_settings(api): - from aare.common.models import BeamlineSettingsModel + from aarecommon.models.models import BeamlineSettingsModel + settings_data = BeamlineSettingsModel().model_dump() - response = api.client.put("/beamline/settings", json=settings_data, headers={"Authorization": "Bearer fake-token"}) + response = api.client.put( + "/beamline/settings", + json=settings_data, + headers={"Authorization": "Bearer fake-token"}, + ) assert response.status_code == 200 assert api.cfg.settings.model_dump() == settings_data def test_get_cryo_settings(api): - from aare.common.models import CryojetSettingsModel + from aarecommon.models.models import CryojetSettingsModel + mock_cryo = CryojetSettingsModel() api.cfg.cryojet_settings = mock_cryo - response = api.client.get("/beamline/cryo_settings", headers={"Authorization": "Bearer fake-token"}) + response = api.client.get( + "/beamline/cryo_settings", headers={"Authorization": "Bearer fake-token"} + ) assert response.status_code == 200 assert response.json() == mock_cryo.model_dump() def test_put_cryo_settings(api): - from aare.common.models import CryojetSettingsModel + from aarecommon.models.models import CryojetSettingsModel + cryo_data = CryojetSettingsModel().model_dump() - response = api.client.put("/beamline/cryo_settings", json=cryo_data, headers={"Authorization": "Bearer fake-token"}) + response = api.client.put( + "/beamline/cryo_settings", + json=cryo_data, + headers={"Authorization": "Bearer fake-token"}, + ) assert response.status_code == 200 assert api.cfg.cryojet_settings.model_dump() == cryo_data def test_baton_status(api): - from aare.common.auth_models import BatonStatus + from aarecommon.models.auth import BatonStatus + mock_baton = BatonStatus(holder=None, request=None, allow_non_staff_request=True) api.cfg.baton_status = mock_baton api.cfg.baton_holder = None api.cfg.queued_baton_transfer = None api.cfg.allow_non_staff_request_from_staff = True - response = api.client.get("/baton/status", headers={"Authorization": "Bearer fake-token"}) + response = api.client.get( + "/baton/status", headers={"Authorization": "Bearer fake-token"} + ) assert response.status_code == 200 assert response.json() == mock_baton.model_dump() def test_baton_request(api, monkeypatch): from aare.daq import server - monkeypatch.setattr(server.auth, "request_baton", lambda cfg, data: {"granted": True}) - response = api.client.post("/baton/request", headers={"Authorization": "Bearer fake-token"}) + + monkeypatch.setattr( + server.auth, "request_baton", lambda cfg, data: {"granted": True} + ) + response = api.client.post( + "/baton/request", headers={"Authorization": "Bearer fake-token"} + ) assert response.status_code == 200 assert response.json() == {"granted": True} def test_baton_release(api, monkeypatch): from aare.daq import server - monkeypatch.setattr(server.auth, "release_baton", lambda cfg, data: {"released": True}) - response = api.client.post("/baton/release", headers={"Authorization": "Bearer fake-token"}) + + monkeypatch.setattr( + server.auth, "release_baton", lambda cfg, data: {"released": True} + ) + response = api.client.post( + "/baton/release", headers={"Authorization": "Bearer fake-token"} + ) assert response.status_code == 200 - assert response.json() == {"released": True} \ No newline at end of file + assert response.json() == {"released": True} diff --git a/tests/unit/daq/test_server_exception_handler.py b/tests/unit/daq/test_server_exception_handler.py index b4a11305..45287694 100644 --- a/tests/unit/daq/test_server_exception_handler.py +++ b/tests/unit/daq/test_server_exception_handler.py @@ -10,33 +10,33 @@ bare-Exception fallback. Each emits the unified response body from §3: from __future__ import annotations import json +from unittest.mock import MagicMock import pytest +from aarecommon.errors.codes import AareErrorCode, AuthErrorCode +from aarecommon.errors.exception_handler import ( + AareAuthError, + AareDBCommunicationError, + AareException, + AareUserError, + AuthenticationException, + AutomationError, + CriticalTellException, + LoopCenteringFailed, + ManualMountException, + MountingFailed, + SampleException, + SmargonCommunicationError, + TellCommunicationError, + UnmountingFailed, + UserRightsException, + WarningTellException, +) from fastapi import FastAPI, HTTPException from starlette.requests import Request from starlette.responses import JSONResponse -from unittest.mock import MagicMock from aare.daq.server_exception_handler import register_exception_handlers -from aare.common.error_codes import AareErrorCode, AuthErrorCode -from aare.common.exception_handler import ( - AareException, - AutomationError, - AareUserError, - AareAuthError, - MountingFailed, - UnmountingFailed, - LoopCenteringFailed, - TellCommunicationError, - CriticalTellException, - WarningTellException, - SmargonCommunicationError, - AareDBCommunicationError, - AuthenticationException, - UserRightsException, - ManualMountException, - SampleException, -) @pytest.fixture @@ -60,6 +60,7 @@ def _body(response: JSONResponse) -> dict: # AutomationError handler -- 503 if critical, 422 otherwise # --------------------------------------------------------------------------- + @pytest.mark.asyncio async def test_automation_error_not_critical_returns_422(app, mock_request): # MountingFailed has class default critical=False in Phase 1 @@ -76,7 +77,9 @@ async def test_automation_error_not_critical_returns_422(app, mock_request): @pytest.mark.asyncio -async def test_automation_error_critical_via_instance_flag_returns_503(app, mock_request): +async def test_automation_error_critical_via_instance_flag_returns_503( + app, mock_request +): # Instance-level override: critical=True exc = MountingFailed("threshold breached", critical=True) handler = app.exception_handlers[AutomationError] @@ -87,9 +90,12 @@ async def test_automation_error_critical_via_instance_flag_returns_503(app, mock @pytest.mark.asyncio -async def test_automation_error_context_contains_endpoint_and_operation(app, mock_request): - exc = TellCommunicationError("timeout", endpoint="/state", operation="GET", - base_url="http://tell:8000") +async def test_automation_error_context_contains_endpoint_and_operation( + app, mock_request +): + exc = TellCommunicationError( + "timeout", endpoint="/state", operation="GET", base_url="http://tell:8000" + ) handler = app.exception_handlers[AutomationError] response = await handler(mock_request, exc) body = _body(response) @@ -100,7 +106,9 @@ async def test_automation_error_context_contains_endpoint_and_operation(app, moc @pytest.mark.asyncio -async def test_automation_error_excludes_critical_and_headers_from_context(app, mock_request): +async def test_automation_error_excludes_critical_and_headers_from_context( + app, mock_request +): # critical kwarg should not appear in context (it has its own field) exc = MountingFailed("x", critical=True) handler = app.exception_handlers[AutomationError] @@ -114,6 +122,7 @@ async def test_automation_error_excludes_critical_and_headers_from_context(app, # AareUserError handler -- always 400 # --------------------------------------------------------------------------- + @pytest.mark.asyncio async def test_user_error_returns_400(app, mock_request): exc = ManualMountException("user must intervene") @@ -140,6 +149,7 @@ async def test_sample_exception_returns_400_as_user_error(app, mock_request): # AareAuthError handler -- 401 / 403, code from instance.code # --------------------------------------------------------------------------- + @pytest.mark.asyncio async def test_authentication_exception_returns_401_with_auth_code(app, mock_request): exc = AuthenticationException("Bad token", code=AuthErrorCode.INVALID_TOKEN) @@ -166,8 +176,9 @@ async def test_user_rights_exception_returns_403(app, mock_request): @pytest.mark.asyncio async def test_authentication_exception_preserves_explicit_status(app, mock_request): # AuthenticationException can be constructed with a custom status_code - exc = AuthenticationException("Forbidden auth path", status_code=403, - code=AuthErrorCode.FORBIDDEN) + exc = AuthenticationException( + "Forbidden auth path", status_code=403, code=AuthErrorCode.FORBIDDEN + ) handler = app.exception_handlers[AareAuthError] response = await handler(mock_request, exc) assert response.status_code == 403 @@ -177,6 +188,7 @@ async def test_authentication_exception_preserves_explicit_status(app, mock_requ # HTTPException handler -- pass-through status, new shape # --------------------------------------------------------------------------- + @pytest.mark.asyncio async def test_http_exception_handler_string_detail(app, mock_request): exc = HTTPException(status_code=418, detail="I'm a teapot") @@ -192,7 +204,9 @@ async def test_http_exception_handler_string_detail(app, mock_request): @pytest.mark.asyncio async def test_http_exception_handler_dict_detail(app, mock_request): - exc = HTTPException(status_code=400, detail={"code": "CUSTOM", "message": "Msg", "field": "foo"}) + exc = HTTPException( + status_code=400, detail={"code": "CUSTOM", "message": "Msg", "field": "foo"} + ) handler = app.exception_handlers[HTTPException] response = await handler(mock_request, exc) assert response.status_code == 400 @@ -215,6 +229,7 @@ async def test_http_exception_5xx_is_critical(app, mock_request): # Bare Exception fallback -- 500, critical=True # --------------------------------------------------------------------------- + @pytest.mark.asyncio async def test_unhandled_exception_is_critical_500(app, mock_request): exc = ValueError("oops") @@ -229,7 +244,9 @@ async def test_unhandled_exception_is_critical_500(app, mock_request): @pytest.mark.asyncio -async def test_unhandled_exception_empty_message_falls_back_to_class_name(app, mock_request): +async def test_unhandled_exception_empty_message_falls_back_to_class_name( + app, mock_request +): exc = ValueError() handler = app.exception_handlers[Exception] response = await handler(mock_request, exc) @@ -241,6 +258,7 @@ async def test_unhandled_exception_empty_message_falls_back_to_class_name(app, m # Handler count -- the design contract is 4 + HTTPException + Exception # --------------------------------------------------------------------------- + def test_only_four_root_handlers_plus_fallbacks(app): """Plan §2: collapse the ~25 per-class handlers down to 4 roots, plus the HTTPException pass-through and the bare-Exception fallback. @@ -250,24 +268,34 @@ def test_only_four_root_handlers_plus_fallbacks(app): counted toward our handler budget. We assert (a) our expected handlers are present and (b) no per-class AareException handlers remain. """ - from aare.common.exception_handler import ( - MountingFailed, UnmountingFailed, TellCommunicationError, - LoopCenteringFailed, CriticalTellException, WarningTellException, - SmargonCommunicationError, AareDBCommunicationError, + from aarecommon.errors.exception_handler import ( + AareDBCommunicationError, + CriticalTellException, + LoopCenteringFailed, + MountingFailed, + SmargonCommunicationError, + TellCommunicationError, + UnmountingFailed, + WarningTellException, ) registered = set(app.exception_handlers.keys()) - expected = {AutomationError, AareUserError, AareAuthError, - HTTPException, Exception} + expected = {AutomationError, AareUserError, AareAuthError, HTTPException, Exception} assert expected.issubset(registered), ( f"Missing required handlers: {expected - registered}" ) # No leftover per-class handlers from the old design - forbidden = {MountingFailed, UnmountingFailed, TellCommunicationError, - LoopCenteringFailed, CriticalTellException, - WarningTellException, SmargonCommunicationError, - AareDBCommunicationError} + forbidden = { + MountingFailed, + UnmountingFailed, + TellCommunicationError, + LoopCenteringFailed, + CriticalTellException, + WarningTellException, + SmargonCommunicationError, + AareDBCommunicationError, + } leftover = forbidden & registered assert not leftover, f"Per-class handlers must be removed: {leftover}" @@ -290,7 +318,9 @@ async def test_automation_error_logs_error_when_critical(app, mock_request, capl @pytest.mark.asyncio -async def test_automation_error_logs_warning_when_not_critical(app, mock_request, caplog): +async def test_automation_error_logs_warning_when_not_critical( + app, mock_request, caplog +): exc = MountingFailed("benign") handler = app.exception_handlers[AutomationError] caplog.clear() @@ -315,7 +345,10 @@ async def test_response_body_shape_uniform(app, mock_request): cases: list[tuple[type, Exception]] = [ (AutomationError, MountingFailed("x")), (AutomationError, LoopCenteringFailed("y")), - (AutomationError, SmargonCommunicationError("z", operation="GET", endpoint="/e")), + ( + AutomationError, + SmargonCommunicationError("z", operation="GET", endpoint="/e"), + ), (AutomationError, AareDBCommunicationError("db", critical=True)), (AareUserError, ManualMountException("m")), (AareUserError, SampleException("s")), diff --git a/tests/unit/daq/test_spreadsheetupdater.py b/tests/unit/daq/test_spreadsheetupdater.py index 48affc5c..9f36e93e 100644 --- a/tests/unit/daq/test_spreadsheetupdater.py +++ b/tests/unit/daq/test_spreadsheetupdater.py @@ -1,14 +1,20 @@ -import pytest import json from types import SimpleNamespace from unittest.mock import MagicMock, patch -from aare.daq.spreadsheetupdater import on_message, get_ws_headers, set_spreadsheet_in_redis -from aare.common.models import SampleShortInfoList + +import pytest +from aarecommon.models.models import SampleShortInfoList + +from aare.daq.spreadsheetupdater import ( + get_ws_headers, + on_message, + set_spreadsheet_in_redis, +) @pytest.fixture def mock_config(): - with patch('aare.daq.spreadsheetupdater.config') as mock: + with patch("aare.daq.spreadsheetupdater.config") as mock: mock._BeamlineConfig__bl = "X10SA" mock._BeamlineConfig__client = MagicMock() # Mocking private attributes access which the code uses @@ -18,20 +24,20 @@ def mock_config(): def test_get_ws_headers_success(): - with patch('os.getenv', return_value="secret"): + with patch("os.getenv", return_value="secret"): headers = get_ws_headers() assert headers == ["X-Shared-Password: secret"] def test_get_ws_headers_fail(): - with patch('os.getenv', return_value=None): + with patch("os.getenv", return_value=None): with pytest.raises(ValueError): get_ws_headers() def test_set_spreadsheet_in_redis(mock_config): data = {"test": "data"} - with patch('aare.daq.spreadsheetupdater.config') as mock_cfg_internal: + with patch("aare.daq.spreadsheetupdater.config") as mock_cfg_internal: mock_client = MagicMock() mock_cfg_internal._BeamlineConfig__client = mock_client mock_cfg_internal.client = mock_client @@ -92,7 +98,9 @@ def test_on_message_success(mock_config): message = json.dumps({"samples": [{}, {}]}) - with patch("aare.daq.spreadsheetupdater.PuckWithTellPosition", side_effect=mock_pucks): + with patch( + "aare.daq.spreadsheetupdater.PuckWithTellPosition", side_effect=mock_pucks + ): on_message(None, message) calls = mock_config._BeamlineConfig__client.set.call_args_list @@ -103,23 +111,25 @@ def test_on_message_success(mock_config): def test_on_message_empty_ref(mock_config): - message = json.dumps({ - "samples": [ - { - "id": 1, - "barcode": "B1", - "position": "P1", - "puck_name": "P1", - "puck_type": "UniPuck", - "puck_location_in_dewar": 1, - "dewar_id": 1, - "pgroup": "p12345", - "dewar_name": "D1", - "tell_position": "A1", - "samples": [] - } - ] - }) + message = json.dumps( + { + "samples": [ + { + "id": 1, + "barcode": "B1", + "position": "P1", + "puck_name": "P1", + "puck_type": "UniPuck", + "puck_location_in_dewar": 1, + "dewar_id": 1, + "pgroup": "p12345", + "dewar_name": "D1", + "tell_position": "A1", + "samples": [], + } + ] + } + ) on_message(None, message) @@ -129,4 +139,4 @@ def test_on_message_empty_ref(mock_config): def test_on_message_invalid_json(mock_config): on_message(None, "invalid json") - mock_config._BeamlineConfig__client.set.assert_not_called() \ No newline at end of file + mock_config._BeamlineConfig__client.set.assert_not_called() diff --git a/tests/unit/daq/test_tell_state_updater.py b/tests/unit/daq/test_tell_state_updater.py index 0d7aafaa..7b3664d1 100644 --- a/tests/unit/daq/test_tell_state_updater.py +++ b/tests/unit/daq/test_tell_state_updater.py @@ -1,4 +1,5 @@ -from aare.common.tell_models import TellActivityEnum, TellPhaseEnum +from aarecommon.models.tell import TellActivityEnum, TellPhaseEnum + from aare.daq.tell_state_machine import advance_tell_state, initial_tell_state @@ -155,4 +156,4 @@ def test_mount_phase_after_old_sample_returned_confirms_auto_unmount_completed() assert state.phase == TellPhaseEnum.PLACING_NEW_SAMPLE state = advance_tell_state(state, "Motion Sync", "Robot Clear after mount") - assert state.phase == TellPhaseEnum.FINALIZING \ No newline at end of file + assert state.phase == TellPhaseEnum.FINALIZING diff --git a/tests/unit/daq/test_workflows.py b/tests/unit/daq/test_workflows.py index 0b80137b..42670b0e 100644 --- a/tests/unit/daq/test_workflows.py +++ b/tests/unit/daq/test_workflows.py @@ -1,10 +1,27 @@ from types import SimpleNamespace +from unittest.mock import MagicMock, patch import pytest -from unittest.mock import MagicMock, patch -from aare.daq.workflows import common_2rse, sa2se, sa2rse, sa2xtal_snapshot, dc2xtal_snapshot, xtal_snapshot2dc, xtal_snapshot2sa, dc2rse, se2sa, sa2dc, dc2sa, sa2xrf, sa2dh, dh2sa, common2dh -from aare.daq.config import ABR_POS_MOUNT, ABR_OMEGA_MOUNT -from aare.common.models import StagePositionEnum +from aarecommon.models.models import StagePositionEnum + +from aare.daq.config import ABR_OMEGA_MOUNT, ABR_POS_MOUNT +from aare.daq.workflows import ( + common2dh, + common_2rse, + dc2rse, + dc2sa, + dc2xtal_snapshot, + dh2sa, + sa2dc, + sa2dh, + sa2rse, + sa2se, + sa2xrf, + sa2xtal_snapshot, + se2sa, + xtal_snapshot2dc, + xtal_snapshot2sa, +) from aare.devices.area_detector import AutoEnum from aare.devices.bec_worker import BeamlineState @@ -29,16 +46,20 @@ def mock_cfg(): def _assert_bec_moved(devs, state): - planner_calls = devs.bec_worker.planner.move_to.call_args_list if devs.bec_worker is not None else [] - direct_calls = devs.bec_worker.move_to.call_args_list if devs.bec_worker is not None else [] + planner_calls = ( + devs.bec_worker.planner.move_to.call_args_list + if devs.bec_worker is not None + else [] + ) + direct_calls = ( + devs.bec_worker.move_to.call_args_list if devs.bec_worker is not None else [] + ) - assert ( - any(call.args == (state,) for call in planner_calls) - or any(call.args == (state,) for call in direct_calls) + assert any(call.args == (state,) for call in planner_calls) or any( + call.args == (state,) for call in direct_calls ), f"BEC was not asked to move to {state}" - def test_common_2rse(mock_devs, mock_cfg): mock_devs.bec_worker = MagicMock() mock_devs.bec_worker.planner = MagicMock() @@ -51,7 +72,9 @@ def test_common_2rse(mock_devs, mock_cfg): assert mock_devs.aerotech_pos == ABR_POS_MOUNT -def test_common_2rse_moves_detector_to_safe_position_when_configured(mock_devs, mock_cfg): +def test_common_2rse_moves_detector_to_safe_position_when_configured( + mock_devs, mock_cfg +): mock_devs.dtz = 200 mock_cfg.dtz_safe_position = 600 @@ -224,4 +247,4 @@ def test_common2dh_skips_dry_when_tell_already_in_ppark(mock_devs, mock_cfg): mock_devs.tell.is_position.assert_called_once_with("pPark") mock_devs.tell.get_mounted_sample.assert_called_once_with() mock_devs.tell.unmount.assert_not_called() - mock_devs.tell.dry.assert_not_called() \ No newline at end of file + mock_devs.tell.dry.assert_not_called() diff --git a/tests/unit/devices/test_aerotech.py b/tests/unit/devices/test_aerotech.py index eda91eb3..5bc97208 100644 --- a/tests/unit/devices/test_aerotech.py +++ b/tests/unit/devices/test_aerotech.py @@ -1,23 +1,30 @@ -import pytest from unittest.mock import MagicMock, patch -from aare.devices.aerotech import AerotechController, AEROTECH_HOME -from aare.common.beamline import MXBeamline -from aare.common.coordinate import AerotechCoordinate, Coordinate -from aare.common.exception_handler import AerotechCommunicationError + +import pytest +from aarecommon.errors.exception_handler import AerotechCommunicationError +from aarecommon.math.coordinate import AerotechCoordinate, Coordinate +from aarecommon.models.beamline import MXBeamline + +from aare.devices.aerotech import AEROTECH_HOME, AerotechController @pytest.fixture def mock_aerotech_api(): - with patch('aarescan_client.ApiClient'), \ - patch('aarescan_client.DefaultApi') as mock_api_class, \ - patch('aarescan_client.Configuration'): + with ( + patch("aarescan_client.ApiClient"), + patch("aarescan_client.DefaultApi") as mock_api_class, + patch("aarescan_client.Configuration"), + ): mock_api = mock_api_class.return_value yield mock_api @pytest.fixture def aerotech_controller(mock_aerotech_api): - with patch("aare.devices.aerotech.cfg_get", return_value="http://mx-x10sa-queue-01.psi.ch:5234"): + with patch( + "aare.devices.aerotech.cfg_get", + return_value="http://mx-x10sa-queue-01.psi.ch:5234", + ): controller = AerotechController(MXBeamline.X10SA) controller._AerotechController__api = mock_aerotech_api @@ -26,9 +33,14 @@ def aerotech_controller(mock_aerotech_api): def test_init_x10sa(mock_aerotech_api): - with patch("aare.devices.aerotech.cfg_get", return_value="http://mx-x10sa-queue-01.psi.ch:5234"): + with patch( + "aare.devices.aerotech.cfg_get", + return_value="http://mx-x10sa-queue-01.psi.ch:5234", + ): controller = AerotechController(MXBeamline.X10SA) - assert controller._AerotechController__base == "http://mx-x10sa-queue-01.psi.ch:5234" + assert ( + controller._AerotechController__base == "http://mx-x10sa-queue-01.psi.ch:5234" + ) assert controller._AerotechController__simulated is False @@ -43,10 +55,10 @@ def test_cancel(aerotech_controller, mock_aerotech_api): def test_is_idle(aerotech_controller, mock_aerotech_api): - mock_aerotech_api.status_get.return_value.state = 'Idle' + mock_aerotech_api.status_get.return_value.state = "Idle" assert aerotech_controller.is_idle() is True - mock_aerotech_api.status_get.return_value.state = 'Busy' + mock_aerotech_api.status_get.return_value.state = "Busy" assert aerotech_controller.is_idle() is False @@ -82,17 +94,24 @@ def test_rotation_scan(aerotech_controller, mock_aerotech_api): def test_grid_scan(aerotech_controller, mock_aerotech_api): - aerotech_controller.grid_scan(grid_elem_count_y=10, grid_elem_size_y_um=10, grid_elem_count_x=10, - grid_elem_size_x_um=10, time_sec=5) + aerotech_controller.grid_scan( + grid_elem_count_y=10, + grid_elem_size_y_um=10, + grid_elem_count_x=10, + grid_elem_size_x_um=10, + time_sec=5, + ) mock_aerotech_api.grid_scan_post.assert_called_once() def test_screening_scan(aerotech_controller, mock_aerotech_api): - aerotech_controller.screening_scan(rotation_deg=10, wedge_deg=2, time_sec=1, steps=5) + aerotech_controller.screening_scan( + rotation_deg=10, wedge_deg=2, time_sec=1, steps=5 + ) mock_aerotech_api.screening_post.assert_called_once() def test_api_error(aerotech_controller, mock_aerotech_api): mock_aerotech_api.status_get.side_effect = Exception("API Error") with pytest.raises(AerotechCommunicationError): - aerotech_controller.is_idle() \ No newline at end of file + aerotech_controller.is_idle() diff --git a/tests/unit/devices/test_experimental_hutch_shutter.py b/tests/unit/devices/test_experimental_hutch_shutter.py index df4e06ec..b03ecf61 100644 --- a/tests/unit/devices/test_experimental_hutch_shutter.py +++ b/tests/unit/devices/test_experimental_hutch_shutter.py @@ -1,8 +1,11 @@ -import pytest from unittest.mock import MagicMock, patch -from aare.common.beamline import MXBeamline + +import pytest +from aarecommon.models.beamline import MXBeamline + from aare.devices.experimental_hutch_shutter import ExperimentalHutchShutter + @patch("aare.devices.experimental_hutch_shutter.PV") def test_shutter_init(mock_pv): shutter = ExperimentalHutchShutter(MXBeamline.X06DA) @@ -13,43 +16,45 @@ def test_shutter_init(mock_pv): assert "X06DA-EH1-PSYS:SH-A-OPEN-SET" in args assert "X06DA-OP-PSH1-EMLS-0010:OPEN" in args + @patch("aare.devices.experimental_hutch_shutter.PV") def test_shutter_state(mock_pv): # Mocking self.__state mock_state_pv = MagicMock() mock_pv.side_effect = [MagicMock(), MagicMock(), mock_state_pv] - + shutter = ExperimentalHutchShutter(MXBeamline.X06DA) - + mock_state_pv.get.return_value = "Open" assert shutter.state() is True - + mock_state_pv.get.return_value = 1 assert shutter.state() is True - + mock_state_pv.get.return_value = "Not Open" assert shutter.state() is False - + mock_state_pv.get.return_value = 0 assert shutter.state() is False - + mock_state_pv.get.return_value = "Unknown" assert shutter.state() is False + @patch("aare.devices.experimental_hutch_shutter.PV") def test_shutter_open_close(mock_pv): mock_close = MagicMock() mock_open = MagicMock() mock_state = MagicMock() mock_pv.side_effect = [mock_close, mock_open, mock_state] - + shutter = ExperimentalHutchShutter(MXBeamline.X06DA) - + shutter.open() assert mock_open.put.call_count == 3 mock_open.put.assert_any_call(0) mock_open.put.assert_any_call(1) - + shutter.close() assert mock_close.put.call_count == 2 mock_close.put.assert_any_call(1) diff --git a/tests/unit/devices/test_filter_transmission.py b/tests/unit/devices/test_filter_transmission.py index 8d1bc3ea..e5d8eb9b 100644 --- a/tests/unit/devices/test_filter_transmission.py +++ b/tests/unit/devices/test_filter_transmission.py @@ -1,28 +1,33 @@ -import pytest from unittest.mock import MagicMock, patch -from aare.common.beamline import MXBeamline + +import pytest +from aarecommon.models.beamline import MXBeamline + from aare.devices.filter_transmission import FilterTransmission + @patch("aare.devices.filter_transmission.PV") def test_filter_init(mock_pv): filters = FilterTransmission(MXBeamline.X06DA) assert mock_pv.call_count == 3 + @patch("aare.devices.filter_transmission.PV") def test_filter_repr_str(mock_pv): mock_set = MagicMock() mock_get = MagicMock() mock_done = MagicMock() mock_pv.side_effect = [mock_set, mock_get, mock_done] - + filters = FilterTransmission(MXBeamline.X06DA) - + mock_set.get.return_value = 0.5 assert "" in repr(filters) - + mock_get.get.return_value = 0.5 assert "0.5000" in str(filters) + @patch("aare.devices.filter_transmission.poll") @patch("aare.devices.filter_transmission.PV") def test_filter_set_wait(mock_pv, mock_poll): @@ -30,34 +35,35 @@ def test_filter_set_wait(mock_pv, mock_poll): mock_get = MagicMock() mock_done = MagicMock() mock_pv.side_effect = [mock_set, mock_get, mock_done] - + filters = FilterTransmission(MXBeamline.X06DA, timeout=0.1) - + # Test set without wait filters.set(0.5) mock_set.put.assert_called_with(0.5) - + # Test set with invalid value with pytest.raises(RuntimeError, match="out of bounds"): filters.set(1.5) # Test set with wait and timeout - mock_done.value = 0 # Busy + mock_done.value = 0 # Busy with pytest.raises(RuntimeError, match="timeout"): filters.set(0.5, wait=True) + @patch("aare.devices.filter_transmission.PV") def test_filter_get(mock_pv): mock_set = MagicMock() mock_get = MagicMock() mock_done = MagicMock() mock_pv.side_effect = [mock_set, mock_get, mock_done] - + filters = FilterTransmission(MXBeamline.X06DA) - - mock_done.value = 1 # Done + + mock_done.value = 1 # Done mock_get.value = 0.123456 - assert filters.get() == 0.12346 # Rounded to 5 - - mock_done.value = 0 # Busy + assert filters.get() == 0.12346 # Rounded to 5 + + mock_done.value = 0 # Busy assert filters.get() is None diff --git a/tests/unit/devices/test_fluorimeter.py b/tests/unit/devices/test_fluorimeter.py index 080ebeb0..5aae0587 100644 --- a/tests/unit/devices/test_fluorimeter.py +++ b/tests/unit/devices/test_fluorimeter.py @@ -1,74 +1,82 @@ -import pytest from unittest.mock import MagicMock, patch -from aare.common.beamline import MXBeamline + +import pytest +from aarecommon.models.beamline import MXBeamline + from aare.devices.fluorimeter import Fluorimeter + @patch("aare.devices.fluorimeter.PV") def test_fluorimeter_init(mock_pv): fluo = Fluorimeter(MXBeamline.X06DA) # Lots of PVs in __init__ assert mock_pv.call_count >= 20 + @patch("aare.devices.fluorimeter.PV") def test_fluorimeter_acquisition(mock_pv): mock_start = MagicMock() mock_stop = MagicMock() mock_erase_start = MagicMock() - - # We need to map which mock is which. + + # We need to map which mock is which. # __init__ order: start, stop, erase_and_start, erase, ... mock_pv.side_effect = [mock_start, mock_stop, mock_erase_start] + [MagicMock()] * 50 - + fluo = Fluorimeter(MXBeamline.X06DA) - + fluo.start_acquisition(erase=False) mock_start.put.assert_called_with(1) - + fluo.start_acquisition(erase=True) mock_erase_start.put.assert_called_with(1) - + fluo.stop_acquisition() mock_stop.put.assert_called_with(1) + @patch("aare.devices.fluorimeter.PV") def test_fluorimeter_status(mock_pv): mock_status = MagicMock() # status is the 6th PV in __init__ mock_pv.side_effect = [MagicMock()] * 5 + [mock_status] + [MagicMock()] * 50 fluo = Fluorimeter(MXBeamline.X06DA) - + mock_status.get.return_value = 0 assert fluo.check_status_done() is True assert fluo.check_status_acquiring() is None - + mock_status.get.return_value = 1 assert fluo.check_status_done() is None assert fluo.check_status_acquiring() is True + @patch("aare.devices.fluorimeter.poll") @patch("aare.devices.fluorimeter.PV") def test_fluorimeter_wait_timeout(mock_pv, mock_poll): mock_status = MagicMock() mock_pv.side_effect = [MagicMock()] * 5 + [mock_status] + [MagicMock()] * 50 fluo = Fluorimeter(MXBeamline.X06DA) - - mock_status.get.return_value = 1 # Acquiring (not done) + + mock_status.get.return_value = 1 # Acquiring (not done) with pytest.raises(TimeoutError, match="timeout waiting for done"): fluo.wait_till_done(timeout_s=0.1) + @patch("aare.devices.fluorimeter.PV") def test_fluorimeter_properties(mock_pv): mock_real_time = MagicMock() # real_time is 7th PV mock_pv.side_effect = [MagicMock()] * 6 + [mock_real_time] + [MagicMock()] * 50 fluo = Fluorimeter(MXBeamline.X06DA) - + mock_real_time.get.return_value = 10.0 assert fluo.real_time == 10.0 - + fluo.real_time = 20.0 mock_real_time.put.assert_called_with(20.0) + def test_fluorimeter_roi_error(): with patch("aare.devices.fluorimeter.PV"): fluo = Fluorimeter(MXBeamline.X06DA) diff --git a/tests/unit/devices/test_jfjoch.py b/tests/unit/devices/test_jfjoch.py index c6d2a9e2..8498cc6d 100644 --- a/tests/unit/devices/test_jfjoch.py +++ b/tests/unit/devices/test_jfjoch.py @@ -1,21 +1,33 @@ -import pytest from unittest.mock import MagicMock, patch + +import pytest +from aarecommon.errors.exception_handler import JFJochCommunicationError +from aarecommon.math.coordinate import Coordinate, SmargonCoordinate +from aarecommon.math.diffraction_geometry import DiffractionGeometry +from aarecommon.math.sample_geometry import SampleGeometryModel +from aarecommon.models.beamline import MXBeamline +from aarecommon.models.models import ( + BeamlineStateEnum, + BeamlineStatus, + DAQStatusModel, + SampleCameraSettings, + SampleShortInfo, + SessionsStateEnum, + SessionStatus, +) +from aarecommon.models.raster_grid import RasterGridRequest +from aarecommon.models.rotation_scan import RotationScanRequest + from aare.devices.jfjoch import JFJochWrapper -from aare.common.beamline import MXBeamline -from aare.common.rotation_scan import RotationScanRequest -from aare.common.raster_grid import RasterGridRequest -from aare.common.models import DAQStatusModel, SampleShortInfo, BeamlineStatus, SessionStatus, BeamlineStateEnum, SessionsStateEnum, SampleCameraSettings -from aare.common.diffraction_geometry import DiffractionGeometry -from aare.common.sample_geometry import SampleGeometryModel -from aare.common.coordinate import Coordinate, SmargonCoordinate -from aare.common.exception_handler import JFJochCommunicationError @pytest.fixture def mock_jfjoch_client(): - with patch('jfjoch_client.ApiClient'), \ - patch('jfjoch_client.DefaultApi') as mock_api_class, \ - patch('jfjoch_client.Configuration'): + with ( + patch("jfjoch_client.ApiClient"), + patch("jfjoch_client.DefaultApi") as mock_api_class, + patch("jfjoch_client.Configuration"), + ): mock_api = mock_api_class.return_value yield mock_api @@ -40,7 +52,7 @@ def create_mock_daq_status(): zoom=1.0, commissioning_mode=False, dtz_min=120.0, - dtz_max=1600.0 + dtz_max=1600.0, ) diff_geom = DiffractionGeometry( energy_keV=12.658, @@ -51,7 +63,7 @@ def create_mock_daq_status(): detector_description="Eiger", detector_serial_number="123", poni_rot1_rad=0.0, - poni_rot2_rad=0.0 + poni_rot2_rad=0.0, ) session = SessionStatus(session=SessionsStateEnum.Vacant, current_pgroup="p12345") @@ -62,7 +74,7 @@ def create_mock_daq_status(): aerotech_meas=Coordinate(x=0, y=0), smargon=SmargonCoordinate(sh_mm=Coordinate(x=0, y=0), phi_deg=0, chi_deg=0), omega_deg=0.0, - beam_size_mm=Coordinate(x=0.01, y=0.01) + beam_size_mm=Coordinate(x=0.01, y=0.01), ) return DAQStatusModel( @@ -71,7 +83,7 @@ def create_mock_daq_status(): bl=bl_status, state=BeamlineStateEnum.SampleAlignment, busy=False, - session=session + session=session, ) @@ -108,10 +120,10 @@ def test_cancel(jfjoch_wrapper, mock_jfjoch_client): def test_is_idle(jfjoch_wrapper, mock_jfjoch_client): - mock_jfjoch_client.status_get.return_value.state = 'Idle' + mock_jfjoch_client.status_get.return_value.state = "Idle" assert jfjoch_wrapper.is_idle() is True - mock_jfjoch_client.status_get.return_value.state = 'Busy' + mock_jfjoch_client.status_get.return_value.state = "Busy" assert jfjoch_wrapper.is_idle() is False @@ -125,7 +137,15 @@ def test_measure_rotation(jfjoch_wrapper, mock_jfjoch_client): dtz=200.0, ) s = create_mock_daq_status() - s.sample = SampleShortInfo(sample_name="S1", user="p12345", db_id=1, dewar_name="D1", puck_name="P1", run_number=1, pin=1) + s.sample = SampleShortInfo( + sample_name="S1", + user="p12345", + db_id=1, + dewar_name="D1", + puck_name="P1", + run_number=1, + pin=1, + ) jfjoch_wrapper.measure_rotation(r, s) mock_jfjoch_client.start_post.assert_called_once() @@ -135,7 +155,15 @@ def test_measure_rotation_error(jfjoch_wrapper, mock_jfjoch_client): mock_jfjoch_client.start_post.side_effect = Exception("API Error") r = RotationScanRequest(exp_time_s=0.01, incr_omega_deg=0.1, steps=10, dtz=200.0) s = create_mock_daq_status() - s.sample = SampleShortInfo(sample_name="S1", user="p12345", db_id=1, dewar_name="D1", puck_name="P1", run_number=1, pin=1) + s.sample = SampleShortInfo( + sample_name="S1", + user="p12345", + db_id=1, + dewar_name="D1", + puck_name="P1", + run_number=1, + pin=1, + ) with pytest.raises(JFJochCommunicationError): jfjoch_wrapper.measure_rotation(r, s) @@ -152,7 +180,15 @@ def test_measure_raster(jfjoch_wrapper, mock_jfjoch_client): dtz=200.0, ) s = create_mock_daq_status() - s.sample = SampleShortInfo(sample_name="S1", user="p12345", db_id=1, dewar_name="D1", puck_name="P1", run_number=1, pin=1) + s.sample = SampleShortInfo( + sample_name="S1", + user="p12345", + db_id=1, + dewar_name="D1", + puck_name="P1", + run_number=1, + pin=1, + ) jfjoch_wrapper.measure_raster(r, s) mock_jfjoch_client.start_post.assert_called_once() @@ -161,7 +197,9 @@ def test_measure_raster(jfjoch_wrapper, mock_jfjoch_client): def test_wait_till_done(jfjoch_wrapper, mock_jfjoch_client): mock_jfjoch_client.result_scan_get.return_value = MagicMock() result = jfjoch_wrapper.wait_till_done(timeout=10) - mock_jfjoch_client.wait_till_done_post_with_http_info.assert_called_once_with(timeout=10) + mock_jfjoch_client.wait_till_done_post_with_http_info.assert_called_once_with( + timeout=10 + ) assert result is not None @@ -181,7 +219,10 @@ def test_get_diffraction_image(jfjoch_wrapper, mock_jfjoch_client): def test_get_diffraction_image_retry(jfjoch_wrapper, mock_jfjoch_client): - mock_jfjoch_client.image_buffer_image_jpeg_get.side_effect = [Exception("Error"), b"image_data"] + mock_jfjoch_client.image_buffer_image_jpeg_get.side_effect = [ + Exception("Error"), + b"image_data", + ] img = jfjoch_wrapper.get_diffraction_image(image_id=1, wait_between_retries_s=0.001) assert img == b"image_data" - assert mock_jfjoch_client.image_buffer_image_jpeg_get.call_count == 2 \ No newline at end of file + assert mock_jfjoch_client.image_buffer_image_jpeg_get.call_count == 2 diff --git a/tests/unit/devices/test_tell_client.py b/tests/unit/devices/test_tell_client.py index 29aa002d..aa4a0b0b 100644 --- a/tests/unit/devices/test_tell_client.py +++ b/tests/unit/devices/test_tell_client.py @@ -1,9 +1,12 @@ -import pytest from unittest.mock import MagicMock -from aare.common.exception_handler import TellCommunicationError -from aare.common.models import DewarAddress, SampleDewarAddress -from aare.devices.tell_client import TellClient, TellEventValueEnum + +import pytest +from aarecommon.errors.exception_handler import TellCommunicationError +from aarecommon.models.models import DewarAddress, SampleDewarAddress + from aare.devices.tell_backend import TellBackend +from aare.devices.tell_client import TellClient, TellEventValueEnum + @pytest.fixture def mock_backend(): @@ -11,21 +14,25 @@ def mock_backend(): backend.get_state.return_value = "Ready" return backend + @pytest.fixture def mock_beamline(): return MagicMock() + def test_get_state(mock_beamline, mock_backend): client = TellClient(mock_beamline, backend=mock_backend) assert client.get_state() == "Ready" mock_backend.get_state.assert_called() + def test_is_in_mount_position_true(mock_beamline, mock_backend): mock_backend.eval.return_value = "True" client = TellClient(mock_beamline, backend=mock_backend) assert client.is_in_mount_position() is True mock_backend.eval.assert_called_with("in_mount_position&") + def test_is_in_mount_position_false(mock_beamline, mock_backend): mock_backend.eval.return_value = "False" client = TellClient(mock_beamline, backend=mock_backend) @@ -98,4 +105,4 @@ def test_mount_state_event_raises_when_command_not_completed( client = TellClient(mock_beamline, backend=mock_backend) with pytest.raises(TellCommunicationError): - client.mount(_mount_address(), wait=True) \ No newline at end of file + client.mount(_mount_address(), wait=True) diff --git a/tests/unit/gui/test_automation_progress_parser.py b/tests/unit/gui/test_automation_progress_parser.py index 20fa1084..49f973ab 100644 --- a/tests/unit/gui/test_automation_progress_parser.py +++ b/tests/unit/gui/test_automation_progress_parser.py @@ -11,9 +11,12 @@ jfjoch_client_scan_result_module.ScanResult = object sys.modules.setdefault("jfjoch_client", jfjoch_client_module) sys.modules.setdefault("jfjoch_client.models", jfjoch_client_models_module) -sys.modules.setdefault("jfjoch_client.models.scan_result", jfjoch_client_scan_result_module) +sys.modules.setdefault( + "jfjoch_client.models.scan_result", jfjoch_client_scan_result_module +) + +from aarecommon.models.automation import StepStatus, WorkflowStateKind -from aare.common.automation_models import StepStatus, WorkflowStateKind from aare.gui.threads.daq_worker import DAQWorker @@ -138,7 +141,9 @@ def test_handle_automation_progress_event_dedups_events_by_timestamp(caplog): worker._handle_automation_progress_event(payload) worker._handle_automation_progress_event(payload) - warning_messages = [record.message for record in caplog.records if record.levelname == "WARNING"] + warning_messages = [ + record.message for record in caplog.records if record.levelname == "WARNING" + ] assert warning_messages.count("Loop centering failed") == 1 worker.cleanup() @@ -146,7 +151,9 @@ def test_handle_automation_progress_event_dedups_events_by_timestamp(caplog): def test_handle_automation_progress_event_trips_recurrence_watcher(): worker = DAQWorker(base_url=None, token="test-token") - worker._recurrence_watchers = [w for w in worker._recurrence_watchers if w.name == "alc"] + worker._recurrence_watchers = [ + w for w in worker._recurrence_watchers if w.name == "alc" + ] trips: list[str] = [] worker.automation_critical_failure.connect(trips.append) @@ -159,8 +166,12 @@ def test_handle_automation_progress_event_trips_recurrence_watcher(): ) worker._handle_automation_progress_event(payload) - worker._handle_automation_progress_event(payload.replace('"seq":1', '"seq":2').replace('10:00:00', '10:00:01')) - worker._handle_automation_progress_event(payload.replace('"seq":1', '"seq":3').replace('10:00:00', '10:00:02')) + worker._handle_automation_progress_event( + payload.replace('"seq":1', '"seq":2').replace("10:00:00", "10:00:01") + ) + worker._handle_automation_progress_event( + payload.replace('"seq":1', '"seq":3').replace("10:00:00", "10:00:02") + ) assert len(trips) == 1 assert "consecutive alc errors" in trips[0] @@ -271,4 +282,4 @@ def test_handle_automation_progress_event_ignores_null_progress(): assert emitted == [] - worker.cleanup() \ No newline at end of file + worker.cleanup() diff --git a/tests/unit/gui/test_camera_thread.py b/tests/unit/gui/test_camera_thread.py index 4c0dcc03..1e3c3950 100644 --- a/tests/unit/gui/test_camera_thread.py +++ b/tests/unit/gui/test_camera_thread.py @@ -1,17 +1,19 @@ -import pytest import json -import numpy as np -import cv2 -import zmq from unittest.mock import MagicMock, patch + +import cv2 +import numpy as np +import pytest +import zmq +from aarecommon.models.models import DAQStatusModel from PySide6.QtGui import QPixmap + from aare.gui.threads.camera_thread import SampleCameraThread -from aare.common.models import DAQStatusModel @pytest.fixture def mock_zmq(): - with patch('zmq.Context') as mock_ctx_class: + with patch("zmq.Context") as mock_ctx_class: mock_ctx = mock_ctx_class.return_value mock_socket = mock_ctx.socket.return_value yield mock_socket @@ -51,7 +53,7 @@ def test_enable_focus_measurement(camera_thread): def test_run_success(camera_thread, mock_zmq, qtbot): img = np.zeros((10, 10, 3), dtype=np.uint8) - _, jpeg_bytes = cv2.imencode('.jpg', img) + _, jpeg_bytes = cv2.imencode(".jpg", img) header = json.dumps({"encoding": "jpeg"}).encode("utf-8") def side_effect(): @@ -106,4 +108,4 @@ def test_stop(camera_thread, mock_zmq): assert camera_thread.running is False if hasattr(camera_thread, "_SampleCameraThread__socket"): - mock_zmq.close.assert_called() \ No newline at end of file + mock_zmq.close.assert_called() diff --git a/tests/unit/gui/test_data_collection_settings.py b/tests/unit/gui/test_data_collection_settings.py index 4f3b17b1..e7f2fbae 100644 --- a/tests/unit/gui/test_data_collection_settings.py +++ b/tests/unit/gui/test_data_collection_settings.py @@ -8,10 +8,10 @@ toggle with the dtz<->resolution coupling that must hold in both modes. import types import pytest +from aarecommon.math.coordinate import Coordinate, SmargonCoordinate +from aarecommon.math.diffraction_geometry import DiffractionGeometry +from aarecommon.models.models import SampleGeometryModel -from aare.common.coordinate import Coordinate, SmargonCoordinate -from aare.common.diffraction_geometry import DiffractionGeometry -from aare.common.models import SampleGeometryModel from aare.gui.panels.raster_data_collection import RasterDataCollectionPanel from aare.gui.panels.rotation_data_collection import RotationDataCollectionPanel from aare.gui.scan_logic.raster_grid_manager import RasterGridManager @@ -68,11 +68,11 @@ def test_db_override_reset_forgets_user_value(qapp): def test_db_override_emits_value_changed_on_toggle(qapp): w = DbOverrideLineEdit(0, 1000, default=200.0, decimals=2) - _edit(w, "300.00") # mine = 300, db = 200 + _edit(w, "300.00") # mine = 300, db = 200 seen = [] w.valueChanged.connect(lambda v: seen.append(v)) w.set_source(DbOverrideLineEdit.SOURCE_DB) - assert seen[-1] == 200.0 # toggling pushes the now-active value downstream + assert seen[-1] == 200.0 # toggling pushes the now-active value downstream w.set_source(DbOverrideLineEdit.SOURCE_MINE) assert seen[-1] == 300.0 @@ -105,7 +105,9 @@ def panel(qapp, diffraction): def test_editing_dtz_updates_resolution_and_switches_to_mine(panel, diffraction): _edit(panel.dtz_enter, "250.00") assert panel._source == DbOverrideLineEdit.SOURCE_MINE - assert abs(panel.high_res_enter.value - diffraction.resolution_angstrom(250.0)) < 0.01 + assert ( + abs(panel.high_res_enter.value - diffraction.resolution_angstrom(250.0)) < 0.01 + ) def test_editing_resolution_updates_dtz(panel, diffraction): @@ -127,24 +129,27 @@ def test_toggle_pushes_active_value_downstream(panel): emitted = [] panel.dtz_updated.connect(lambda v: emitted.append(round(v, 2))) - _edit(panel.dtz_enter, "250.00") # mine dtz = 250 + _edit(panel.dtz_enter, "250.00") # mine dtz = 250 mine_dtz = panel.dtz_enter.value panel.set_source(DbOverrideLineEdit.SOURCE_DB) db_dtz = panel.dtz_enter.value - assert emitted[-1] == round(db_dtz, 2) # downstream got the db value + assert emitted[-1] == round(db_dtz, 2) # downstream got the db value panel.set_source(DbOverrideLineEdit.SOURCE_MINE) - assert panel.dtz_enter.value == mine_dtz # user value recovered - assert emitted[-1] == round(mine_dtz, 2) # downstream got the user value + assert panel.dtz_enter.value == mine_dtz # user value recovered + assert emitted[-1] == round(mine_dtz, 2) # downstream got the user value def test_user_override_persists_across_samples(panel, diffraction): # Sample 1 loads a database resolution. panel._sample = types.SimpleNamespace(db_id=1) panel._params = types.SimpleNamespace( - targetresolution=2.5, transmission=0.5, - totalrange=180.0, oscillation=0.1, exposure=0.02, + targetresolution=2.5, + transmission=0.5, + totalrange=180.0, + oscillation=0.1, + exposure=0.02, ) panel.update_data_collection_parameters() @@ -155,8 +160,11 @@ def test_user_override_persists_across_samples(panel, diffraction): # Sample 2 arrives with a different database exposure. panel._sample = types.SimpleNamespace(db_id=2) panel._params = types.SimpleNamespace( - targetresolution=1.8, transmission=1.0, - totalrange=360.0, oscillation=0.2, exposure=0.01, + targetresolution=1.8, + transmission=1.0, + totalrange=360.0, + oscillation=0.2, + exposure=0.01, ) panel.update_data_collection_parameters() diff --git a/tests/unit/gui/test_gui_main.py b/tests/unit/gui/test_gui_main.py index 23459944..607bd619 100644 --- a/tests/unit/gui/test_gui_main.py +++ b/tests/unit/gui/test_gui_main.py @@ -1,7 +1,7 @@ # import sys # from unittest.mock import MagicMock, patch # -# from aare.common.beamline import MXBeamline +# from aarecommon.models.beamline import MXBeamline # from aare.gui.gui import main # # @@ -46,4 +46,4 @@ # mock_splash.show.assert_called_once() # mock_splash.finish.assert_called_once_with(mock_window) # -# mock_sys_exit.assert_called_once_with(0) \ No newline at end of file +# mock_sys_exit.assert_called_once_with(0) diff --git a/tests/unit/gui/test_main_window.py b/tests/unit/gui/test_main_window.py index 0f6e026a..231ddb2d 100644 --- a/tests/unit/gui/test_main_window.py +++ b/tests/unit/gui/test_main_window.py @@ -1,6 +1,8 @@ -import pytest from unittest.mock import MagicMock, patch + +import pytest from PySide6.QtCore import Qt + from aare.gui.main_window import MainWindow @@ -11,16 +13,22 @@ def mock_ui_state(): def test_main_window_init(qtbot, mock_ui_state): - with patch("requests.get") as mock_get, \ - patch("aare.gui.main_window.DAQWorker"), \ - patch("aare.gui.main_window.PredictionSubscriber"), \ - patch("aare.gui.main_window.VideoThread"), \ - patch("aare.gui.main_window.JFJochDBusClient"), \ - patch("aare.gui.main_window.jwt.decode") as mock_jwt: - + with ( + patch("requests.get") as mock_get, + patch("aare.gui.main_window.DAQWorker"), + patch("aare.gui.main_window.PredictionSubscriber"), + patch("aare.gui.main_window.VideoThread"), + patch("aare.gui.main_window.JFJochDBusClient"), + patch("aare.gui.main_window.jwt.decode") as mock_jwt, + ): mock_get.return_value.status_code = 200 mock_get.return_value.json.return_value = {"status": "ok"} - mock_jwt.return_value = {"sub": "testuser", "staff": True, "pgroups": ["p123"], "session": 15} + mock_jwt.return_value = { + "sub": "testuser", + "staff": True, + "pgroups": ["p123"], + "session": 15, + } fake_token = "header.payload.signature" @@ -32,7 +40,7 @@ def test_main_window_init(qtbot, mock_ui_state): pred_zmq_addr=None, beamline_cam_addr="localhost", gonio_cam_addr="localhost", - gonio_cam_id=1 + gonio_cam_id=1, ) win.show() qtbot.addWidget(win) @@ -40,15 +48,22 @@ def test_main_window_init(qtbot, mock_ui_state): assert win.windowTitle() == "AareGUI" assert win.isVisible() -def test_main_window_mount_view(qtbot, mock_ui_state): - with patch("requests.get"), \ - patch("aare.gui.main_window.DAQWorker"), \ - patch("aare.gui.main_window.PredictionSubscriber"), \ - patch("aare.gui.main_window.VideoThread"), \ - patch("aare.gui.main_window.JFJochDBusClient"), \ - patch("aare.gui.main_window.jwt.decode") as mock_jwt: - mock_jwt.return_value = {"sub": "testuser", "staff": True, "pgroups": ["p123"], "session": 15} +def test_main_window_mount_view(qtbot, mock_ui_state): + with ( + patch("requests.get"), + patch("aare.gui.main_window.DAQWorker"), + patch("aare.gui.main_window.PredictionSubscriber"), + patch("aare.gui.main_window.VideoThread"), + patch("aare.gui.main_window.JFJochDBusClient"), + patch("aare.gui.main_window.jwt.decode") as mock_jwt, + ): + mock_jwt.return_value = { + "sub": "testuser", + "staff": True, + "pgroups": ["p123"], + "session": 15, + } fake_token = "header.payload.signature" win = MainWindow( @@ -59,21 +74,28 @@ def test_main_window_mount_view(qtbot, mock_ui_state): pred_zmq_addr=None, beamline_cam_addr=None, gonio_cam_addr=None, - gonio_cam_id=None + gonio_cam_id=None, ) qtbot.addWidget(win) win.mount_view() -def test_mark_user_interaction_reports_backend(qtbot, mock_ui_state): - with patch("requests.get"), \ - patch("aare.gui.main_window.DAQWorker") as mock_daq_cls, \ - patch("aare.gui.main_window.PredictionSubscriber"), \ - patch("aare.gui.main_window.VideoThread"), \ - patch("aare.gui.main_window.JFJochDBusClient"), \ - patch("aare.gui.main_window.jwt.decode") as mock_jwt: - mock_jwt.return_value = {"sub": "testuser", "staff": True, "pgroups": ["p123"], "session": 15} +def test_mark_user_interaction_reports_backend(qtbot, mock_ui_state): + with ( + patch("requests.get"), + patch("aare.gui.main_window.DAQWorker") as mock_daq_cls, + patch("aare.gui.main_window.PredictionSubscriber"), + patch("aare.gui.main_window.VideoThread"), + patch("aare.gui.main_window.JFJochDBusClient"), + patch("aare.gui.main_window.jwt.decode") as mock_jwt, + ): + mock_jwt.return_value = { + "sub": "testuser", + "staff": True, + "pgroups": ["p123"], + "session": 15, + } fake_token = "header.payload.signature" win = MainWindow( @@ -84,7 +106,7 @@ def test_mark_user_interaction_reports_backend(qtbot, mock_ui_state): pred_zmq_addr=None, beamline_cam_addr=None, gonio_cam_addr=None, - gonio_cam_id=None + gonio_cam_id=None, ) qtbot.addWidget(win) @@ -93,28 +115,35 @@ def test_mark_user_interaction_reports_backend(qtbot, mock_ui_state): win.daq.report_gui_interaction.assert_called_once_with(15) + def test_automation_activity_refreshes_idle_timestamp(qtbot, mock_ui_state): - with patch("requests.get"), \ - patch("aare.gui.main_window.DAQWorker"), \ - patch("aare.gui.main_window.PredictionSubscriber"), \ - patch("aare.gui.main_window.VideoThread"), \ - patch("aare.gui.main_window.JFJochDBusClient"), \ - patch("aare.gui.main_window.jwt.decode") as mock_jwt: - - from aare.common.models import ( - DAQStatusModel, - SessionStatus, - BeamlineStatus, - SampleCameraSettings, - CrystalSize, - SessionsStateEnum, + with ( + patch("requests.get"), + patch("aare.gui.main_window.DAQWorker"), + patch("aare.gui.main_window.PredictionSubscriber"), + patch("aare.gui.main_window.VideoThread"), + patch("aare.gui.main_window.JFJochDBusClient"), + patch("aare.gui.main_window.jwt.decode") as mock_jwt, + ): + from aarecommon.math.coordinate import Coordinate, SmargonCoordinate + from aarecommon.math.diffraction_geometry import DiffractionGeometry + from aarecommon.math.sample_geometry import SampleGeometryModel + from aarecommon.models.models import ( BeamlineStateEnum, + BeamlineStatus, + CrystalSize, + DAQStatusModel, + SampleCameraSettings, + SessionsStateEnum, + SessionStatus, ) - from aare.common.coordinate import Coordinate, SmargonCoordinate - from aare.common.diffraction_geometry import DiffractionGeometry - from aare.common.sample_geometry import SampleGeometryModel - mock_jwt.return_value = {"sub": "testuser", "staff": True, "pgroups": ["p123"], "session": 15} + mock_jwt.return_value = { + "sub": "testuser", + "staff": True, + "pgroups": ["p123"], + "session": 15, + } fake_token = "header.payload.signature" win = MainWindow( @@ -125,7 +154,7 @@ def test_automation_activity_refreshes_idle_timestamp(qtbot, mock_ui_state): pred_zmq_addr=None, beamline_cam_addr=None, gonio_cam_addr=None, - gonio_cam_id=None + gonio_cam_id=None, ) qtbot.addWidget(win) @@ -171,7 +200,9 @@ def test_automation_activity_refreshes_idle_timestamp(qtbot, mock_ui_state): ), state=BeamlineStateEnum.SampleAlignment, busy=False, - session=SessionStatus(session=SessionsStateEnum.Vacant, current_pgroup="p123", staff=True), + session=SessionStatus( + session=SessionsStateEnum.Vacant, current_pgroup="p123", staff=True + ), crystal_size=CrystalSize(x=0, y=0, z=0), ) @@ -180,15 +211,22 @@ def test_automation_activity_refreshes_idle_timestamp(qtbot, mock_ui_state): assert win._last_user_interaction_ts == 250.0 -def test_idle_timeout_closes_when_inactive_and_not_running(qtbot, mock_ui_state): - with patch("requests.get"), \ - patch("aare.gui.main_window.DAQWorker"), \ - patch("aare.gui.main_window.PredictionSubscriber"), \ - patch("aare.gui.main_window.VideoThread"), \ - patch("aare.gui.main_window.JFJochDBusClient"), \ - patch("aare.gui.main_window.jwt.decode") as mock_jwt: - mock_jwt.return_value = {"sub": "testuser", "staff": True, "pgroups": ["p123"], "session": 15} +def test_idle_timeout_closes_when_inactive_and_not_running(qtbot, mock_ui_state): + with ( + patch("requests.get"), + patch("aare.gui.main_window.DAQWorker"), + patch("aare.gui.main_window.PredictionSubscriber"), + patch("aare.gui.main_window.VideoThread"), + patch("aare.gui.main_window.JFJochDBusClient"), + patch("aare.gui.main_window.jwt.decode") as mock_jwt, + ): + mock_jwt.return_value = { + "sub": "testuser", + "staff": True, + "pgroups": ["p123"], + "session": 15, + } fake_token = "header.payload.signature" win = MainWindow( @@ -199,7 +237,7 @@ def test_idle_timeout_closes_when_inactive_and_not_running(qtbot, mock_ui_state) pred_zmq_addr=None, beamline_cam_addr=None, gonio_cam_addr=None, - gonio_cam_id=None + gonio_cam_id=None, ) qtbot.addWidget(win) @@ -212,15 +250,22 @@ def test_idle_timeout_closes_when_inactive_and_not_running(qtbot, mock_ui_state) win.close.assert_called_once() -def test_idle_timeout_does_not_close_while_automation_active(qtbot, mock_ui_state): - with patch("requests.get"), \ - patch("aare.gui.main_window.DAQWorker"), \ - patch("aare.gui.main_window.PredictionSubscriber"), \ - patch("aare.gui.main_window.VideoThread"), \ - patch("aare.gui.main_window.JFJochDBusClient"), \ - patch("aare.gui.main_window.jwt.decode") as mock_jwt: - mock_jwt.return_value = {"sub": "testuser", "staff": True, "pgroups": ["p123"], "session": 15} +def test_idle_timeout_does_not_close_while_automation_active(qtbot, mock_ui_state): + with ( + patch("requests.get"), + patch("aare.gui.main_window.DAQWorker"), + patch("aare.gui.main_window.PredictionSubscriber"), + patch("aare.gui.main_window.VideoThread"), + patch("aare.gui.main_window.JFJochDBusClient"), + patch("aare.gui.main_window.jwt.decode") as mock_jwt, + ): + mock_jwt.return_value = { + "sub": "testuser", + "staff": True, + "pgroups": ["p123"], + "session": 15, + } fake_token = "header.payload.signature" win = MainWindow( @@ -231,7 +276,7 @@ def test_idle_timeout_does_not_close_while_automation_active(qtbot, mock_ui_stat pred_zmq_addr=None, beamline_cam_addr=None, gonio_cam_addr=None, - gonio_cam_id=None + gonio_cam_id=None, ) qtbot.addWidget(win) @@ -245,15 +290,22 @@ def test_idle_timeout_does_not_close_while_automation_active(qtbot, mock_ui_stat win.close.assert_not_called() -def test_cleanup_returns_from_portrait_mode(qtbot, mock_ui_state): - with patch("requests.get"), \ - patch("aare.gui.main_window.DAQWorker"), \ - patch("aare.gui.main_window.PredictionSubscriber"), \ - patch("aare.gui.main_window.VideoThread"), \ - patch("aare.gui.main_window.JFJochDBusClient"), \ - patch("aare.gui.main_window.jwt.decode") as mock_jwt: - mock_jwt.return_value = {"sub": "testuser", "staff": True, "pgroups": ["p123"], "session": 15} +def test_cleanup_returns_from_portrait_mode(qtbot, mock_ui_state): + with ( + patch("requests.get"), + patch("aare.gui.main_window.DAQWorker"), + patch("aare.gui.main_window.PredictionSubscriber"), + patch("aare.gui.main_window.VideoThread"), + patch("aare.gui.main_window.JFJochDBusClient"), + patch("aare.gui.main_window.jwt.decode") as mock_jwt, + ): + mock_jwt.return_value = { + "sub": "testuser", + "staff": True, + "pgroups": ["p123"], + "session": 15, + } fake_token = "header.payload.signature" win = MainWindow( @@ -264,7 +316,7 @@ def test_cleanup_returns_from_portrait_mode(qtbot, mock_ui_state): pred_zmq_addr=None, beamline_cam_addr=None, gonio_cam_addr=None, - gonio_cam_id=None + gonio_cam_id=None, ) qtbot.addWidget(win) @@ -277,14 +329,20 @@ def test_cleanup_returns_from_portrait_mode(qtbot, mock_ui_state): def test_cleanup_returns_from_compact_automation_view(qtbot, mock_ui_state): - with patch("requests.get"), \ - patch("aare.gui.main_window.DAQWorker"), \ - patch("aare.gui.main_window.PredictionSubscriber"), \ - patch("aare.gui.main_window.VideoThread"), \ - patch("aare.gui.main_window.JFJochDBusClient"), \ - patch("aare.gui.main_window.jwt.decode") as mock_jwt: - - mock_jwt.return_value = {"sub": "testuser", "staff": True, "pgroups": ["p123"], "session": 15} + with ( + patch("requests.get"), + patch("aare.gui.main_window.DAQWorker"), + patch("aare.gui.main_window.PredictionSubscriber"), + patch("aare.gui.main_window.VideoThread"), + patch("aare.gui.main_window.JFJochDBusClient"), + patch("aare.gui.main_window.jwt.decode") as mock_jwt, + ): + mock_jwt.return_value = { + "sub": "testuser", + "staff": True, + "pgroups": ["p123"], + "session": 15, + } fake_token = "header.payload.signature" win = MainWindow( @@ -295,7 +353,7 @@ def test_cleanup_returns_from_compact_automation_view(qtbot, mock_ui_state): pred_zmq_addr=None, beamline_cam_addr=None, gonio_cam_addr=None, - gonio_cam_id=None + gonio_cam_id=None, ) qtbot.addWidget(win) @@ -304,4 +362,4 @@ def test_cleanup_returns_from_compact_automation_view(qtbot, mock_ui_state): win.cleanup() - assert win.content_stack.currentWidget() is win._standard_main_page \ No newline at end of file + assert win.content_stack.currentWidget() is win._standard_main_page diff --git a/tests/unit/gui/test_models.py b/tests/unit/gui/test_models.py index d13322c3..d5b3624b 100644 --- a/tests/unit/gui/test_models.py +++ b/tests/unit/gui/test_models.py @@ -1,20 +1,50 @@ import pytest +from aarecommon.models.models import DewarAddress, SampleShortInfo from PySide6.QtCore import Qt + from aare.gui.models.sample_queue_model import SampleQueueSpreadsheet from aare.gui.models.user_sample_model import UserSampleSpreadsheet -from aare.common.models import SampleShortInfo, DewarAddress + @pytest.fixture def sample_list(): return [ - SampleShortInfo(db_id=1, puck_name="P1", dewar_name="D1", sample_name="S1", run_number=1, user="U1", pin=1, location=DewarAddress(segment="A", pos=1)), - SampleShortInfo(db_id=2, puck_name="P2", dewar_name="D2", sample_name="S2", run_number=2, user="U2", pin=2, location=DewarAddress(segment="A", pos=2)), - SampleShortInfo(db_id=3, puck_name="P1", dewar_name="D1", sample_name="S3", run_number=3, user="U1", pin=3, location=DewarAddress(segment="A", pos=1)), + SampleShortInfo( + db_id=1, + puck_name="P1", + dewar_name="D1", + sample_name="S1", + run_number=1, + user="U1", + pin=1, + location=DewarAddress(segment="A", pos=1), + ), + SampleShortInfo( + db_id=2, + puck_name="P2", + dewar_name="D2", + sample_name="S2", + run_number=2, + user="U2", + pin=2, + location=DewarAddress(segment="A", pos=2), + ), + SampleShortInfo( + db_id=3, + puck_name="P1", + dewar_name="D1", + sample_name="S3", + run_number=3, + user="U1", + pin=3, + location=DewarAddress(segment="A", pos=1), + ), ] + def test_user_sample_model_init(sample_list): model = UserSampleSpreadsheet(samples=sample_list) - model.set_show_all_pgroups(True) # Ensure all pgroups are shown for testing + model.set_show_all_pgroups(True) # Ensure all pgroups are shown for testing assert model.rowCount() == 3 # Check if filtering works # "User" is column 5 @@ -23,6 +53,7 @@ def test_user_sample_model_init(sample_list): model.clear_filter() assert model.rowCount() == 3 + def test_user_sample_model_column_filter(sample_list): model = UserSampleSpreadsheet(samples=sample_list) model.set_show_all_pgroups(True) @@ -32,6 +63,7 @@ def test_user_sample_model_column_filter(sample_list): model.clear_all_column_filters() assert model.rowCount() == 3 + def test_user_sample_model_unique_values(sample_list): model = UserSampleSpreadsheet(samples=sample_list) model.set_show_all_pgroups(True) @@ -41,6 +73,7 @@ def test_user_sample_model_unique_values(sample_list): assert "U2" in users assert len(users) == 2 + def test_sample_queue_model_init(sample_list): model = SampleQueueSpreadsheet(samples=sample_list[:2]) assert model.rowCount() == 2 @@ -48,23 +81,27 @@ def test_sample_queue_model_init(sample_list): assert model.data(model.index(0, 0), Qt.ItemDataRole.DisplayRole) == "D1" assert model.data(model.index(1, 2), Qt.ItemDataRole.DisplayRole) == "S2" + def test_sample_queue_model_update(sample_list): model = SampleQueueSpreadsheet() assert model.rowCount() == 0 model.updateData(sample_list[:2]) assert model.rowCount() == 2 + def test_sample_queue_model_remove(sample_list): model = SampleQueueSpreadsheet(samples=sample_list[:2]) model.remove_sample(1) assert model.rowCount() == 1 assert model.samples[0].db_id == 2 + def test_sample_queue_model_clear(sample_list): model = SampleQueueSpreadsheet(samples=sample_list[:2]) model.clearSamples() assert model.rowCount() == 0 + def test_sample_queue_model_set_running(sample_list): model = SampleQueueSpreadsheet(samples=sample_list[:2]) # Background color role for first row @@ -73,12 +110,19 @@ def test_sample_queue_model_set_running(sample_list): color_running = model.data(model.index(0, 0), Qt.ItemDataRole.BackgroundRole) assert color_not_running != color_running + def test_sample_queue_model_header(sample_list): model = SampleQueueSpreadsheet(samples=sample_list[:2]) - assert model.headerData(0, Qt.Orientation.Horizontal, Qt.ItemDataRole.DisplayRole) == "Dewar" - assert model.headerData(1, Qt.Orientation.Vertical, Qt.ItemDataRole.DisplayRole) == "2" + assert ( + model.headerData(0, Qt.Orientation.Horizontal, Qt.ItemDataRole.DisplayRole) + == "Dewar" + ) + assert ( + model.headerData(1, Qt.Orientation.Vertical, Qt.ItemDataRole.DisplayRole) == "2" + ) + def test_sample_queue_model_flags(sample_list): model = SampleQueueSpreadsheet(samples=sample_list[:2]) flags = model.flags(model.index(0, 0)) - assert flags & Qt.ItemFlag.ItemIsDropEnabled \ No newline at end of file + assert flags & Qt.ItemFlag.ItemIsDropEnabled diff --git a/tests/unit/gui/test_panels.py b/tests/unit/gui/test_panels.py index 41bf5495..1b208b93 100644 --- a/tests/unit/gui/test_panels.py +++ b/tests/unit/gui/test_panels.py @@ -1,20 +1,28 @@ import pytest +from aarecommon.math.coordinate import Coordinate, SmargonCoordinate +from aarecommon.math.diffraction_geometry import DiffractionGeometry +from aarecommon.math.sample_geometry import SampleGeometryModel +from aarecommon.models.beamline import MXBeamline +from aarecommon.models.models import ( + BeamlineStateEnum, + BeamlineStatus, + DAQStatusModel, + SampleCameraSettings, + SessionStatus, +) from PySide6.QtCore import Qt + from aare.gui.panels.status_panel import StatusPanel -from aare.common.models import DAQStatusModel, BeamlineStatus, SessionStatus, SampleCameraSettings, BeamlineStateEnum -from aare.common.diffraction_geometry import DiffractionGeometry -from aare.common.sample_geometry import SampleGeometryModel -from aare.common.coordinate import Coordinate, SmargonCoordinate -from aare.common.beamline import MXBeamline + @pytest.fixture def mock_daq_status(): status = DAQStatusModel( bl=BeamlineStatus( - name="X10SA", - ring_current_mA=400.0, - flux_ph_s=1e12, - transmission=1.0, + name="X10SA", + ring_current_mA=400.0, + flux_ph_s=1e12, + transmission=1.0, cryojet_K=100.0, front_light=50.0, back_light=50.0, @@ -24,57 +32,61 @@ def mock_daq_status(): zoom=1.0, commissioning_mode=False, dtz_min=120.0, - dtz_max=1600.0 + dtz_max=1600.0, ), diffraction=DiffractionGeometry( - detector_description="EIGER", - detector_serial_number="123", - dtz_mm=200.0, + detector_description="EIGER", + detector_serial_number="123", + dtz_mm=200.0, pixel_size_mm=0.075, energy_keV=12.658, beam_center_pxl=(1000.0, 1000.0), detector_size_pxl=(4000, 4000), poni_rot1_rad=0.0, - poni_rot2_rad=0.0 + poni_rot2_rad=0.0, ), geom=SampleGeometryModel( beam_location_pxl=Coordinate(x=1000, y=1000), pixel_in_mm=0.001, aerotech=Coordinate(x=0, y=0), aerotech_meas=Coordinate(x=0, y=0), - smargon=SmargonCoordinate(sh_mm=Coordinate(x=0,y=0,z=0), phi_deg=0, chi_deg=0), + smargon=SmargonCoordinate( + sh_mm=Coordinate(x=0, y=0, z=0), phi_deg=0, chi_deg=0 + ), omega_deg=10.5, - beam_size_mm=Coordinate(x=0.02, y=0.01) + beam_size_mm=Coordinate(x=0.02, y=0.01), ), session=SessionStatus(), state=BeamlineStateEnum.Maintenance, - busy=False + busy=False, ) return status + def test_status_panel_update(qtbot, mock_daq_status): panel = StatusPanel() qtbot.addWidget(panel) - + # Initial state check (some values from __init__ defaults) assert "400" in panel.ring_current.text() - + # Update with mock status panel.update_daq_status(mock_daq_status) - + assert panel.omega.text() == "10.50" - assert "20.0" in panel.beam_size.text() # 0.02 * 1000 - assert "10.0" in panel.beam_size.text() # 0.01 * 1000 + assert "20.0" in panel.beam_size.text() # 0.02 * 1000 + assert "10.0" in panel.beam_size.text() # 0.01 * 1000 assert "400.0" in panel.ring_current.text() - assert "1000.00" in panel.flux.text() # 1e12 / 1e9 = 1000 + assert "1000.00" in panel.flux.text() # 1e12 / 1e9 = 1000 assert "0.979" in panel.wavelength.text() + def test_status_panel_low_current(qtbot, mock_daq_status): panel = StatusPanel() qtbot.addWidget(panel) - + mock_daq_status.bl.ring_current_mA = 300.0 panel.update_daq_status(mock_daq_status) - + assert "color: red" in panel.ring_current.text() assert "300.0" in panel.ring_current.text() diff --git a/tests/unit/gui/test_threads_logic.py b/tests/unit/gui/test_threads_logic.py index 1eb80707..18a3f4be 100644 --- a/tests/unit/gui/test_threads_logic.py +++ b/tests/unit/gui/test_threads_logic.py @@ -1,47 +1,78 @@ -import pytest from unittest.mock import MagicMock + +import pytest +from aarecommon.models.models import DAQStatusModel, DewarAddress, SampleShortInfo + from aare.gui.scan_logic.sample_mount_logic import SampleMountLogic -from aare.common.models import DAQStatusModel, SampleShortInfo, DewarAddress + def test_sample_mount_logic_emits_on_change(): logic = SampleMountLogic() mock_slot = MagicMock() logic.sample_changed.connect(mock_slot) - - sample1 = SampleShortInfo(db_id=1, puck_name="P1", dewar_name="D1", sample_name="S1", run_number=1, user="U1", pin=1, location=DewarAddress(segment="A", pos=1)) - sample2 = SampleShortInfo(db_id=2, puck_name="P1", dewar_name="D1", sample_name="S2", run_number=1, user="U1", pin=2, location=DewarAddress(segment="A", pos=1)) - + + sample1 = SampleShortInfo( + db_id=1, + puck_name="P1", + dewar_name="D1", + sample_name="S1", + run_number=1, + user="U1", + pin=1, + location=DewarAddress(segment="A", pos=1), + ) + sample2 = SampleShortInfo( + db_id=2, + puck_name="P1", + dewar_name="D1", + sample_name="S2", + run_number=1, + user="U1", + pin=2, + location=DewarAddress(segment="A", pos=1), + ) + status = MagicMock(spec=DAQStatusModel) status.sample = sample1 - + # First update, should emit logic.update_daq_status(status) mock_slot.assert_called_once_with(sample1) mock_slot.reset_mock() - + # Second update same sample, should NOT emit logic.update_daq_status(status) mock_slot.assert_not_called() - + # Third update different sample, should emit status.sample = sample2 logic.update_daq_status(status) mock_slot.assert_called_once_with(sample2) + def test_sample_mount_logic_none_to_sample(): logic = SampleMountLogic() mock_slot = MagicMock() logic.sample_changed.connect(mock_slot) - - sample1 = SampleShortInfo(db_id=1, puck_name="P1", dewar_name="D1", sample_name="S1", run_number=1, user="U1", pin=1, location=DewarAddress(segment="A", pos=1)) - + + sample1 = SampleShortInfo( + db_id=1, + puck_name="P1", + dewar_name="D1", + sample_name="S1", + run_number=1, + user="U1", + pin=1, + location=DewarAddress(segment="A", pos=1), + ) + status = MagicMock(spec=DAQStatusModel) status.sample = None - + logic.update_daq_status(status) # Initial __sample is None, so if status.sample is also None, it won't emit mock_slot.assert_not_called() - + status.sample = sample1 logic.update_daq_status(status) mock_slot.assert_called_once_with(sample1)