from unittest.mock import MagicMock, patch import numpy as np import pytest from aarecommon.math.coordinate import Coordinate, SmargonCoordinate from aarecommon.math.sample_geometry import SampleGeometryModel from aarecommon.models.models import DewarAddress, PuckLoadedInfo, SampleShortInfo from aarecommon.models.raster_grid import CenterOfMassModel, RasterGridRequest from aarecommon.models.rotation_scan import RotationScanRequest from jfjoch_client.models import ScanResult from aare.daq.aaredb import AareWrapper @pytest.fixture def mock_bl(): bl = MagicMock() bl.value = "X10SA" bl.name = "X10SA" return bl @pytest.fixture def sample_info(): return SampleShortInfo( db_id=123, puck_name="P1", dewar_name="D1", sample_name="S1", run_number=1, user="testuser", pin=5, location=DewarAddress(segment="A", pos=1), ) @pytest.fixture def daq_status(mock_bl): status = MagicMock() status.bl = mock_bl status.diffraction.detector_description = "EIGER2 16M" status.diffraction.detector_serial_number = "1234" status.diffraction.dtz_mm = 200.0 status.diffraction.max_resolution_angstrom = 1.5 status.diffraction.beam_center_pxl = [1000.0, 1000.0] status.diffraction.pixel_size_mm = 0.075 status.diffraction.energy_keV = 12.658 status.diffraction.wavelength_angstrom = 0.979 status.bl.ring_current_mA = 400.0 status.bl.transmission = 1.0 status.bl.flux_ph_s = 1e12 status.bl.cryojet_K = 100.0 status.geom.beam_size_mm = Coordinate(x=0.01, y=0.01) status.geom.smargon = SmargonCoordinate(sh_mm=Coordinate(x=0, y=0, z=0), phi_deg=0, chi_deg=0) return status @pytest.fixture def geom_model(): return SampleGeometryModel( beam_location_pxl=Coordinate(x=1000, y=1000), pixel_in_mm=0.001, aerotech=Coordinate(x=0, y=0), aerotech_meas=Coordinate(x=0, y=0), smargon=SmargonCoordinate(sh_mm=Coordinate(x=0, y=0, z=0), phi_deg=0, chi_deg=0), omega_deg=0, beam_size_mm=Coordinate(x=0.01, y=0.01), ) @patch("aareDB.ApiClient") @patch("aareDB.TellsRunnerApi") @patch("aareDB.SamplesRunnerApi") @patch("aareDB.ProcessingsRunnerApi") @patch("aareDB.GridscanRunnerApi") def test_aare_wrapper_init(mock_grid, mock_proc, mock_sample, mock_tell, mock_api, mock_bl): wrapper = AareWrapper(bl=mock_bl) mock_api.assert_called_once() assert wrapper._bl == mock_bl @patch("aareDB.ApiClient") @patch("aareDB.TellsRunnerApi") @patch("aareDB.SamplesRunnerApi") @patch("aareDB.ProcessingsRunnerApi") @patch("aareDB.GridscanRunnerApi") def test_set_pucks_beamline(mock_grid, mock_proc, mock_sample, mock_tell, mock_api, mock_bl): wrapper = AareWrapper(bl=mock_bl) pucks = [PuckLoadedInfo(puck_name="P1", location=DewarAddress(segment="A", pos=1))] wrapper.set_pucks_beamline(pucks) mock_tell.return_value.set_tell_positions.assert_called_once() @patch("aareDB.ApiClient") @patch("aareDB.SamplesRunnerApi") def test_create_manual_sample(mock_sample, mock_api, mock_bl, sample_info): wrapper = AareWrapper(bl=mock_bl) mock_sample.return_value.insert_sample.return_value.id = 456 wrapper.create_manual_sample(sample_info) assert sample_info.db_id == 456 mock_sample.return_value.insert_sample.assert_called_once() @patch("aareDB.ApiClient") @patch("aareDB.SamplesRunnerApi") def test_send_sample_event(mock_sample, mock_api, mock_bl, sample_info): from aareDB import SampleEventType wrapper = AareWrapper(bl=mock_bl) wrapper.send_sample_event(sample_info.db_id, SampleEventType.MOUNTED, "Test comment") mock_sample.return_value.create_sample_event.assert_called_once() # Test None sample mock_sample.return_value.create_sample_event.reset_mock() wrapper.send_sample_event(None, SampleEventType.MOUNTED) mock_sample.return_value.create_sample_event.assert_not_called() # Test invalid ID sample_info.db_id = -1 wrapper.send_sample_event(sample_info.db_id, SampleEventType.MOUNTED) mock_sample.return_value.create_sample_event.assert_not_called() @patch("aareDB.ApiClient") @patch("aareDB.SamplesRunnerApi") def test_create_manual_sample_error(mock_sample, mock_api, mock_bl, sample_info, caplog): import logging caplog.set_level(logging.ERROR) wrapper = AareWrapper(bl=mock_bl) mock_sample.return_value.insert_sample.side_effect = Exception("DB Error") wrapper.create_manual_sample(sample_info) assert "Error inserting sample" in caplog.text # the cause now arrives via the logged traceback rather than the message assert "DB Error" in caplog.text @patch("aareDB.ApiClient") @patch("aareDB.SamplesRunnerApi") def test_send_sample_event_error(mock_sample, mock_api, mock_bl, sample_info, caplog): import logging from aareDB import SampleEventType caplog.set_level(logging.ERROR) wrapper = AareWrapper(bl=mock_bl) mock_sample.return_value.create_sample_event.side_effect = Exception("Event Error") wrapper.send_sample_event(sample_info.db_id, SampleEventType.MOUNTED) assert "Error sending sample event" in caplog.text assert "MOUNTED" in caplog.text assert "Event Error" in caplog.text @patch("aareDB.ApiClient") @patch("requests.post") def test_upload_image(mock_post, mock_api, mock_bl): wrapper = AareWrapper(bl=mock_bl) mock_post.return_value.status_code = 200 img = np.zeros((100, 100, 3), dtype=np.uint8) wrapper.upload_image(123, "test_img", img, "test message") mock_post.assert_called_once() assert "test_img.jpg" in str(mock_post.call_args) @patch("aareDB.ApiClient") @patch("requests.post") def test_upload_jpg(mock_post, mock_api, mock_bl): wrapper = AareWrapper(bl=mock_bl) mock_post.return_value.status_code = 200 wrapper.upload_jpg(123, "test_img", b"fake jpeg bytes") mock_post.assert_called_once() @patch("aareDB.ApiClient") @patch("aareDB.SamplesRunnerApi") def test_create_rotation_run(mock_sample, mock_api, mock_bl, sample_info, daq_status): wrapper = AareWrapper(bl=mock_bl) # Standard rotation req = RotationScanRequest( exp_time_s=0.1, incr_omega_deg=0.1, steps=100, file_prefix="test_prefix", dtz=200.0 ) wrapper.create_rotation_run(sample_info, req, daq_status) mock_sample.return_value.create_experiment_parameters_for_sample.assert_called_once() # None sample - should return early mock_sample.return_value.create_experiment_parameters_for_sample.reset_mock() wrapper.create_rotation_run(None, req, daq_status) mock_sample.return_value.create_experiment_parameters_for_sample.assert_not_called() @patch("aareDB.ApiClient") @patch("aareDB.SamplesRunnerApi") def test_create_rotation_run_screening(mock_sample, mock_api, mock_bl, sample_info, daq_status): wrapper = AareWrapper(bl=mock_bl) req = RotationScanRequest( exp_time_s=0.1, incr_omega_deg=0.1, steps=100, file_prefix="test_prefix", screening=True, wedge_omega_deg=5.0, dtz=200.0, ) wrapper.create_rotation_run(sample_info, req, daq_status) mock_sample.return_value.create_experiment_parameters_for_sample.assert_called_once() @patch("aareDB.ApiClient") @patch("aareDB.SamplesRunnerApi") def test_create_rotation_run_error(mock_sample, mock_api, mock_bl, sample_info, daq_status, caplog): wrapper = AareWrapper(bl=mock_bl) req = RotationScanRequest( exp_time_s=0.1, incr_omega_deg=0.1, steps=100, file_prefix="test_prefix", dtz=200.0 ) mock_sample.return_value.create_experiment_parameters_for_sample.side_effect = Exception( "Rotation error" ) wrapper.create_rotation_run(sample_info, req, daq_status) assert "Rotation error" in caplog.text @patch("aareDB.ApiClient") @patch("aareDB.SamplesRunnerApi") def test_create_gridscan_run(mock_sample, mock_api, mock_bl, sample_info, daq_status): wrapper = AareWrapper(bl=mock_bl) req = RasterGridRequest( exp_time_s=0.1, n_x=10, n_y=10, grid_size_mm=Coordinate(x=0.01, y=0.01), smargon_top_left=None, file_prefix="test_grid_prefix", dtz=200.0, ) wrapper.create_gridscan_run(sample_info, req, daq_status) mock_sample.return_value.create_experiment_parameters_for_sample.assert_called_once() # None sample mock_sample.return_value.create_experiment_parameters_for_sample.reset_mock() wrapper.create_gridscan_run(None, req, daq_status) mock_sample.return_value.create_experiment_parameters_for_sample.assert_not_called() @patch("aareDB.ApiClient") @patch("aareDB.SamplesRunnerApi") def test_create_gridscan_run_error(mock_sample, mock_api, mock_bl, sample_info, daq_status, caplog): wrapper = AareWrapper(bl=mock_bl) req = RasterGridRequest( exp_time_s=0.1, n_x=10, n_y=10, grid_size_mm=Coordinate(x=0.01, y=0.01), smargon_top_left=None, file_prefix="test_grid_prefix", dtz=200.0, ) mock_sample.return_value.create_experiment_parameters_for_sample.side_effect = Exception( "Grid error" ) wrapper.create_gridscan_run(sample_info, req, daq_status) assert "Grid error" in caplog.text @patch("aareDB.ApiClient") @patch("requests.post") def test_ingest_gridscan(mock_post, mock_api, mock_bl, sample_info, geom_model): wrapper = AareWrapper(bl=mock_bl) mock_post.return_value.status_code = 200 raster_result = MagicMock(spec=ScanResult) raster_request = RasterGridRequest( exp_time_s=0.1, n_x=10, n_y=10, grid_size_mm=Coordinate(x=0.01, y=0.01), smargon_top_left=None, ) com = CenterOfMassModel(n_x=5.0, n_y=5.0) wrapper.ingest_gridscan( sample_info, raster_result, raster_request, geom_model, com, (500.0, 500.0) ) mock_post.assert_called_once() # None sample mock_post.reset_mock() wrapper.ingest_gridscan(None, raster_result, raster_request, geom_model, com, (500.0, 500.0)) mock_post.assert_not_called() @patch("aareDB.ApiClient") @patch("requests.post") def test_ingest_scan(mock_post, mock_api, mock_bl, sample_info, geom_model): wrapper = AareWrapper(bl=mock_bl) mock_post.return_value.status_code = 200 result = MagicMock(spec=ScanResult) wrapper.ingest_scan(sample_info, result, geom_model, (500.0, 500.0)) mock_post.assert_called_once() # None sample mock_post.reset_mock() wrapper.ingest_scan(None, result, geom_model, (500.0, 500.0)) mock_post.assert_not_called() @patch("aareDB.ApiClient") def test_format_gridscan_payload_no_com(mock_api, mock_bl, sample_info, geom_model): wrapper = AareWrapper(bl=mock_bl) raster_result = MagicMock(spec=ScanResult) raster_request = RasterGridRequest( exp_time_s=0.1, n_x=10, n_y=10, grid_size_mm=Coordinate(x=0.01, y=0.01), smargon_top_left=None, ) payload = wrapper.format_gridscan_payload( sample_info, raster_result, raster_request, geom_model, None, (500.0, 500.0) ) assert payload.center_pxl is None assert payload.sample_id == sample_info.db_id @patch("aareDB.ApiClient") def test_format_gridscan_payload_with_top_left(mock_api, mock_bl, sample_info, geom_model): wrapper = AareWrapper(bl=mock_bl) raster_result = MagicMock(spec=ScanResult) top_left = SmargonCoordinate(sh_mm=Coordinate(x=0.1, y=0.1, z=0.1), phi_deg=0, chi_deg=0) raster_request = RasterGridRequest( exp_time_s=0.1, n_x=10, n_y=10, grid_size_mm=Coordinate(x=0.01, y=0.01), smargon_top_left=top_left, ) payload = wrapper.format_gridscan_payload( sample_info, raster_result, raster_request, geom_model, None, (500.0, 500.0) ) assert payload.sample_id == sample_info.db_id @patch("aareDB.ApiClient") def test_ingest_gridscan_payload_none(mock_api, mock_bl, sample_info, geom_model): wrapper = AareWrapper(bl=mock_bl) # We want format_gridscan_payload to return None # Looking at aaredb.py, it returns None if an exception occurs and it's caught # But wait, it raises 'e' after logging. Oh, I see. # Lines 351-352 in aaredb.py: if payload is None: return # This happens if format_gridscan_payload returns None. with patch.object(AareWrapper, "format_gridscan_payload", return_value=None): wrapper.ingest_gridscan(sample_info, None, None, geom_model, None, (0, 0)) # type: ignore # intended with patch.object(AareWrapper, "format_scan_payload", return_value=None): wrapper.ingest_scan(sample_info, None, geom_model, (0, 0)) # type: ignore # intended @patch("aareDB.ApiClient") def test_format_gridscan_payload_error(mock_api, mock_bl, sample_info, caplog): wrapper = AareWrapper(bl=mock_bl) # Passing None for geom_model should trigger an error in smargon_to_picture with pytest.raises(AttributeError): wrapper.format_gridscan_payload(sample_info, None, None, None, None, (0, 0)) # type: ignore # intended assert "NoneType" in caplog.text @patch("aareDB.ApiClient") def test_format_scan_payload_error(mock_api, mock_bl, sample_info, caplog): wrapper = AareWrapper(bl=mock_bl) with pytest.raises(AttributeError): # Passing None for geom should trigger error when accessing geom.beam_size_mm wrapper.format_scan_payload(sample_info, None, None, (0, 0)) # type: ignore # intended assert "NoneType" in caplog.text