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

244 lines
9.5 KiB
Python

from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from aareDB import SampleEventType
from aare.common.models import DewarAddress, SampleShortInfo, DAQOperation
from aare.common.tell_models import TellActivityEnum, TellPhaseEnum, TellStateModel
from aare.daq.config import ABR_POS_MOUNT, BeamlineStateEnum
from aare.daq.daq import AareDAQ
from aare.daq.operations.mounting.models import MountingResult, MountingContext
from aare.daq.operations.mounting.service import MountingService
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._AareDAQ__cfg = SimpleNamespace(current_sample=previous_sample)
daq._AareDAQ__aare = MagicMock()
daq._AareDAQ__devs = MagicMock()
daq._AareDAQ__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._AareDAQ__cfg = SimpleNamespace()
daq._AareDAQ__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._AareDAQ__cfg
assert service.ctx.deps.devs is bare_daq._AareDAQ__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(timezone.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(timezone.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._AareDAQ__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._AareDAQ__set_state.call_args_list[0].args[0] == BeamlineStateEnum.RobotSampleExchange
assert daq._AareDAQ__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._AareDAQ__cfg.current_sample is None
send_calls = daq._AareDAQ__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._AareDAQ__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._AareDAQ__cfg.current_sample == previous_sample
send_calls = daq._AareDAQ__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._AareDAQ__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._AareDAQ__set_state.call_args_list[0].args[0] == BeamlineStateEnum.RobotSampleExchange
assert daq._AareDAQ__set_state.call_args_list[-1].args[0] == BeamlineStateEnum.SampleAlignment
def test_create_loop_centering_service_uses_shared_screenshot_service():
daq = object.__new__(AareDAQ)
daq._AareDAQ__cfg = SimpleNamespace()
daq._AareDAQ__devs = SimpleNamespace()
daq._AareDAQ__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._AareDAQ__cfg = SimpleNamespace()
daq._AareDAQ__devs = SimpleNamespace()
daq._AareDAQ__mlbox = SimpleNamespace()
daq._AareDAQ__jfjoch = SimpleNamespace()
daq._AareDAQ__aare = SimpleNamespace()
daq._AareDAQ__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