From fcc5aea1a0a8d7c417ad09280f288d16daddd2db Mon Sep 17 00:00:00 2001 From: perl_d Date: Thu, 3 Sep 2026 09:48:56 +0200 Subject: [PATCH 1/3] feat: introduce switching gridscan analysis modes with raster service refactor --- pyproject.toml | 4 +- src/aare/daq/config.py | 10 + src/aare/daq/operations/raster/service.py | 415 ++++++++++------------ src/aare/daq/server.py | 16 + 4 files changed, 209 insertions(+), 236 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 124fa5cb..c5ef8a7f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -7,12 +7,10 @@ requires-python = ">=3.11" dependencies = [ "uv", "gunicorn", - # >=0.7: DataCollectionParameters.transmission is a 0-1 fraction, which - # the scan panels rely on (older releases held an int percentage). "aarecommon>=0.7.3", "pydantic>=2.11", "numpy", - "jfjoch_client>=1.0.0rc165", + "jfjoch_client>=1.0.0rc166", "pyJWT", "pyzmq", "opencv-python-headless", diff --git a/src/aare/daq/config.py b/src/aare/daq/config.py index 1d0bfdda..4ddd4566 100644 --- a/src/aare/daq/config.py +++ b/src/aare/daq/config.py @@ -10,6 +10,7 @@ import numpy as np import redis import redis_lock from aarecommon.config.beamline import cfg_get +from aarecommon.config.config_models import GridscanAnalysisMode from aarecommon.config.logger import setup_logger from aarecommon.errors.exception_handler import BeamlineBusyException from aarecommon.math.coordinate import AerotechCoordinate, Coordinate @@ -656,6 +657,15 @@ class BeamlineConfig: self._client.set(f"{self._bl}:beam_size_x", data.x) self._client.set(f"{self._bl}:beam_size_y", data.y) + @property + def gridscan_analysis_mode(self) -> GridscanAnalysisMode: + data = self._client.get(f"{self._bl}:gridscan_analysis_mode") + return GridscanAnalysisMode(data) if data is not None else GridscanAnalysisMode.FindXtal + + @gridscan_analysis_mode.setter + def gridscan_analysis_mode(self, val: GridscanAnalysisMode) -> GridscanAnalysisMode: + return self._client.set(f"{self._bl}:gridscan_analysis_mode", str(val)) + def _get_settings(self) -> BeamlineSettingsModel: tmp = self._client.get(f"{self._bl}:settings") if tmp is None: diff --git a/src/aare/daq/operations/raster/service.py b/src/aare/daq/operations/raster/service.py index 6715e7f4..9c0b03ad 100644 --- a/src/aare/daq/operations/raster/service.py +++ b/src/aare/daq/operations/raster/service.py @@ -1,8 +1,10 @@ import copy +import re import time from math import ceil, floor from aarecommon.config.beamline import cfg_get +from aarecommon.config.config_models import GridscanAnalysisMode from aarecommon.config.logger_events import ( geom_log_context, log_ml_bundle_meta, @@ -13,22 +15,29 @@ from aarecommon.config.logger_events import ( from aarecommon.errors.exception_handler import RasterScanException from aarecommon.math.coordinate import AerotechCoordinate, Coordinate, SmargonCoordinate from aarecommon.math.find_xtal import ( + center_of_grid, compute_crystal_score_array, get_best_b_factor, get_best_res, get_xtal_size, raster_highest_score, ) +from aarecommon.math.jfjoch_gridscan_union import GridScanResult, Thresholds, analyse from aarecommon.math.raster_grid import grid_to_image_id from aarecommon.models.models import BeamlineStateEnum, MLBoxType from aarecommon.models.raster_grid import ( + CenterOfMassModel, CompletedRasterGrid, CompletedRasterGridElem, RasterGridRequest, ) from aareDB import SampleEventType +from bec_lib.redis_connector.streams import dataclass +from jfjoch_client import ScanResult from jfjoch_client.exceptions import NotFoundException +from jfjoch_client.models.grid_scan import GridScan +from aare.daq.daq import DAQStatusModel from aare.daq.mlbox import MLBoxPredictionResult from aare.daq.operations.common.ml_bounding_box import ( MLRasterPlan, @@ -39,10 +48,92 @@ from aare.daq.operations.raster.models import RasterContext from aare.devices.area_detector import AutoEnum +@dataclass +class _AnalysisResults: + center_smargon_coord: SmargonCoordinate + com: CenterOfMassModel + v2_results: GridScanResult | None = None + + class RasterService: def __init__(self, *, context: RasterContext, logger): self.ctx = context self.logger = logger + self.ingestor = self.ctx.services.ingestion or self.ctx.deps.aare + self.send_sample_event = ( + self.ctx.services.events.send + if self.ctx.services.events is not None + else self.ctx.deps.aare.send_sample_event + ) + + def _get_anaylsis_results( + self, + mode: GridscanAnalysisMode, + scan_result: ScanResult, + request: RasterGridRequest, + *, + daq: DAQStatusModel | None = None, + thresholds: Thresholds | None = None, + z_um: float | None = None, + ): + if mode == GridscanAnalysisMode.FindXtal: + com = raster_highest_score(scan_result.images) + smargon_coord = self._smargon_offset_from_scan_result(scan_result, request, com) + v2_results = None + else: + step_x_um = (request.grid_size_mm.x / request.n_x) * 1000 + step_y_um = (request.grid_size_mm.y / request.n_y) * 1000 + grid_scan = GridScan( + n_fast=request.n_x, step_x_um=step_x_um, step_y_um=step_y_um, snake=True + ) + v2_results = analyse(scan_result, grid_scan, daq=daq, thresholds=thresholds, z_um=z_um) + self.logger.info( + f"experimental gridscan analaysis results: \n {v2_results.model_dump_json(indent=2)}" + ) + if v2_results.centre is not None: + com = CenterOfMassModel( + n_x=v2_results.centre.nx, + n_y=v2_results.centre.ny, + max_image=v2_results.centre.image_number, + ) + smargon_coord = SmargonCoordinate( + sh_mm=request.smargon_top_left.sh_mm + v2_results.centre.offset_mm(), + phi_deg=request.smargon_top_left.phi_deg, + chi_deg=request.smargon_top_left.chi_deg, + ) + else: + com = center_of_grid(request.n_x, request.n_y) + smargon_coord = self._smargon_offset_from_scan_result(scan_result, request, com) + return _AnalysisResults(smargon_coord, com, v2_results) + + def _smargon_offset_from_scan_result( + self, scan_result: ScanResult, request: RasterGridRequest, com: CenterOfMassModel + ): + target_coor = com.get_com_mm(request) + target_coor_offset = self.ctx.sample_geometry.smargon_nudge(target_coor) + self.logger.info( + "Calculated raster centre offset", + extra=merge_log_context( + sample_log_context(self.ctx.sample), + raster_request_log_context(request), + { + "centre_offset_x_mm": target_coor_offset.x, + "centre_offset_y_mm": target_coor_offset.y, + "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), + }, + ), + ) + self.logger.info(f"moving Smargon to grid centre offset {target_coor_offset}") + return SmargonCoordinate( + sh_mm=request.smargon_top_left.sh_mm + target_coor_offset, + phi_deg=request.smargon_top_left.phi_deg, + chi_deg=request.smargon_top_left.chi_deg, + ) @staticmethod def _scan_result_image_count(scan_result) -> int | None: @@ -165,61 +256,6 @@ class RasterService: x1, y1, x2, y2 = box 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 - clipped at the frame edge, otherwise loop_face.""" - if plan.loop_all_box is not None and not self._box_touches_frame_edge( - plan.loop_all_box, plan.image_width, plan.image_height - ): - return plan.loop_all_box - return plan.loop_face_box - - def _apply_zoom_to_fit(self, plan: MLRasterPlan) -> None: - """Zoom so the chosen loop box covers at most 1/3 of the frame.""" - box = self._zoom_to_fit_box(plan) - if box is None or not plan.image_width or not plan.image_height: - return - - x1, y1, x2, y2 = box - box_px = max(x2 - x1, y2 - y1) - if box_px <= 0: - return - screen_px = min(plan.image_width, plan.image_height) - - current_zoom = self.ctx.deps.devs.zoom - pixel_in_mm_now = self.ctx.deps.cfg.pixel_to_mm(current_zoom) - box_mm = box_px * pixel_in_mm_now - # Target mm-per-pixel so the box spans exactly screen/3. - target_pixel_in_mm = 3.0 * box_mm / screen_px - try: - target_zoom = self.ctx.deps.cfg.zoom_for_pixel_to_mm(target_pixel_in_mm) - except (ValueError, ZeroDivisionError) as e: - self.logger.warning(f"Skipping auto-center zoom-to-fit: {e}") - return - - zoom_min = float(cfg_get("daq.hardware.zoom_min", 1.0)) - zoom_max = float(cfg_get("daq.hardware.zoom_max", 1000.0)) - target_zoom = max(zoom_min, min(zoom_max, target_zoom)) - - self.logger.info( - "Auto-center zoom-to-fit", - extra=merge_log_context( - sample_log_context(self.ctx.sample), - { - "box_px": box_px, - "screen_px": screen_px, - "current_zoom": current_zoom, - "target_zoom": target_zoom, - "used_loop_all": box is plan.loop_all_box, - }, - ), - ) - - self.ctx.deps.devs.samcam_auto(AutoEnum.AUTO) - self.ctx.deps.devs.set_zoom(target_zoom, wait=True) - time.sleep(0.2) - self.ctx.deps.devs.samcam_auto(AutoEnum.ONCE) - def auto_center_line_scan_top_left( self, *, @@ -386,17 +422,72 @@ class RasterService: return top_left, n_y - def execute( - self, request: RasterGridRequest, wait_for_screenshot: float | None = None - ) -> CompletedRasterGridElem: + def _do_gridscan_move_on_aerotech(self, request: RasterGridRequest, expected_time: float): + self.ctx.deps.devs.set_dtz(request.dtz, wait=True) + self.ctx.deps.devs.aerotech.grid_scan( + grid_elem_size_y_um=request.grid_size_mm.y * 1000, + grid_elem_size_x_um=request.grid_size_mm.x * 1000, + grid_elem_count_x=request.n_x, + grid_elem_count_y=request.n_y, + time_sec=request.exp_time_s, + run_async=True, + ) + self.ctx.deps.devs.aerotech.wait_till_done( + timeout=int(round(expected_time + expected_time * 0.1 + 60, 0)) + ) + + if isinstance(self.ctx.deps.cfg.abr_meas_pos, Coordinate): + coord = self.ctx.deps.cfg.abr_meas_pos + else: + coord = self.ctx.deps.cfg.abr_meas_pos.at_mm + self.ctx.deps.devs.aerotech_pos = AerotechCoordinate( + at_mm=coord, omega_deg=self.ctx.deps.devs.aerotech_omega + ) + self.ctx.deps.devs.aerotech.wait_till_done(timeout=60) + + def _check_setup(self, request: RasterGridRequest, wait_for_screenshot_s: float | None): total_time = request.exp_time_s * request.n_x * request.n_y + request.n_y * 0.3 - if wait_for_screenshot is None: - wait_for_screenshot = self.ctx.settings.wait_for_screenshot_s + if wait_for_screenshot_s is None: + wait_for_screenshot_s = self.ctx.settings.wait_for_screenshot_s if total_time > 1200: raise RasterScanException(f"Raster scan is too long {total_time}s > 20min") + return total_time, wait_for_screenshot_s + + def _initial_gridscan_ingestion( + self, + request: RasterGridRequest, + scan_result: ScanResult, + wait_for_screenshot_s: float, + sample_id: int, + ): + 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) + + time.sleep(wait_for_screenshot_s) + + self.ctx.services.screenshots.save_to_db( + sample_id, f"{sample_id}_post_raster_{int(request.omega_deg)}deg" + ) + + self.ingestor.ingest_gridscan( + sample=self.ctx.sample, + raster_result=scan_result, + 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), + ) + + def _additional_analysis_data_ingestion(self, results: GridScanResult): ... + def execute( + self, request: RasterGridRequest, wait_for_screenshot: float | None = None + ) -> CompletedRasterGridElem: + total_time, wait_for_screenshot = self._check_setup(request, wait_for_screenshot) status = self.ctx.status smargon_top_left = request.smargon_top_left @@ -421,43 +512,11 @@ class RasterService: if self.ctx.sample is not None and self.ctx.sample.db_id is not None: self.ctx.deps.aare.create_gridscan_run(self.ctx.sample, request, status) - 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.ctx.deps.jfjoch.wait_till_running(timeout=60.0) self.logger.debug( f"Starting grid scan with {request.n_x}x{request.n_y} points, exp time {request.exp_time_s}s" ) - if not self.ctx.deps.cfg.simulated_detector: - # The detector move (sa2dc) is non-blocking; ensure it has reached - # the requested distance before triggering the scan, so data is - # collected at the dtz already reported to JFJoch/DB. - self.ctx.deps.devs.set_dtz(request.dtz, wait=True) - self.ctx.deps.devs.aerotech.grid_scan( - grid_elem_size_y_um=request.grid_size_mm.y * 1000, - grid_elem_size_x_um=request.grid_size_mm.x * 1000, - grid_elem_count_x=request.n_x, - grid_elem_count_y=request.n_y, - time_sec=request.exp_time_s, - run_async=True, - ) - - self.ctx.deps.devs.aerotech.wait_till_done( - timeout=int(round(total_time + total_time * 0.1 + 60, 0)) - ) - - if isinstance(self.ctx.deps.cfg.abr_meas_pos, Coordinate): - coord = self.ctx.deps.cfg.abr_meas_pos - else: - coord = self.ctx.deps.cfg.abr_meas_pos.at_mm - self.ctx.deps.devs.aerotech_pos = AerotechCoordinate( - at_mm=coord, omega_deg=self.ctx.deps.devs.aerotech_omega - ) - self.ctx.deps.devs.aerotech.wait_till_done(timeout=60) - - x = None - y = None + self._do_gridscan_move_on_aerotech(request, total_time) scan_result = self.ctx.deps.jfjoch.wait_till_done(60) if scan_result is None: @@ -474,64 +533,10 @@ class RasterService: self.logger.info( f"Raster scan results from JFJoch: file {scan_result.file_prefix} with images {scan_result.images}" ) - com = raster_highest_score(scan_result.images) - if com is None: - self.logger.info("Calcualted COM is None -> using centre image") - if request.n_x == 1: - x = request.grid_size_mm.x / 2.0 - else: - 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)) - - self.logger.info( - "Calculated raster centre offset", - extra=merge_log_context( - sample_log_context(self.ctx.sample), - raster_request_log_context(request), - { - "centre_offset_x_mm": target_coor_offset.x, - "centre_offset_y_mm": target_coor_offset.y, - "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), - }, - ), - ) - else: - self.logger.info("Calcualted COM is not None, proceeding") - target_coor = com.get_com_mm(request) - target_coor_offset = self.ctx.sample_geometry.smargon_nudge(target_coor) - self.logger.info( - "Calculated raster centre offset", - extra=merge_log_context( - sample_log_context(self.ctx.sample), - raster_request_log_context(request), - { - "centre_offset_x_mm": target_coor_offset.x, - "centre_offset_y_mm": target_coor_offset.y, - "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), - }, - ), - ) - - 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, - chi_deg=request.smargon_top_left.chi_deg, + analysis_results = self._get_anaylsis_results( + self.ctx.deps.cfg.gridscan_analysis_mode, scan_result, request, daq=self.ctx.status ) - - self.ctx.deps.devs.smargon_pos = target_smargon + self.ctx.deps.devs.smargon_pos = analysis_results.center_smargon_coord self.ctx.deps.devs.smargon_wait(timeout=180) self.logger.info( @@ -555,68 +560,30 @@ class RasterService: else None ) if sample_id: - 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) + self._initial_gridscan_ingestion( + request, scan_result, wait_for_screenshot, sample_id + ) + if analysis_results.v2_results is not None: + self._additional_analysis_data_ingestion(analysis_results.v2_results) - if wait_for_screenshot and wait_for_screenshot > 0: - time.sleep(wait_for_screenshot) - - self.ctx.services.screenshots.save_to_db( - sample_id, f"{sample_id}_post_raster_{int(request.omega_deg)}deg" + diffraction_image_id = analysis_results.com.max_image + diffraction_image_filename = ( + f"{sample_id}_best_diffraction_from_raster_image_{diffraction_image_id}" ) - if self.ctx.services.ingestion is not None: - self.ctx.services.ingestion.ingest_gridscan( - sample=self.ctx.sample, - raster_result=scan_result, - 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), - ) - else: - self.ctx.deps.aare.ingest_gridscan( - sample=self.ctx.sample, - raster_result=scan_result, - 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), - ) - - 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_id = com.max_image - else: - 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, request=request - ) - self._upload_raster_diffraction_preview( - sample_id=self.ctx.sample.db_id, + sample_id=sample_id, filename=diffraction_image_filename, image_id=diffraction_image_id, scan_result=scan_result, request=request, ) - if self.ctx.services.events is not None: - self.ctx.services.events.send( - self.ctx.sample.db_id, - SampleEventType.RASTERED, - comment=f"Raster completed at {request.omega_deg:.1f} deg", - ) - else: - self.ctx.deps.aare.send_sample_event( - self.ctx.sample.db_id, - event_type=SampleEventType.RASTERED, - comment=f"Raster completed at {request.omega_deg:.1f} deg", - ) + self.send_sample_event( + sample_id, + SampleEventType.RASTERED, + comment=f"Raster completed at {request.omega_deg:.1f} deg", + ) self.logger.info( "Raster finished", @@ -671,18 +638,9 @@ class RasterService: ), ) - if self.ctx.services.events is not None: - self.ctx.services.events.send( - sample.db_id, - SampleEventType.RASTERING, - comment=f"Raster at {geom.omega_deg:.1f} deg", - ) - else: - self.ctx.deps.aare.send_sample_event( - sample.db_id, - SampleEventType.RASTERING, - comment=f"Raster at {geom.omega_deg:.1f} deg", - ) + self.send_sample_event( + sample.db_id, SampleEventType.RASTERING, comment=f"Raster at {geom.omega_deg:.1f} deg" + ) plan = self.ml_raster_plan(sample.db_id, f"ml_{geom.omega_deg:.2f}deg") @@ -773,18 +731,12 @@ class RasterService: res1 = self.execute(grid) grid.omega_deg += 90 - if self.ctx.services.events is not None: - self.ctx.services.events.send( - self.ctx.sample.db_id, - SampleEventType.RASTERING, - comment=f"Raster at {geom.omega_deg:.1f} deg", - ) - else: - self.ctx.deps.aare.send_sample_event( - self.ctx.sample.db_id, - SampleEventType.RASTERING, - comment=f"Raster at {geom.omega_deg:.1f} deg", - ) + self.send_sample_event( + self.ctx.sample.db_id, + SampleEventType.RASTERING, + comment=f"Raster at {geom.omega_deg:.1f} deg", + ) + self.ctx.deps.devs.aerotech_omega = grid.omega_deg grid.n_x = 1 @@ -813,12 +765,9 @@ class RasterService: ) status = self.ctx.status - if not self.ctx.deps.cfg.simulated_detector: - self.logger.info(f"initialise detector for raster at {grid.omega_deg}") - 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(f"initialise detector for raster at {grid.omega_deg}") + self.ctx.deps.jfjoch.measure_raster(grid, status) + self.logger.info("detector initialised") if self.ctx.services.state is None: raise RuntimeError("RasterService requires services.state") diff --git a/src/aare/daq/server.py b/src/aare/daq/server.py index ce4a1c6d..f56e5e70 100644 --- a/src/aare/daq/server.py +++ b/src/aare/daq/server.py @@ -846,6 +846,22 @@ async def local_contact_restart_device(device: str, token: str = Depends(oauth2_ return result +@app.get("/local_contact/get_gridscan_mode") +async def get_gridscan_mode(token: str = Depends(oauth2_scheme)) -> GridscanAnalysisMode: + """ + Toggle the TELL blower. Staff only. + + Args: + token: OAuth2 access token. + + Returns: + Dictionary with status and message. + """ + data = auth.parse_token(token) + auth.check_jwt_staff_only(data) + return cfg.get_gridscan_analysis_mode() + + @app.post("/local_contact/resync/detector_metadata") async def local_contact_resync_detector_metadata(token: str = Depends(oauth2_scheme)) -> dict: """ -- 2.54.0 From e17db09f1df5e93b89decd56399ee181188a00b8 Mon Sep 17 00:00:00 2001 From: perl_d Date: Thu, 3 Sep 2026 13:37:32 +0200 Subject: [PATCH 2/3] feat: propagate gridscan analysis choice to GUI --- src/aare/daq/server.py | 24 ++++++++++++++++--- src/aare/gui/panels/local_contact_panel.py | 27 ++++++++++++++++++++++ src/aare/gui/threads/daq_worker.py | 21 +++++++++++++++++ 3 files changed, 69 insertions(+), 3 deletions(-) diff --git a/src/aare/daq/server.py b/src/aare/daq/server.py index f56e5e70..58c777ad 100644 --- a/src/aare/daq/server.py +++ b/src/aare/daq/server.py @@ -10,6 +10,7 @@ from typing import Any, ClassVar import uvicorn from aarecommon.config.beamline import mx_beamline +from aarecommon.config.config_models import GridscanAnalysisMode 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 ( @@ -846,10 +847,10 @@ async def local_contact_restart_device(device: str, token: str = Depends(oauth2_ return result -@app.get("/local_contact/get_gridscan_mode") +@app.get("/local_contact/gridscan_analysis_mode") async def get_gridscan_mode(token: str = Depends(oauth2_scheme)) -> GridscanAnalysisMode: """ - Toggle the TELL blower. Staff only. + Get the gridscan analysis mode Args: token: OAuth2 access token. @@ -859,7 +860,24 @@ async def get_gridscan_mode(token: str = Depends(oauth2_scheme)) -> GridscanAnal """ data = auth.parse_token(token) auth.check_jwt_staff_only(data) - return cfg.get_gridscan_analysis_mode() + return cfg.gridscan_analysis_mode + + +@app.put("/local_contact/gridscan_analysis_mode") +async def set_gridscan_mode(val: GridscanAnalysisMode, token: str = Depends(oauth2_scheme)) -> str: + """ + Set the gridscan analysis mode + + Args: + token: OAuth2 access token. + + Returns: + Dictionary with status and message. + """ + data = auth.parse_token(token) + auth.check_jwt_staff_only(data) + cfg.gridscan_analysis_mode = val + return "OK" @app.post("/local_contact/resync/detector_metadata") diff --git a/src/aare/gui/panels/local_contact_panel.py b/src/aare/gui/panels/local_contact_panel.py index 3361d978..3f12e3aa 100644 --- a/src/aare/gui/panels/local_contact_panel.py +++ b/src/aare/gui/panels/local_contact_panel.py @@ -3,6 +3,7 @@ from __future__ import annotations from collections.abc import Callable from typing import ClassVar +from aarecommon.config.config_models import GridscanAnalysisMode from aarecommon.config.logger import setup_logger from aarecommon.models.models import DAQStatusModel from PySide6.QtCore import QUrl, Slot @@ -616,8 +617,34 @@ class LocalContactPanel(QFrame): button_layout.addWidget(reload_button) button_layout.addStretch(1) + analysis_settings = QGroupBox("Analysis settings", tab) + analysis_settings_layout = QGridLayout(analysis_settings) + form_layout.setContentsMargins(10, 12, 10, 10) + form_layout.setHorizontalSpacing(8) + form_layout.setVerticalSpacing(8) + analysis_choice_label = QLabel("Gridscan analysis mode:", analysis_settings) + find_xtal_button = QPushButton('classic "find_xtal.py"', analysis_settings) + jfj_gs_u_button = QPushButton('Meitian\'s "jfjoch_gridscan_union.py"', analysis_settings) + current_label = QLabel("Current:", analysis_settings) + current_value_label = QLabel("Unknown/default", analysis_settings) + analysis_settings_layout.addWidget(analysis_choice_label, 1, 0) + analysis_settings_layout.addWidget(find_xtal_button, 2, 0) + analysis_settings_layout.addWidget(jfj_gs_u_button, 2, 1) + analysis_settings_layout.addWidget(current_label, 2, 2) + analysis_settings_layout.addWidget(current_value_label, 2, 3) + + find_xtal_button.clicked.connect( + lambda: self._daq.set_gridscan_analysis_mode(GridscanAnalysisMode.FindXtal) + ) + jfj_gs_u_button.clicked.connect( + lambda: self._daq.set_gridscan_analysis_mode(GridscanAnalysisMode.JfjochGridscanUnion) + ) + self._daq.gridscan_analysis_mode.connect(current_value_label.setText) + self._daq.refresh_gridscan_analysis_mode() + layout.addWidget(form_box) layout.addWidget(button_row) + layout.addWidget(analysis_settings) layout.addStretch(1) return tab diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py index 81579a2b..7fdd3ffc 100644 --- a/src/aare/gui/threads/daq_worker.py +++ b/src/aare/gui/threads/daq_worker.py @@ -10,6 +10,7 @@ from dataclasses import dataclass from datetime import datetime from typing import ClassVar, Literal, cast +from aarecommon.config.config_models import GridscanAnalysisMode from aarecommon.config.logger import setup_logger from aarecommon.errors.codes import AareErrorCode, AuthErrorCode, export_error_codes from aarecommon.math.coordinate import AerotechCoordinate, SmargonCoordinate @@ -136,6 +137,8 @@ class DAQWorker(QObject): steer_beam_available = Signal(bool) + gridscan_analysis_mode = Signal(str) + def __init__(self, base_url: str | None, token: str, parent=None): """ Initialize the DAQWorker. @@ -1694,6 +1697,24 @@ class DAQWorker(QObject): logger.exception(message) self.local_contact_transfer_error.emit(message) + def refresh_gridscan_analysis_mode(self): + request = QNetworkRequest(QUrl(f"{self._base_url}/local_contact/gridscan_analysis_mode")) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) + reply = self._net_manager.get(request) + reply.finished.connect( + lambda: self.gridscan_analysis_mode.emit(self.handle_response(reply)) + ) + + def set_gridscan_analysis_mode(self, value: GridscanAnalysisMode): + request = QNetworkRequest( + QUrl(f"{self._base_url}/local_contact/gridscan_analysis_mode?val={value!s}") + ) + request.setRawHeader(b"Authorization", f"Bearer {self._token}".encode()) + request.setRawHeader(b"Content-Type", b"application/json") + body = QByteArray(str(value).encode("utf-8")) + reply = self._net_manager.put(request, body) + reply.finished.connect(lambda: self.refresh_gridscan_analysis_mode()) + @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()}") -- 2.54.0 From f6f68f17e52ea3fd7fcefbe4f350c22b641971c2 Mon Sep 17 00:00:00 2001 From: perl_d Date: Thu, 3 Sep 2026 16:38:52 +0200 Subject: [PATCH 3/3] feat: forward v2 scan results decision to DB and move calculation to analysis stage --- src/aare/daq/aaredb.py | 110 +++++++++------------ src/aare/daq/daq.py | 17 +++- src/aare/daq/operations/common/services.py | 29 +++++- src/aare/daq/operations/raster/service.py | 14 ++- 4 files changed, 100 insertions(+), 70 deletions(-) diff --git a/src/aare/daq/aaredb.py b/src/aare/daq/aaredb.py index 70f14953..ea84ea69 100644 --- a/src/aare/daq/aaredb.py +++ b/src/aare/daq/aaredb.py @@ -13,6 +13,7 @@ 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.gridscan_decision import GridScanDecision from aarecommon.models.models import ( DAQStatusModel, PuckLoadedInfo, @@ -35,6 +36,7 @@ from aareDB import ( ) from aareDB import Detector as DetectorParameters from jfjoch_client.models import ScanResult +from numpy.typing import NDArray from pydantic import StrictInt logger = setup_logger("aareDAQ") @@ -303,20 +305,27 @@ class AareWrapper: sample: SampleShortInfo | None, raster_result: ScanResult, raster_request: RasterGridRequest, + analysis_image_scores: list[float | None] | None, geom: SampleGeometryModel, - com: CenterOfMassModel | None, + com: CenterOfMassModel, beam_mark_pxl: tuple[float, float], + decision: GridScanDecision | None, ): if sample is None: + logger.error("No sample for DB gridscan ingestion request!") return - 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() + payload = self.format_gridscan_payload( + sample, + raster_result, + raster_request, + geom, + com, + analysis_image_scores, + beam_mark_pxl, + decision, + ).model_dump_json() url = f"{self._host}/protected_router/gridscan_runner/ingest" headers = { @@ -327,14 +336,13 @@ class AareWrapper: url, auth=(os.getenv("AAREDB_USERNAME"), os.getenv("AAREDB_PASSWORD")), headers=headers, - data=json.dumps(payload), + data=payload, timeout=30, verify=self._ssl_ca_cert, cert=(self._cert_file, self._key_file), ) response.raise_for_status() - - logger.info(f"Response status code: {response.status_code}") + logger.info(f"AareDB gridscan ingestion response status code: {response.status_code}") def format_gridscan_payload( self, @@ -342,64 +350,40 @@ class AareWrapper: raster_result: ScanResult, raster_request: RasterGridRequest, geom: SampleGeometryModel, - com: CenterOfMassModel | None, + com: CenterOfMassModel, + analysis_image_scores: list[float | None] | None, beam_mark_pxl: tuple[float, float], - ) -> RasterPayloadModel | None: + decision: GridScanDecision | None, + ) -> RasterPayloadModel: + 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, + ) - 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, - ) + if raster_request.smargon_top_left is None: + beam_loc = geom.smargon.sh_mm + else: + beam_loc = raster_request.smargon_top_left.sh_mm - if raster_request.smargon_top_left is None: - beam_loc = geom.smargon.sh_mm - else: - beam_loc = raster_request.smargon_top_left.sh_mm + start_pxl = geom.smargon_to_picture(beam_loc) - start_pxl = geom.smargon_to_picture(beam_loc) + com_grid_pxl = com.get_com_pxl(raster_request, geom) + center_pxl = Coordinate(x=start_pxl.x + com_grid_pxl.x, y=start_pxl.y + com_grid_pxl.y) - if com: - 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) - 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 - ] - except Exception as e: - logger.warning( - f"raster score computation failed, sending null score: {e}", exc_info=True - ) - 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, - ) - - return payload - - except Exception as e: - logger.error(e) - raise + return RasterPayloadModel( + request=raster_request, + result=raster_result, + sample_id=sample.db_id, + attach_image=True, + centre_of_mass=com, + raster_score=analysis_image_scores, + 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, + decision=decision, + ) @log_timing(logger, "AareDB call") def ingest_scan( diff --git a/src/aare/daq/daq.py b/src/aare/daq/daq.py index a9adf94e..c9b836ea 100644 --- a/src/aare/daq/daq.py +++ b/src/aare/daq/daq.py @@ -41,6 +41,7 @@ from aarecommon.errors.exception_handler import ( ) from aarecommon.math.coordinate import AerotechCoordinate, Coordinate, SmargonCoordinate from aarecommon.math.diffraction_geometry import DiffractionGeometry +from aarecommon.math.find_xtal import CenterOfMassModel from aarecommon.math.sample_geometry import SampleGeometryModel from aarecommon.models.automation import ( AutomationProgress, @@ -64,10 +65,11 @@ from aarecommon.models.models import ( SimpleScanParameters, ZoomModeEnum, ) -from aarecommon.models.raster_grid import CompletedRasterGrid, RasterGridRequest +from aarecommon.models.raster_grid import CompletedRasterGrid, GridScanDecision, 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 from aare.beamline_dispatch.protocols import BeamlineDispatch from aare.daq import workflows @@ -214,15 +216,26 @@ class _DAQScanIngestor: ) def ingest_gridscan( - self, *, sample, raster_result, raster_request, geom, com, beam_mark_pxl + self, + *, + sample: SampleShortInfo | None, + raster_result: ScanResult, + raster_request: RasterGridRequest, + analysis_image_scores: list[float | None] | None, + geom: SampleGeometryModel, + com: CenterOfMassModel, + beam_mark_pxl: tuple[float, float], + decision: GridScanDecision | None, ) -> None: self._daq._aare.ingest_gridscan( sample=sample, raster_result=raster_result, raster_request=raster_request, + analysis_image_scores=analysis_image_scores, geom=geom, com=com, beam_mark_pxl=beam_mark_pxl, + decision=decision, ) diff --git a/src/aare/daq/operations/common/services.py b/src/aare/daq/operations/common/services.py index 0b9252ef..20baf241 100644 --- a/src/aare/daq/operations/common/services.py +++ b/src/aare/daq/operations/common/services.py @@ -1,6 +1,11 @@ from dataclasses import dataclass from typing import Protocol +from aarecommon.math.find_xtal import CenterOfMassModel +from aarecommon.models.models import SampleGeometryModel, SampleShortInfo +from aarecommon.models.raster_grid import GridScanDecision, RasterGridRequest +from jfjoch_client import ScanResult + from aare.daq.config import BeamlineStateEnum from aare.daq.operations.screenshot.service import ScreenshotService @@ -29,7 +34,16 @@ class ScanIngestor(Protocol): def ingest_scan(self, *, sample, result, geom, beam_mark_pxl) -> None: ... def ingest_gridscan( - self, *, sample, raster_result, raster_request, geom, com, beam_mark_pxl + self, + *, + sample: SampleShortInfo | None, + raster_result: ScanResult, + raster_request: RasterGridRequest, + analysis_image_scores: list[float | None] | None, + geom: SampleGeometryModel, + com: CenterOfMassModel, + beam_mark_pxl: tuple[float, float], + decision: GridScanDecision | None, ) -> None: ... @@ -87,15 +101,26 @@ class ScanIngestionService: ) def ingest_gridscan( - self, *, sample, raster_result, raster_request, geom, com, beam_mark_pxl + self, + *, + sample: SampleShortInfo | None, + raster_result: ScanResult, + raster_request: RasterGridRequest, + analysis_image_scores: list[float | None] | None, + geom: SampleGeometryModel, + com: CenterOfMassModel, + beam_mark_pxl: tuple[float, float], + decision: GridScanDecision | None, ) -> None: self.ingestor.ingest_gridscan( sample=sample, raster_result=raster_result, raster_request=raster_request, + analysis_image_scores=analysis_image_scores, geom=geom, com=com, beam_mark_pxl=beam_mark_pxl, + decision=decision, ) diff --git a/src/aare/daq/operations/raster/service.py b/src/aare/daq/operations/raster/service.py index 9c0b03ad..d157aaae 100644 --- a/src/aare/daq/operations/raster/service.py +++ b/src/aare/daq/operations/raster/service.py @@ -45,13 +45,13 @@ from aare.daq.operations.common.ml_bounding_box import ( get_ml_bounding_box, ) from aare.daq.operations.raster.models import RasterContext -from aare.devices.area_detector import AutoEnum @dataclass class _AnalysisResults: center_smargon_coord: SmargonCoordinate com: CenterOfMassModel + scores: list[float | None] | None v2_results: GridScanResult | None = None @@ -59,7 +59,7 @@ class RasterService: def __init__(self, *, context: RasterContext, logger): self.ctx = context self.logger = logger - self.ingestor = self.ctx.services.ingestion or self.ctx.deps.aare + self.ingestor = self.ctx.deps.aare self.send_sample_event = ( self.ctx.services.events.send if self.ctx.services.events is not None @@ -78,6 +78,13 @@ class RasterService: ): if mode == GridscanAnalysisMode.FindXtal: com = raster_highest_score(scan_result.images) + _scores = compute_crystal_score_array(scan_result.images) + scores = [ + float(_scores[img.nx, img.ny]) + if img.nx is not None and img.ny is not None + else None + for img in scan_result.images + ] smargon_coord = self._smargon_offset_from_scan_result(scan_result, request, com) v2_results = None else: @@ -104,7 +111,8 @@ class RasterService: else: com = center_of_grid(request.n_x, request.n_y) smargon_coord = self._smargon_offset_from_scan_result(scan_result, request, com) - return _AnalysisResults(smargon_coord, com, v2_results) + scores = None + return _AnalysisResults(smargon_coord, com, scores, v2_results) def _smargon_offset_from_scan_result( self, scan_result: ScanResult, request: RasterGridRequest, com: CenterOfMassModel -- 2.54.0