From d1ddae4922753bc1207ff816c28657cce1f8ba07 Mon Sep 17 00:00:00 2001 From: appel_c Date: Tue, 21 Apr 2026 16:10:19 +0200 Subject: [PATCH] w --- csaxs_bec/scans/scans_v4/cont_line_scan.py | 2 +- tests/tests_scans/__init__.py | 0 tests/tests_scans/conftest.py | 302 ++++++++++++++++++ tests/tests_scans/scan_test_utils.py | 138 ++++++++ .../tests_scans/test_cont_line_scan_cSAXS.py | 12 + 5 files changed, 453 insertions(+), 1 deletion(-) create mode 100644 tests/tests_scans/__init__.py create mode 100644 tests/tests_scans/conftest.py create mode 100644 tests/tests_scans/scan_test_utils.py create mode 100644 tests/tests_scans/test_cont_line_scan_cSAXS.py diff --git a/csaxs_bec/scans/scans_v4/cont_line_scan.py b/csaxs_bec/scans/scans_v4/cont_line_scan.py index 221f4fb..4dcd0ec 100644 --- a/csaxs_bec/scans/scans_v4/cont_line_scan.py +++ b/csaxs_bec/scans/scans_v4/cont_line_scan.py @@ -104,7 +104,7 @@ class ContLineScanCSAXS(ScanBase): super().__init__(**kwargs) self.fast_motor = self.dev[fast_motor] if isinstance(fast_motor, str) else fast_motor - self.ddg = self.dev.get("ddg", None) + self.ddg = self.dev.get("ddg1", None) if self.ddg is None: raise ScanAbortion( "Main delay generator named 'ddg1' device is required for the continuous line scan." diff --git a/tests/tests_scans/__init__.py b/tests/tests_scans/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/tests_scans/conftest.py b/tests/tests_scans/conftest.py new file mode 100644 index 0000000..ff8a16e --- /dev/null +++ b/tests/tests_scans/conftest.py @@ -0,0 +1,302 @@ +import importlib +import inspect +import pkgutil +from types import SimpleNamespace +from typing import get_type_hints +from unittest import mock + +import pytest +from bec_lib import messages +from bec_lib.device import DeviceBase +from bec_lib.tests.utils import ConnectorMock +from bec_server.scan_server.instruction_handler import InstructionHandler +from bec_server.scan_server.scan_assembler import ScanAssembler +from bec_server.scan_server.scan_gui_models import GUIInput +from bec_server.scan_server.scans import ScanArgType +from bec_server.scan_server.scans.scans_v4 import ScanBase as ScanBaseV4 + + +class _DoneAfterNthCheckStatusMock: + def __init__(self, resolve_after: int = 1, result=None) -> None: + self.resolve_after = max(resolve_after, 1) + self.result = result + self.wait_calls = 0 + self._done_checks = 0 + + @property + def done(self): + self._done_checks += 1 + return self._done_checks >= self.resolve_after + + def wait(self, *args, **kwargs): + self.wait_calls += 1 + return self + + +@pytest.fixture +def nth_done_status_mock(): + def _build(resolve_after: int = 1, result=None): + return _DoneAfterNthCheckStatusMock(resolve_after=resolve_after, result=result) + + return _build + + +@pytest.fixture +def readout_priority(): + return SimpleNamespace( + monitored=[], baseline=["samx", "samy", "samz"], on_request=[], async_=[] + ) + + +class Signal: + def __init__(self, name, value=0.0): + self.name = name + self._value = value + + def read(self): + return {self.name: {"value": self._value}} + + def put(self, value): + self._value = value + + def get(self): + return self._value + + +class _MockV4Device(DeviceBase): + def __init__(self, name: str, limits=(-10.0, 10.0), value: float = 0.0): + info = { + "device_info": { + "signals": { + name: {"obj_name": name, "kind_str": "hinted", "describe": {"precision": 3}} + } + } + } + super().__init__(name=name, info=info) + self._limits = limits + self._value = value + self._enabled = True + self._precision = 3 + self.velocity = Signal(f"{name}_velocity", value=10.0) + self.acceleration = Signal(f"{name}_acceleration", value=0.5) + + def read(self, *args, **kwargs): + return {self.full_name: {"value": self._value}} + + @property + def root(self): + return self + + @property + def full_name(self): + return self.name + + @property + def limits(self): + return self._limits + + @property + def enabled(self): + return self._enabled + + @property + def precision(self): + return self._precision + + +class _MockV4Devices(dict): + def __init__(self, devices: dict[str, _MockV4Device], readout_priority: dict | None = None): + super().__init__(devices) + readout_priority = readout_priority or {} + self._base_readout_priority = { + "baseline": list(readout_priority.get("baseline", [])), + "monitored": list(readout_priority.get("monitored", [])), + "on_request": list(readout_priority.get("on_request", [])), + "async": list(readout_priority.get("async", [])), + } + + @property + def enabled_devices(self): + return list(self.values()) + + def _applied_readout_priority(self, readout_priority=None) -> dict[str, list[str]]: + groups = { + group_name: [device_name for device_name in device_names if device_name in self] + for group_name, device_names in self._base_readout_priority.items() + } + + for group_name in ["baseline", "monitored", "on_request", "async"]: + for device_name in (readout_priority or {}).get(group_name, []): + if device_name not in self: + continue + for existing_group in groups.values(): + if device_name in existing_group: + existing_group.remove(device_name) + groups[group_name].append(device_name) + + for group_name, device_names in groups.items(): + groups[group_name] = sorted(set(device_names)) + return groups + + def monitored_devices(self, readout_priority=None): + monitored = self._applied_readout_priority(readout_priority)["monitored"] + return [self[name] for name in monitored if name in self] + + def baseline_devices(self, readout_priority=None): + baseline = self._applied_readout_priority(readout_priority)["baseline"] + return [self[name] for name in baseline if name in self] + + def async_devices(self, readout_priority=None): + async_devices = self._applied_readout_priority(readout_priority)["async"] + return [self[name] for name in async_devices if name in self] + + def on_request_devices(self, readout_priority=None): + on_request = self._applied_readout_priority(readout_priority)["on_request"] + return [self[name] for name in on_request if name in self] + + def continuous_devices(self, readout_priority=None): + return [] + + def get_software_triggered_devices(self): + return [] + + +def _infer_v4_device_names(scan_cls, scan_args: tuple, scan_kwargs: dict) -> list[str]: + arg_input = getattr(scan_cls, "arg_input", {}) or {} + if not arg_input: + type_hints = get_type_hints(scan_cls.__init__) + signature = inspect.signature(scan_cls) + arg_input = { + name: type_hints.get(name, parameter.annotation) + for name, parameter in signature.parameters.items() + if name not in {"args", "kwargs"} + and parameter.annotation is not inspect.Parameter.empty + } + if not arg_input: + return [] + + device_names = [] + bundle_size = scan_cls.arg_bundle_size["bundle"] + + def _is_device_arg(arg_type) -> bool: + converted = GUIInput.convert_to_legacy_scan_arg_type(arg_type) + if converted == ScanArgType.DEVICE: + return True + return inspect.isclass(converted) and issubclass(converted, DeviceBase) + + if bundle_size > 0: + arg_names = list(arg_input.keys()) + for bundle_start in range(0, len(scan_args), bundle_size): + for offset, arg_name in enumerate(arg_names): + arg_index = bundle_start + offset + if arg_index >= len(scan_args): + break + if _is_device_arg(arg_input.get(arg_name)): + device_names.append(scan_args[arg_index]) + else: + bound = inspect.signature(scan_cls).bind_partial(*scan_args, **scan_kwargs) + for arg_name, value in bound.arguments.items(): + if arg_name == "args": + continue + if _is_device_arg(arg_input.get(arg_name)): + device_names.append(value) + + for arg_name, arg_type in arg_input.items(): + if _is_device_arg(arg_type) and arg_name in scan_kwargs: + device_names.append(scan_kwargs[arg_name]) + + return [name for name in device_names if isinstance(name, str)] + + +def _base_readout_priority(readout_priority) -> dict[str, list[str]]: + return { + "monitored": list(readout_priority.monitored), + "baseline": list(readout_priority.baseline), + "on_request": list(readout_priority.on_request), + "async": list(readout_priority.async_), + } + + +def _get_v4_scan_classes() -> dict[str, type[ScanBaseV4]]: + import bec_server.scan_server.scans as scans_v4_module + + import csaxs_bec.scans.scans_v4 as scans_v4_csaxs_module + + scan_classes = {} + for mod in (scans_v4_module, scans_v4_csaxs_module): + for module_info in pkgutil.iter_modules(mod.__path__, prefix=f"{mod.__name__}."): + module = importlib.import_module(module_info.name) + for _, scan_cls in inspect.getmembers(module, predicate=inspect.isclass): + if scan_cls.__module__ != module.__name__: + continue + if not issubclass(scan_cls, ScanBaseV4): + continue + scan_name = getattr(scan_cls, "scan_name", None) + if not scan_name or scan_name == "_v4_base_scan": + continue + scan_classes[scan_name] = scan_cls + if scan_name.startswith("_v4_"): + scan_classes[scan_name.removeprefix("_v4_")] = scan_cls + return scan_classes + + +@pytest.fixture +def v4_scan_assembler(readout_priority): + scan_classes = _get_v4_scan_classes() + + def _assemble_scan(scan_type, *scan_args, **scan_kwargs): + scan_id = scan_kwargs.pop("scan_id", "scan-id-test") + + try: + scan_cls = scan_classes[scan_type] + except KeyError as exc: + available = ", ".join(sorted(scan_classes)) + raise KeyError(f"Unknown v4 scan type '{scan_type}'. Available: {available}") from exc + + connector = ConnectorMock("") + instruction_handler = InstructionHandler(connector) + base_readout_priority = _base_readout_priority(readout_priority) + device_names = sorted( + set(_infer_v4_device_names(scan_cls, scan_args, scan_kwargs)) + | set(base_readout_priority["monitored"]) + | set(base_readout_priority["baseline"]) + | set(base_readout_priority["on_request"]) + | set(base_readout_priority["async"]) + ) + # Add ddg1 + device_names.append("ddg1") + devices = _MockV4Devices( + {name: _MockV4Device(name) for name in device_names}, + readout_priority=base_readout_priority, + ) + # Mock method for ddg1 shutter open delay call + devices["ddg1"].get_shutter_open_delay = mock.MagicMock(return_value=0.02) + device_manager = SimpleNamespace(devices=devices, connector=connector) + resolved_scan_kwargs = { + "system_config": {"file_directory": "/tmp/data/S00000"}, + **scan_kwargs, + } + + parent = mock.MagicMock() + parent.device_manager = device_manager + parent.connector = connector + parent.queue_manager.instruction_handler = instruction_handler + parent.scan_manager = SimpleNamespace(scan_dict={scan_type: scan_cls}) + + assembler = ScanAssembler(parent=parent) + msg = messages.ScanQueueMessage( + metadata={"RID": "rid-test"}, + scan_type=scan_type, + parameter={"args": list(scan_args), "kwargs": resolved_scan_kwargs}, + queue="primary", + ) + scan = assembler.assemble_direct_scan(msg, scan_id) + scan._test = SimpleNamespace( + connector=connector, + instruction_handler=instruction_handler, + device_manager=device_manager, + assembler=assembler, + ) + return scan + + return _assemble_scan diff --git a/tests/tests_scans/scan_test_utils.py b/tests/tests_scans/scan_test_utils.py new file mode 100644 index 0000000..d2855ac --- /dev/null +++ b/tests/tests_scans/scan_test_utils.py @@ -0,0 +1,138 @@ +from unittest import mock + + +def assert_prepare_scan_reads_baseline_devices(scan): + baseline_status = mock.MagicMock() + scan.actions.read_baseline_devices = mock.MagicMock(return_value=baseline_status) + + scan.prepare_scan() + + scan.actions.read_baseline_devices.assert_called_once_with(wait=False) + assert scan._baseline_readout_status is baseline_status + + +def assert_prepare_scan_starts_premove_move(scan): + premove_status = mock.MagicMock() + scan.actions.set = mock.MagicMock(return_value=premove_status) + + scan.prepare_scan() + + assert scan.actions.set.call_count >= 1 + assert any(call.kwargs.get("wait") is False for call in scan.actions.set.call_args_list) + assert scan._premove_motor_status is premove_status + + +def assert_scan_open_called(scan): + scan.actions.open_scan = mock.MagicMock() + + scan.open_scan() + + scan.actions.open_scan.assert_called_once_with() + + +def assert_stage_all_devices_called(scan): + scan.actions.stage_all_devices = mock.MagicMock() + + scan.stage() + + scan.actions.stage_all_devices.assert_called_once_with() + + +def assert_pre_scan_called(scan): + scan._premove_motor_status = mock.MagicMock() + scan.actions.pre_scan = mock.MagicMock() + + scan.pre_scan() + + scan.actions.pre_scan.assert_called_once_with() + + +def assert_pre_scan_waits_for_premove(scan): + premove_status = mock.MagicMock() + scan._premove_motor_status = premove_status + scan.actions.pre_scan = mock.MagicMock() + + scan.pre_scan() + + premove_status.wait.assert_called_once_with() + scan.actions.pre_scan.assert_called_once_with() + + +def assert_unstage_all_devices_called(scan): + scan.actions.unstage_all_devices = mock.MagicMock() + + scan.unstage() + + scan.actions.unstage_all_devices.assert_called_once_with() + + +def assert_close_scan_waits_for_baseline_and_closes(scan, nth_done_status_mock): + baseline_status = nth_done_status_mock(resolve_after=2) + scan._baseline_readout_status = baseline_status + scan.actions.close_scan = mock.MagicMock() + scan.actions.check_for_unchecked_statuses = mock.MagicMock() + + scan.close_scan() + + assert baseline_status.wait_calls == 1 + scan.actions.close_scan.assert_called_once_with() + scan.actions.check_for_unchecked_statuses.assert_called_once_with() + + +def assert_scan_core_delegates_to_step_scan(scan): + scan.prepare_scan() + scan.components.step_scan = mock.MagicMock() + + scan.scan_core() + + scan.components.step_scan.assert_called_once() + args, kwargs = scan.components.step_scan.call_args + assert args == (scan.motors, scan.scan_info.positions) + assert kwargs["at_each_point"] == scan.at_each_point + if "last_positions" in kwargs: + assert (kwargs["last_positions"] == scan.positions[0]).all() + + +def assert_post_scan_waits_for_completion_and_moves_back_when_relative(scan, nth_done_status_mock): + completion_status = nth_done_status_mock(resolve_after=3) + scan.relative = True + scan.start_positions = [1.2, -0.7] + scan.actions.complete_all_devices = mock.MagicMock(return_value=completion_status) + scan.components.move_and_wait = mock.MagicMock() + + scan.post_scan() + + scan.actions.complete_all_devices.assert_called_once_with(wait=False) + scan.components.move_and_wait.assert_called_once_with(scan.motors, scan.start_positions) + assert completion_status.wait_calls == 1 + + +DEFAULT_HOOK_TESTS = [ + ("prepare_scan", [assert_prepare_scan_reads_baseline_devices]), + ("open_scan", [assert_scan_open_called]), + ("stage", [assert_stage_all_devices_called]), + ("pre_scan", [assert_pre_scan_called]), + ("unstage", [assert_unstage_all_devices_called]), + ("close_scan", [assert_close_scan_waits_for_baseline_and_closes]), +] + + +PREMOVE_HOOK_TESTS = [ + ("prepare_scan", [assert_prepare_scan_starts_premove_move]), + ("pre_scan", [assert_pre_scan_waits_for_premove]), +] + + +STANDARD_STEP_SCAN_TESTS = [ + ("scan_core", [assert_scan_core_delegates_to_step_scan]), + ("post_scan", [assert_post_scan_waits_for_completion_and_moves_back_when_relative]), +] + + +def run_scan_tests(scan, tests, nth_done_status_mock=None): + for test_name, assertions in tests: + for assertion in assertions: + if test_name in {"close_scan", "post_scan"}: + assertion(scan, nth_done_status_mock) + else: + assertion(scan) diff --git a/tests/tests_scans/test_cont_line_scan_cSAXS.py b/tests/tests_scans/test_cont_line_scan_cSAXS.py new file mode 100644 index 0000000..eeed7da --- /dev/null +++ b/tests/tests_scans/test_cont_line_scan_cSAXS.py @@ -0,0 +1,12 @@ +import pytest + +from .scan_test_utils import DEFAULT_HOOK_TESTS, run_scan_tests + + +@pytest.mark.parametrize(("hook_name", "hook_tests"), DEFAULT_HOOK_TESTS) +def test_cont_line_scan(v4_scan_assembler, nth_done_status_mock, hook_name, hook_tests): + scan = v4_scan_assembler( + "cont_line_scan_csaxs", fast_motor="samx", start=0, stop=5, step_size=0.5, exp_time=0.1 + ) + + run_scan_tests(scan, [(hook_name, hook_tests)], nth_done_status_mock)