WIP agent step 5: migrate devices to v4 scan info
CI for superxas_bec / test (push) Successful in 1m31s
CI for superxas_bec / test (pull_request) Successful in 1m28s

This commit is contained in:
2026-05-28 08:44:21 +02:00
parent e42dc75c38
commit e47ce01cc3
7 changed files with 462 additions and 89 deletions
+173 -75
View File
@@ -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)
+41 -14
View File
@@ -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
+2
View File
@@ -0,0 +1,2 @@
"""Utility helpers for SuperXAS devices."""
+26
View File
@@ -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
@@ -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
@@ -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
@@ -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()