diff --git a/superxas_bec/devices/mo1_bragg/mo1_bragg.py b/superxas_bec/devices/mo1_bragg/mo1_bragg.py index e6311fc..89b1abf 100644 --- a/superxas_bec/devices/mo1_bragg/mo1_bragg.py +++ b/superxas_bec/devices/mo1_bragg/mo1_bragg.py @@ -12,13 +12,13 @@ from typing import Any, Literal from bec_lib.devicemanager import ScanInfo from bec_lib.logger import bec_logger +from bec_server.scan_server.scans.scan_base import ScanInfo as ScanServerScanInfo from ophyd import Component as Cpt from ophyd import DeviceStatus, Signal, StatusBase from ophyd.status import SubscriptionStatus, WaitTimeoutError from ophyd_devices import CompareStatus, ProgressSignal, TransitionStatus from ophyd_devices.interfaces.base_classes.psi_device_base import PSIDeviceBase from ophyd_devices.utils.errors import DeviceStopError -from pydantic import BaseModel, Field from typeguard import typechecked from superxas_bec.devices.mo1_bragg.mo1_bragg_devices import Mo1BraggPositioner @@ -33,6 +33,7 @@ from superxas_bec.devices.mo1_bragg.mo1_bragg_enums import ( TriggerControlSource, ) from superxas_bec.devices.mo1_bragg.mo1_bragg_utils import compute_spline +from superxas_bec.devices.utils.utils import fetch_scan_info # Initialise logger logger = bec_logger.logger @@ -44,34 +45,6 @@ class Mo1BraggError(Exception): """Exception for the Mo1 Bragg positioner""" -########## Scan Parameter Model ########## - - -class ScanParameter(BaseModel): - """Dataclass to store the scan parameters for the Mo1 Bragg positioner. - This needs to be in sync with the kwargs of the MO1 Bragg scans from SuperXAS, to - ensure that the scan parameters are correctly set. Any changes in the scan kwargs, - i.e. renaming or adding new parameters, need to be represented here as well.""" - - scan_time: float | None = Field(None, description="Scan time for a half oscillation") - scan_duration: float | None = Field(None, description="Duration of the scan") - xrd_enable_low: bool | None = Field( - None, description="XRD enabled for low, should be PV trig_ena_lo_enum" - ) # trig_enable_low: bool = None - xrd_enable_high: bool | None = Field( - None, description="XRD enabled for high, should be PV trig_ena_hi_enum" - ) # trig_enable_high: bool = None - exp_time_low: float | None = Field(None, description="Exposure time low energy/angle") - exp_time_high: float | None = Field(None, description="Exposure time high energy/angle") - cycle_low: int | None = Field(None, description="Cycle for low energy/angle") - cycle_high: int | None = Field(None, description="Cycle for high energy/angle") - start: float | None = Field(None, description="Start value for energy/angle") - stop: float | None = Field(None, description="Stop value for energy/angle") - p_kink: float | None = Field(None, description="P Kink") - e_kink: float | None = Field(None, description="Energy Kink") - model_config: dict = {"validate_assignment": True} - - ########### Mo1 Bragg Motor Class ########### @@ -83,7 +56,7 @@ class Mo1Bragg(PSIDeviceBase, Mo1BraggPositioner): progress_signal = Cpt(ProgressSignal, name="progress_signal") - USER_ACCESS = ["set_advanced_xas_settings", "set_xtal"] + USER_ACCESS = ["set_advanced_xas_settings", "set_xtal", "convert_angle_energy"] def __init__(self, name: str, prefix: str = "", scan_info: ScanInfo | None = None, **kwargs): # type: ignore """ @@ -94,8 +67,15 @@ class Mo1Bragg(PSIDeviceBase, Mo1BraggPositioner): scan_info (ScanInfo): The scan info to use. """ super().__init__(name=name, scan_info=scan_info, prefix=prefix, **kwargs) - self.scan_parameter = ScanParameter() + self.scan_parameters: ScanServerScanInfo | None = None self.timeout_for_pvwait = 7.5 + self.valid_scan_names = [ + "xas_simple_scan", + "xas_simple_scan_with_xrd", + "xas_advanced_scan", + "xas_advanced_scan_with_xrd", + "nidaq_continuous_scan", + ] ######################################## # Beamline Specific Implementations # @@ -122,6 +102,7 @@ class Mo1Bragg(PSIDeviceBase, Mo1BraggPositioner): Information about the upcoming scan can be accessed from the scan_info (self.scan_info.msg) object. """ + self.scan_parameters = fetch_scan_info(self.scan_info) if self.scan_control.scan_msg.get() != ScanControlLoadMessage.PENDING: status = CompareStatus(self.scan_control.scan_msg, ScanControlLoadMessage.PENDING) @@ -129,14 +110,22 @@ class Mo1Bragg(PSIDeviceBase, Mo1BraggPositioner): self.scan_control.scan_val_reset.put(1) status.wait(timeout=self.timeout_for_pvwait) - scan_name = self.scan_info.msg.scan_name - self._update_scan_parameter() + scan_name = self.scan_parameters.scan_name + if not self._check_if_scan_name_is_valid(self.scan_parameters): + return None + start, stop = self._get_start_stop() + scan_time = self._scan_parameter("scan_time") + scan_duration = self._scan_parameter("scan_duration") if scan_name == "xas_simple_scan": + self._raise_for_missing( + scan_name, ["start", "stop", "scan_time", "scan_duration"], + [start, stop, scan_time, scan_duration], + ) self.set_xas_settings( - low=self.scan_parameter.start, - high=self.scan_parameter.stop, - scan_time=self.scan_parameter.scan_time, + low=start, + high=stop, + scan_time=scan_time, ) self.set_trig_settings( enable_low=False, @@ -147,32 +136,72 @@ class Mo1Bragg(PSIDeviceBase, Mo1BraggPositioner): cycle_high=0, ) self.set_scan_control_settings( - mode=ScanControlMode.SIMPLE, scan_duration=self.scan_parameter.scan_duration + mode=ScanControlMode.SIMPLE, scan_duration=scan_duration ) elif scan_name == "xas_simple_scan_with_xrd": + xrd_enable_low = self._scan_parameter("xrd_enable_low", "break_enable_low") + xrd_enable_high = self._scan_parameter("xrd_enable_high", "break_enable_high") + exp_time_low = self._scan_parameter("exp_time_low", "break_time_low") + exp_time_high = self._scan_parameter("exp_time_high", "break_time_high") + cycle_low = self._scan_parameter("cycle_low") + cycle_high = self._scan_parameter("cycle_high") + self._raise_for_missing( + scan_name, + [ + "start", + "stop", + "scan_time", + "scan_duration", + "xrd_enable_low", + "xrd_enable_high", + "exp_time_low", + "exp_time_high", + "cycle_low", + "cycle_high", + ], + [ + start, + stop, + scan_time, + scan_duration, + xrd_enable_low, + xrd_enable_high, + exp_time_low, + exp_time_high, + cycle_low, + cycle_high, + ], + ) self.set_xas_settings( - low=self.scan_parameter.start, - high=self.scan_parameter.stop, - scan_time=self.scan_parameter.scan_time, + low=start, + high=stop, + scan_time=scan_time, ) self.set_trig_settings( - enable_low=self.scan_parameter.xrd_enable_low, # enable_low=self.scan_parameter.trig_enable_low, - enable_high=self.scan_parameter.xrd_enable_high, # enable_high=self.scan_parameter.trig_enable_high, - exp_time_low=self.scan_parameter.exp_time_low, - exp_time_high=self.scan_parameter.exp_time_high, - cycle_low=self.scan_parameter.cycle_low, - cycle_high=self.scan_parameter.cycle_high, + enable_low=xrd_enable_low, + enable_high=xrd_enable_high, + exp_time_low=exp_time_low, + exp_time_high=exp_time_high, + cycle_low=cycle_low, + cycle_high=cycle_high, ) self.set_scan_control_settings( - mode=ScanControlMode.SIMPLE, scan_duration=self.scan_parameter.scan_duration + mode=ScanControlMode.SIMPLE, scan_duration=scan_duration ) elif scan_name == "xas_advanced_scan": + p_kink = self._scan_parameter("p_kink") + e_kink = self._scan_parameter("e_kink") + self._raise_for_missing( + scan_name, + ["start", "stop", "scan_time", "scan_duration", "p_kink", "e_kink"], + [start, stop, scan_time, scan_duration, p_kink, e_kink], + ) self.set_advanced_xas_settings( - low=self.scan_parameter.start, - high=self.scan_parameter.stop, - scan_time=self.scan_parameter.scan_time, - p_kink=self.scan_parameter.p_kink, - e_kink=self.scan_parameter.e_kink, + low=start, + high=stop, + scan_time=scan_time, + p_kink=p_kink, + e_kink=e_kink, ) self.set_trig_settings( enable_low=False, @@ -183,26 +212,65 @@ class Mo1Bragg(PSIDeviceBase, Mo1BraggPositioner): cycle_high=0, ) self.set_scan_control_settings( - mode=ScanControlMode.ADVANCED, scan_duration=self.scan_parameter.scan_duration + mode=ScanControlMode.ADVANCED, scan_duration=scan_duration ) elif scan_name == "xas_advanced_scan_with_xrd": + p_kink = self._scan_parameter("p_kink") + e_kink = self._scan_parameter("e_kink") + xrd_enable_low = self._scan_parameter("xrd_enable_low", "break_enable_low") + xrd_enable_high = self._scan_parameter("xrd_enable_high", "break_enable_high") + exp_time_low = self._scan_parameter("exp_time_low", "break_time_low") + exp_time_high = self._scan_parameter("exp_time_high", "break_time_high") + cycle_low = self._scan_parameter("cycle_low") + cycle_high = self._scan_parameter("cycle_high") + self._raise_for_missing( + scan_name, + [ + "start", + "stop", + "scan_time", + "scan_duration", + "p_kink", + "e_kink", + "xrd_enable_low", + "xrd_enable_high", + "exp_time_low", + "exp_time_high", + "cycle_low", + "cycle_high", + ], + [ + start, + stop, + scan_time, + scan_duration, + p_kink, + e_kink, + xrd_enable_low, + xrd_enable_high, + exp_time_low, + exp_time_high, + cycle_low, + cycle_high, + ], + ) self.set_advanced_xas_settings( - low=self.scan_parameter.start, - high=self.scan_parameter.stop, - scan_time=self.scan_parameter.scan_time, - p_kink=self.scan_parameter.p_kink, - e_kink=self.scan_parameter.e_kink, + low=start, + high=stop, + scan_time=scan_time, + p_kink=p_kink, + e_kink=e_kink, ) self.set_trig_settings( - enable_low=self.scan_parameter.xrd_enable_low, # enable_low=self.scan_parameter.trig_enable_low, - enable_high=self.scan_parameter.xrd_enable_high, # enable_high=self.scan_parameter.trig_enable_high, - exp_time_low=self.scan_parameter.exp_time_low, - exp_time_high=self.scan_parameter.exp_time_high, - cycle_low=self.scan_parameter.cycle_low, - cycle_high=self.scan_parameter.cycle_high, + enable_low=xrd_enable_low, + enable_high=xrd_enable_high, + exp_time_low=exp_time_low, + exp_time_high=exp_time_high, + cycle_low=cycle_low, + cycle_high=cycle_high, ) self.set_scan_control_settings( - mode=ScanControlMode.ADVANCED, scan_duration=self.scan_parameter.scan_duration + mode=ScanControlMode.ADVANCED, scan_duration=scan_duration ) else: return @@ -284,6 +352,46 @@ class Mo1Bragg(PSIDeviceBase, Mo1BraggPositioner): self.stopped = True # Needs to be set to stop motion ######### Utility Methods ######### + def _check_if_scan_name_is_valid(self, scan_parameters: ScanServerScanInfo) -> bool: + """Check if the scan is within the list of scans supported by the backend.""" + return scan_parameters.scan_name in self.valid_scan_names + + def _scan_parameter(self, *names: str): + """Fetch a scan parameter from v4 metadata, with legacy fallbacks.""" + if self.scan_parameters is None: + return None + sources = [ + self.scan_parameters.additional_scan_parameters, + getattr(self.scan_parameters, "metadata", {}), + getattr(self.scan_info.msg, "scan_parameters", {}), + ] + request_inputs = self.scan_parameters.request_inputs or {} + sources.extend( + [ + request_inputs.get("inputs", {}), + request_inputs.get("kwargs", {}), + ] + ) + for source in sources: + for name in names: + if isinstance(source, dict) and name in source: + return source[name] + return None + + def _get_start_stop(self): + """Return scan start/stop from v4 positions or legacy request inputs.""" + if self.scan_parameters is not None and self.scan_parameters.positions is not None: + if len(self.scan_parameters.positions) == 2: + return self.scan_parameters.positions + return self._scan_parameter("start"), self._scan_parameter("stop") + + def _raise_for_missing(self, scan_name: str, names: list[str], values: list) -> None: + if any(value is None for value in values): + raise Mo1BraggError( + f"Missing scan parameters for {scan_name}. Required parameters: " + f"{', '.join(names)} in scan info {self.scan_parameters}" + ) + def _progress_update(self, value, **kwargs) -> None: """Callback method to update the scan progress, runs a callback to SUB_PROGRESS subscribers, i.e. BEC. @@ -454,13 +562,3 @@ class Mo1Bragg(PSIDeviceBase, Mo1BraggPositioner): for s in status_list: s.wait(timeout=self.timeout_for_pvwait) - - - def _update_scan_parameter(self): - """Get the scan_info parameters for the scan.""" - for key, value in self.scan_info.msg.request_inputs["inputs"].items(): - if hasattr(self.scan_parameter, key): - setattr(self.scan_parameter, key, value) - for key, value in self.scan_info.msg.request_inputs["kwargs"].items(): - if hasattr(self.scan_parameter, key): - setattr(self.scan_parameter, key, value) diff --git a/superxas_bec/devices/nidaq/nidaq.py b/superxas_bec/devices/nidaq/nidaq.py index 17fcbc4..3ac1664 100644 --- a/superxas_bec/devices/nidaq/nidaq.py +++ b/superxas_bec/devices/nidaq/nidaq.py @@ -1,9 +1,10 @@ from __future__ import annotations import time -from typing import TYPE_CHECKING, Literal, cast +from typing import TYPE_CHECKING, Literal from bec_lib.logger import bec_logger +from bec_server.scan_server.scans.scan_base import ScanInfo as ScanServerScanInfo from ophyd import Component as Cpt from ophyd import Device, DeviceStatus, EpicsSignal, EpicsSignalRO, Kind, StatusBase from ophyd.status import SubscriptionStatus, WaitTimeoutError @@ -19,6 +20,7 @@ from superxas_bec.devices.nidaq.nidaq_enums import ( ScanRates, ScanType, ) +from superxas_bec.devices.utils.utils import fetch_scan_info if TYPE_CHECKING: # pragma: no cover from bec_lib.devicemanager import ScanInfo @@ -425,7 +427,7 @@ class Nidaq(PSIDeviceBase, NidaqControl): def __init__(self, prefix: str = "", *, name: str, scan_info: ScanInfo = None, **kwargs): super().__init__(name=name, prefix=prefix, scan_info=scan_info, **kwargs) - self.scan_info: ScanInfo + self.scan_parameters: ScanServerScanInfo | None = None self.timeout_wait_for_signal = 5 # put 5s firsts self._timeout_wait_for_pv = 5 # 5s timeout for pv calls. editted due to timeout issues persisting self.valid_scan_names = [ @@ -440,9 +442,11 @@ class Nidaq(PSIDeviceBase, NidaqControl): # Beamline Methods # ######################################## - def _check_if_scan_name_is_valid(self) -> bool: + def _check_if_scan_name_is_valid(self, scan_parameters: ScanServerScanInfo | None) -> bool: """Check if the scan is within the list of scans for which the backend is working""" - scan_name = self.scan_info.msg.scan_name + if scan_parameters is None: + return False + scan_name = scan_parameters.scan_name if scan_name in self.valid_scan_names: return True return False @@ -600,7 +604,8 @@ class Nidaq(PSIDeviceBase, NidaqControl): Information about the upcoming scan can be accessed from the scan_info (self.scan_info.msg) object. If the upcoming scan is not in the list of valid scans, return immediately. """ - if not self._check_if_scan_name_is_valid(): + self.scan_parameters = fetch_scan_info(self.scan_info) + if not self._check_if_scan_name_is_valid(self.scan_parameters): return None if self.state.get() != NidaqState.STANDBY: @@ -610,16 +615,16 @@ class Nidaq(PSIDeviceBase, NidaqControl): status.wait(timeout=self.timeout_wait_for_signal) # If scan is not part of the valid_scan_names, - if self.scan_info.msg.scan_name != "nidaq_continuous_scan": + if self.scan_parameters.scan_name != "nidaq_continuous_scan": self.scan_type.set(ScanType.TRIGGERED).wait(timeout=self._timeout_wait_for_pv) self.scan_duration.set(0).wait(timeout=self._timeout_wait_for_pv) self.enable_compression.set(1).wait(timeout=self._timeout_wait_for_pv) else: self.scan_type.set(ScanType.CONTINUOUS).wait(timeout=self._timeout_wait_for_pv) - self.scan_duration.set(self.scan_info.msg.scan_parameters["scan_duration"]).wait( + self.scan_duration.set(self._scan_parameter("scan_duration")).wait( timeout=self._timeout_wait_for_pv ) - self.enable_compression.set(self.scan_info.msg.scan_parameters["compression"]).wait( + self.enable_compression.set(self._scan_parameter("compression")).wait( timeout=self._timeout_wait_for_pv ) @@ -632,7 +637,7 @@ class Nidaq(PSIDeviceBase, NidaqControl): # self.stage_call.set(1).wait(timeout=self._timeout_wait_for_pv) self.stage_call.put(1) status.wait(timeout=self.timeout_wait_for_signal) - if self.scan_info.msg.scan_name != "nidaq_continuous_scan": + if self.scan_parameters.scan_name != "nidaq_continuous_scan": status = self.on_kickoff() self.cancel_on_stop(status) status.wait(timeout=self._timeout_wait_for_pv) @@ -663,10 +668,10 @@ class Nidaq(PSIDeviceBase, NidaqControl): before the motor starts its oscillation. This is needed for being properly homed. The NIDAQ should go into Acquiring mode. """ - if not self._check_if_scan_name_is_valid(): + if not self._check_if_scan_name_is_valid(self.scan_parameters): return None - if self.scan_info.msg.scan_name == "nidaq_continuous_scan": + if self.scan_parameters.scan_name == "nidaq_continuous_scan": logger.info(f"Device {self.name} ready to be kicked off for nidaq_continuous_scan") return None @@ -687,15 +692,35 @@ class Nidaq(PSIDeviceBase, NidaqControl): For the NIDAQ we use this method to stop the backend since it would not stop by itself in its current implementation since the number of points are not predefined. """ - if not self._check_if_scan_name_is_valid(): + if not self._check_if_scan_name_is_valid(self.scan_parameters): return None status = CompareStatus(self.state, NidaqState.STANDBY) self.cancel_on_stop(status) - if self.scan_info.msg.scan_name != "nidaq_continuous_scan": + if self.scan_parameters.scan_name != "nidaq_continuous_scan": self.on_stop() return status + def _scan_parameter(self, name: str): + """Fetch a scan parameter from v4 metadata, with legacy fallbacks.""" + if self.scan_parameters is None: + return None + sources = [ + self.scan_parameters.additional_scan_parameters, + getattr(self.scan_info.msg, "scan_parameters", {}), + ] + request_inputs = self.scan_parameters.request_inputs or {} + sources.extend( + [ + request_inputs.get("inputs", {}), + request_inputs.get("kwargs", {}), + ] + ) + for source in sources: + if isinstance(source, dict) and name in source: + return source[name] + return None + def _progress_update(self, value, **kwargs) -> None: """Callback method to update the scan progress, runs a callback to SUB_PROGRESS subscribers, i.e. BEC. @@ -703,7 +728,9 @@ class Nidaq(PSIDeviceBase, NidaqControl): Args: value (int) : current progress value """ - scan_duration = self.scan_info.msg.scan_parameters.get("scan_duration", None) + if self.scan_parameters is None: + return + scan_duration = self._scan_parameter("scan_duration") if not isinstance(scan_duration, (int, float)): return value = scan_duration - value diff --git a/superxas_bec/devices/utils/__init__.py b/superxas_bec/devices/utils/__init__.py new file mode 100644 index 0000000..eb07fbf --- /dev/null +++ b/superxas_bec/devices/utils/__init__.py @@ -0,0 +1,2 @@ +"""Utility helpers for SuperXAS devices.""" + diff --git a/superxas_bec/devices/utils/utils.py b/superxas_bec/devices/utils/utils.py new file mode 100644 index 0000000..6cff1eb --- /dev/null +++ b/superxas_bec/devices/utils/utils.py @@ -0,0 +1,26 @@ +"""Utility functions for SuperXAS devices.""" + +from copy import deepcopy + +import numpy as np +from bec_lib.devicemanager import ScanInfo +from bec_server.scan_server.scans.scan_base import ScanInfo as ScanServerScanInfo +from pydantic import ValidationError + + +def fetch_scan_info(scan_info: ScanInfo) -> ScanServerScanInfo: + """Normalize BEC scan info into the v4 scan info model.""" + info = deepcopy(scan_info.msg.info) + if isinstance(info.get("positions"), list): + info["positions"] = np.array(info["positions"]) + try: + msg = ScanServerScanInfo.model_validate(info) + except ValidationError: + if info.get("scan_type") == "fly": + info["scan_type"] = "hardware_triggered" + else: + info["scan_type"] = "software_triggered" + msg = ScanServerScanInfo.model_validate(info) + + return msg + diff --git a/tests/tests_devices/test_device_scan_info_utils.py b/tests/tests_devices/test_device_scan_info_utils.py new file mode 100644 index 0000000..5fea25e --- /dev/null +++ b/tests/tests_devices/test_device_scan_info_utils.py @@ -0,0 +1,53 @@ +# pylint: skip-file +from types import SimpleNamespace + +import numpy as np + +from superxas_bec.devices.utils.utils import fetch_scan_info + + +def test_fetch_scan_info_accepts_v4_scan_info_with_positions_list(): + scan_info = SimpleNamespace( + msg=SimpleNamespace( + info={ + "scan_name": "xas_simple_scan", + "scan_id": "scan-id-test", + "scan_type": "hardware_triggered", + "positions": [8000.0, 9000.0], + "additional_scan_parameters": { + "scan_time": 1.0, + "scan_duration": 10.0, + }, + } + ) + ) + + msg = fetch_scan_info(scan_info) + + assert msg.scan_name == "xas_simple_scan" + assert msg.scan_type == "hardware_triggered" + np.testing.assert_array_equal(msg.positions, np.array([8000.0, 9000.0])) + assert msg.additional_scan_parameters["scan_duration"] == 10.0 + + +def test_fetch_scan_info_converts_legacy_fly_scan_type(): + scan_info = SimpleNamespace( + msg=SimpleNamespace( + info={ + "scan_name": "xas_simple_scan", + "scan_id": "scan-id-test", + "scan_type": "fly", + "positions": [8000.0, 9000.0], + "request_inputs": { + "inputs": {}, + "kwargs": {"scan_time": 1.0, "scan_duration": 10.0}, + }, + } + ) + ) + + msg = fetch_scan_info(scan_info) + + assert msg.scan_type == "hardware_triggered" + assert msg.request_inputs["kwargs"]["scan_time"] == 1.0 + diff --git a/tests/tests_devices/test_mo1_bragg_v4_scan_info.py b/tests/tests_devices/test_mo1_bragg_v4_scan_info.py new file mode 100644 index 0000000..2cb6fa7 --- /dev/null +++ b/tests/tests_devices/test_mo1_bragg_v4_scan_info.py @@ -0,0 +1,100 @@ +# pylint: skip-file +from unittest import mock + +import numpy as np +import pytest +from bec_server.scan_server.scans.scan_base import ScanInfo as ScanServerScanInfo +from ophyd_devices.tests.utils import patched_device + +from superxas_bec.devices.mo1_bragg.mo1_bragg import Mo1Bragg, ScanControlLoadMessage + + +@pytest.fixture(scope="function") +def mock_bragg(): + with patched_device( + Mo1Bragg, name="mo1_bragg", prefix="X10DA-OP-MO1:BRAGG:" + ) as dev: + yield dev + + +def _set_scan_info(dev, scan_info): + dev.scan_info.msg.info.update(scan_info.model_dump()) + + +def _mock_status(): + status = mock.MagicMock() + status.wait = mock.MagicMock() + return status + + +def test_mo1_bragg_stage_uses_v4_simple_scan_info(mock_bragg): + scan_info = ScanServerScanInfo( + scan_name="xas_simple_scan", + scan_id="scan-id-test", + scan_type="hardware_triggered", + positions=np.array([8000.0, 9000.0]), + additional_scan_parameters={"scan_time": 1.0, "scan_duration": 10.0}, + ) + _set_scan_info(mock_bragg, scan_info) + mock_bragg.scan_control.scan_msg._read_pv.mock_data = ScanControlLoadMessage.PENDING + + with ( + mock.patch.object(mock_bragg, "set_xas_settings") as set_xas_settings, + mock.patch.object(mock_bragg, "set_trig_settings") as set_trig_settings, + mock.patch.object(mock_bragg, "set_scan_control_settings") as set_scan_control_settings, + mock.patch.object(mock_bragg, "cancel_on_stop"), + mock.patch( + "superxas_bec.devices.mo1_bragg.mo1_bragg.CompareStatus", + return_value=_mock_status(), + ), + ): + mock_bragg.on_stage() + + set_xas_settings.assert_called_once_with(low=8000.0, high=9000.0, scan_time=1.0) + set_trig_settings.assert_called_once_with( + enable_low=False, + enable_high=False, + exp_time_low=0, + exp_time_high=0, + cycle_low=0, + cycle_high=0, + ) + assert set_scan_control_settings.call_args.kwargs["scan_duration"] == 10.0 + + +def test_mo1_bragg_stage_uses_v4_advanced_scan_info(mock_bragg): + scan_info = ScanServerScanInfo( + scan_name="xas_advanced_scan", + scan_id="scan-id-test", + scan_type="hardware_triggered", + positions=np.array([8000.0, 9000.0]), + additional_scan_parameters={ + "scan_time": 1.0, + "scan_duration": 10.0, + "p_kink": 50.0, + "e_kink": 8500.0, + }, + ) + _set_scan_info(mock_bragg, scan_info) + mock_bragg.scan_control.scan_msg._read_pv.mock_data = ScanControlLoadMessage.PENDING + + with ( + mock.patch.object(mock_bragg, "set_advanced_xas_settings") as set_advanced_xas_settings, + mock.patch.object(mock_bragg, "set_trig_settings"), + mock.patch.object(mock_bragg, "set_scan_control_settings") as set_scan_control_settings, + mock.patch.object(mock_bragg, "cancel_on_stop"), + mock.patch( + "superxas_bec.devices.mo1_bragg.mo1_bragg.CompareStatus", + return_value=_mock_status(), + ), + ): + mock_bragg.on_stage() + + set_advanced_xas_settings.assert_called_once_with( + low=8000.0, + high=9000.0, + scan_time=1.0, + p_kink=50.0, + e_kink=8500.0, + ) + assert set_scan_control_settings.call_args.kwargs["scan_duration"] == 10.0 diff --git a/tests/tests_devices/test_nidaq_v4_scan_info.py b/tests/tests_devices/test_nidaq_v4_scan_info.py new file mode 100644 index 0000000..53a5f7b --- /dev/null +++ b/tests/tests_devices/test_nidaq_v4_scan_info.py @@ -0,0 +1,67 @@ +# pylint: skip-file +from unittest import mock + +import pytest +from bec_server.scan_server.scans.scan_base import ScanInfo as ScanServerScanInfo +from ophyd_devices.tests.utils import patched_device + +from superxas_bec.devices.nidaq.nidaq import Nidaq, NidaqState, ScanType +from superxas_bec.devices.utils.utils import fetch_scan_info + + +@pytest.fixture(scope="function") +def mock_nidaq(): + with patched_device( + Nidaq, name="nidaq", prefix="X10DA-CPCL-SCANSERVER:" + ) as dev: + yield dev + + +def _set_scan_info(dev, scan_info): + dev.scan_info.msg.info.update(scan_info.model_dump()) + dev.scan_parameters = fetch_scan_info(dev.scan_info) + + +def test_nidaq_check_scan_name_uses_normalized_scan_info(mock_nidaq): + valid = ScanServerScanInfo(scan_name="xas_simple_scan", scan_id="scan-id-test") + invalid = ScanServerScanInfo(scan_name="line_scan", scan_id="scan-id-test") + + assert mock_nidaq._check_if_scan_name_is_valid(valid) + assert not mock_nidaq._check_if_scan_name_is_valid(invalid) + + +def test_nidaq_progress_update_uses_v4_additional_parameters(mock_nidaq): + scan_info = ScanServerScanInfo( + scan_name="nidaq_continuous_scan", + scan_id="scan-id-test", + additional_scan_parameters={"scan_duration": 10.0, "compression": False}, + ) + _set_scan_info(mock_nidaq, scan_info) + mock_nidaq.progress_signal.put = mock.MagicMock() + + mock_nidaq._progress_update(4.0) + + mock_nidaq.progress_signal.put.assert_called_once_with(value=6.0, max_value=10.0, done=False) + + +def test_nidaq_stage_uses_v4_continuous_scan_info(mock_nidaq): + scan_info = ScanServerScanInfo( + scan_name="nidaq_continuous_scan", + scan_id="scan-id-test", + additional_scan_parameters={"scan_duration": 10.0, "compression": False}, + ) + mock_nidaq.scan_info.msg.info.update(scan_info.model_dump()) + mock_nidaq.state.put(NidaqState.STANDBY) + + with ( + mock.patch("superxas_bec.devices.nidaq.nidaq.CompareStatus") as status_cls, + mock.patch.object(mock_nidaq, "cancel_on_stop"), + mock.patch.object(mock_nidaq, "on_kickoff") as on_kickoff, + ): + status_cls.return_value.wait = mock.MagicMock() + mock_nidaq.on_stage() + + assert mock_nidaq.scan_type.get() == ScanType.CONTINUOUS + assert mock_nidaq.scan_duration.get() == 10.0 + assert mock_nidaq.enable_compression.get() is False + on_kickoff.assert_not_called()