WIP agent step 5: migrate devices to v4 scan info
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
"""Utility helpers for SuperXAS devices."""
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user