Files
AareDAQ/tests/unit/daq/test_mount.py
T

310 lines
11 KiB
Python

from datetime import UTC, datetime, timedelta
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from aarecommon.errors.exception_handler import BECCommunicationError, JFJochCommunicationError
from aarecommon.models.models import DAQOperation, DewarAddress, SampleShortInfo
from aarecommon.models.tell import TellActivityEnum, TellPhaseEnum, TellStateModel
from aareDB import SampleEventType
from aare.daq.config import ABR_POS_MOUNT, BeamlineStateEnum
from aare.daq.daq import AareDAQ
from aare.daq.operations.mounting.models import MountingContext, MountingResult
from aare.daq.operations.screenshot.service import ScreenshotService
def _make_sample(sample_id: int, name: str) -> SampleShortInfo:
return SampleShortInfo(
db_id=sample_id,
puck_name="puck1",
dewar_name="dew1",
sample_name=name,
run_number=1,
user="group1",
pin=sample_id,
location=DewarAddress(segment="A", pos=1),
)
def _make_daq(previous_sample: SampleShortInfo | None) -> AareDAQ:
daq = object.__new__(AareDAQ)
daq._cfg = SimpleNamespace(current_sample=previous_sample)
daq._aare = MagicMock()
daq._devs = MagicMock()
daq._set_state = MagicMock()
daq._handle_operation_error = MagicMock()
daq.save_screenshot_db = MagicMock()
daq.sync_current_sample_from_tell = MagicMock(return_value=previous_sample)
daq._create_mounting_service = MagicMock()
return daq
@pytest.fixture
def bare_daq():
daq = object.__new__(AareDAQ)
daq._cfg = SimpleNamespace()
daq._devs = SimpleNamespace()
return daq
def test_create_mounting_service_builds_expected_context(mock_logger, bare_daq):
service = bare_daq._create_mounting_service()
assert isinstance(service.ctx, MountingContext)
assert service.ctx.deps.cfg is bare_daq._cfg
assert service.ctx.deps.devs is bare_daq._devs
assert service.ctx.settings.mount_position == ABR_POS_MOUNT
def test_was_previous_sample_unmounted_since_prefers_tell_phase_confirmation():
previous_sample = _make_sample(1, "old")
daq = _make_daq(previous_sample)
started_at = datetime.now(UTC) - timedelta(seconds=5)
daq._safe_tell_state = MagicMock(
return_value=TellStateModel(
activity=TellActivityEnum.MOUNTING,
operation="mount",
phase=TellPhaseEnum.PICKING_NEW_SAMPLE,
last_update_ts=datetime.now(UTC).isoformat(),
last_event_class="Motion Sync",
last_event_value="Sample get on Puck",
)
)
daq._get_tell_events_from_redis = MagicMock(return_value=[])
assert daq._was_previous_sample_unmounted_since(started_at) is True
def test_execute_mount_and_prepare_success_uses_mounting_result_fields():
previous_sample = _make_sample(1, "old")
target_sample = _make_sample(2, "new")
daq = _make_daq(previous_sample)
service = MagicMock()
service.execute.return_value = MountingResult(
success=True,
mounted_sample=target_sample,
previous_sample=previous_sample,
did_unmount_previous=True,
)
daq._create_mounting_service.return_value = service
result = daq._execute_mount_and_prepare(target_sample)
assert result is True
send_calls = daq._aare.send_sample_event.call_args_list
assert send_calls[0].args[0] == previous_sample.db_id
assert send_calls[0].args[1] == SampleEventType.UNMOUNTING
assert send_calls[1].args[0] == target_sample.db_id
assert send_calls[1].args[1] == SampleEventType.MOUNTING
assert send_calls[2].args[0] == previous_sample.db_id
assert send_calls[2].args[1] == SampleEventType.UNMOUNTED
assert send_calls[3].args[0] == target_sample.db_id
assert send_calls[3].args[1] == SampleEventType.MOUNTED
daq.save_screenshot_db.assert_called_once_with(
target_sample.db_id, f"{target_sample.db_id}_mounted"
)
assert daq._set_state.call_args_list[0].args[0] == BeamlineStateEnum.RobotSampleExchange
assert daq._set_state.call_args_list[-1].args[0] == BeamlineStateEnum.SampleAlignment
def test_execute_mount_and_prepare_marks_previous_sample_unmounted_when_mount_fails_after_auto_unmount():
previous_sample = _make_sample(1, "old")
target_sample = _make_sample(2, "new")
daq = _make_daq(previous_sample)
service = MagicMock()
service.execute.return_value = MountingResult(
success=False,
error=RuntimeError("mount failed after unmount"),
comment="mount failed after unmount",
previous_sample=previous_sample,
)
daq._create_mounting_service.return_value = service
daq._was_previous_sample_unmounted_since = MagicMock(return_value=True)
result = daq._execute_mount_and_prepare(target_sample)
assert result is False
assert daq._cfg.current_sample is None
send_calls = daq._aare.send_sample_event.call_args_list
assert send_calls[0].args[0] == previous_sample.db_id
assert send_calls[0].args[1] == SampleEventType.UNMOUNTING
assert send_calls[1].args[0] == target_sample.db_id
assert send_calls[1].args[1] == SampleEventType.MOUNTING
assert send_calls[2].args[0] == previous_sample.db_id
assert send_calls[2].args[1] == SampleEventType.UNMOUNTED
assert send_calls[2].kwargs["comment"] == "Auto-unmount succeeded before mount failed"
daq._handle_operation_error.assert_called_once()
assert daq._handle_operation_error.call_args.kwargs["operation"] == DAQOperation.MOUNT
assert daq._handle_operation_error.call_args.kwargs["sample"] == target_sample
assert daq._handle_operation_error.call_args.kwargs["event_type"] == SampleEventType.MOUNTFAILED
assert daq._set_state.call_args_list[-1].args[0] == BeamlineStateEnum.SampleAlignment
def test_execute_mount_and_prepare_does_not_mark_previous_sample_unmounted_when_not_confirmed():
previous_sample = _make_sample(1, "old")
target_sample = _make_sample(2, "new")
daq = _make_daq(previous_sample)
service = MagicMock()
service.execute.return_value = MountingResult(
success=False,
error=RuntimeError("mount failed before unmount confirmation"),
comment="mount failed before unmount confirmation",
previous_sample=previous_sample,
)
daq._create_mounting_service.return_value = service
daq._was_previous_sample_unmounted_since = MagicMock(return_value=False)
result = daq._execute_mount_and_prepare(target_sample)
assert result is False
assert daq._cfg.current_sample == previous_sample
send_calls = daq._aare.send_sample_event.call_args_list
assert len(send_calls) == 2
assert send_calls[0].args[0] == previous_sample.db_id
assert send_calls[0].args[1] == SampleEventType.UNMOUNTING
assert send_calls[1].args[0] == target_sample.db_id
assert send_calls[1].args[1] == SampleEventType.MOUNTING
daq._handle_operation_error.assert_called_once()
assert daq._handle_operation_error.call_args.kwargs["event_type"] == SampleEventType.MOUNTFAILED
def test_execute_mount_and_prepare_unmount_success_uses_unmount_operation():
previous_sample = _make_sample(1, "old")
daq = _make_daq(previous_sample)
service = MagicMock()
service.execute.return_value = MountingResult(
success=True,
mounted_sample=None,
previous_sample=previous_sample,
did_unmount_previous=True,
)
daq._create_mounting_service.return_value = service
result = daq._execute_mount_and_prepare(None)
assert result is True
send_calls = daq._aare.send_sample_event.call_args_list
assert len(send_calls) == 2
assert send_calls[0].args[0] == previous_sample.db_id
assert send_calls[0].args[1] == SampleEventType.UNMOUNTING
assert send_calls[1].args[0] == previous_sample.db_id
assert send_calls[1].args[1] == SampleEventType.UNMOUNTED
daq.save_screenshot_db.assert_not_called()
assert daq._set_state.call_args_list[0].args[0] == BeamlineStateEnum.RobotSampleExchange
assert daq._set_state.call_args_list[-1].args[0] == BeamlineStateEnum.SampleAlignment
def test_raise_if_critical_jfjoch_detector_error_preserves_jfjoch_exception_family():
daq = object.__new__(AareDAQ)
original = JFJochCommunicationError(
"DAQ state error: detector must be idle",
operation="measure",
endpoint="/measurement/start",
base_url="http://detector",
status_code=500,
)
with pytest.raises(JFJochCommunicationError) as exc_info:
daq._raise_if_critical_jfjoch_detector_error(original, command="wait_till_done")
exc = exc_info.value
assert "Critical detector error while running JFJoch command 'wait_till_done'" in str(exc)
assert exc.operation == "measure"
assert exc.endpoint == "/measurement/start"
assert exc.base_url == "http://detector"
assert exc.status_code == 500
assert exc.critical is True
assert exc.__cause__ is original
def test_raise_if_critical_bec_error_preserves_bec_exception_family():
daq = object.__new__(AareDAQ)
original_exc = RuntimeError("planner disconnected")
original = BECCommunicationError(
"BEC operation failed",
operation="planner.move_to:data_collection",
endpoint="/bec",
base_url="redis://bec",
exception=original_exc,
)
with pytest.raises(BECCommunicationError) as exc_info:
daq._raise_if_critical_bec_error(original, command="planner.move_to:data_collection")
exc = exc_info.value
assert "Critical BEC error while running 'planner.move_to:data_collection'" in str(exc)
assert exc.operation == "planner.move_to:data_collection"
assert exc.endpoint == "/bec"
assert exc.base_url == "redis://bec"
assert exc.exception is original_exc
assert exc.critical is True
assert exc.__cause__ is original
def test_raise_if_critical_jfjoch_detector_error_ignores_non_critical_jfjoch_errors():
daq = object.__new__(AareDAQ)
original = JFJochCommunicationError(
"temporary detector hiccup",
operation="measure",
endpoint="/measurement/start",
status_code=503,
)
assert daq._raise_if_critical_jfjoch_detector_error(original, command="wait_till_done") is None
def test_create_loop_centering_service_uses_shared_screenshot_service():
daq = object.__new__(AareDAQ)
daq._cfg = SimpleNamespace()
daq._devs = SimpleNamespace()
daq._mlbox = SimpleNamespace(predict_all_best=MagicMock())
daq._screenshot_service = MagicMock(spec=ScreenshotService)
daq._append_smargon_trace = MagicMock()
type(daq).sample = property(lambda self: None)
type(daq).sample_geometry = property(lambda self: SimpleNamespace())
type(daq).status = property(lambda self: SimpleNamespace())
service = daq._create_loop_centering_service()
assert service.ctx.services.screenshots is daq._screenshot_service
def test_create_raster_service_uses_shared_screenshot_service():
daq = object.__new__(AareDAQ)
daq._cfg = SimpleNamespace()
daq._devs = SimpleNamespace()
daq._mlbox = SimpleNamespace()
daq._jfjoch = SimpleNamespace()
daq._aare = SimpleNamespace()
daq._set_state = MagicMock()
daq._screenshot_service = MagicMock(spec=ScreenshotService)
type(daq).sample = property(lambda self: None)
type(daq).sample_geometry = property(lambda self: SimpleNamespace())
type(daq).status = property(lambda self: SimpleNamespace())
service = daq._create_raster_service()
assert service.ctx.services.screenshots is daq._screenshot_service