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