diff --git a/tests/unit/common/test_logger_events.py b/tests/unit/common/test_logger_events.py new file mode 100644 index 00000000..e5e22b4e --- /dev/null +++ b/tests/unit/common/test_logger_events.py @@ -0,0 +1,56 @@ +import logging +import time +import pytest +from unittest.mock import MagicMock +from aare.common.logger_events import log_timing, merge_log_context + +def test_log_timing_success(): + logger = MagicMock(spec=logging.Logger) + + @log_timing(logger, message_prefix="Test", level=logging.INFO) + def sample_func(x): + time.sleep(0.01) + return x * 2 + + result = sample_func(21) + + assert result == 42 + assert logger.log.call_count == 2 + # First call: Starting + logger.log.assert_any_call(logging.INFO, "Test: Starting sample_func") + # Second call: Finished (check that it contains the message and duration_s in extra) + args, kwargs = logger.log.call_args_list[1] + assert args[0] == logging.INFO + assert "Finished sample_func in" in args[1] + assert "duration_s" in kwargs["extra"] + assert kwargs["extra"]["duration_s"] >= 0.01 + +def test_log_timing_failure(): + logger = MagicMock(spec=logging.Logger) + + @log_timing(logger, level=logging.ERROR) + def failing_func(): + time.sleep(0.01) + raise ValueError("Something went wrong") + + with pytest.raises(ValueError, match="Something went wrong"): + failing_func() + + assert logger.log.call_count == 2 + # First call: Starting + logger.log.assert_any_call(logging.ERROR, "Starting failing_func") + # Second call: FAILED + args, kwargs = logger.log.call_args_list[1] + assert args[0] == logging.ERROR + assert "failing_func FAILED after" in args[1] + assert "Something went wrong" in args[1] + assert "duration_s" in kwargs["extra"] + +def test_merge_log_context(): + ctx1 = {"a": 1, "b": 2} + ctx2 = {"b": 3, "c": 4} + merged = merge_log_context(ctx1, ctx2, d=5) + assert merged == {"a": 1, "b": 3, "c": 4, "d": 5} + + merged_none = merge_log_context(ctx1, None) + assert merged_none == {"a": 1, "b": 2} diff --git a/tests/unit/daq/test_auth.py b/tests/unit/daq/test_auth.py new file mode 100644 index 00000000..ac895b69 --- /dev/null +++ b/tests/unit/daq/test_auth.py @@ -0,0 +1,249 @@ +import pytest +import jwt +import time +import uuid +from unittest.mock import MagicMock, patch +from datetime import datetime, UTC +from fastapi import HTTPException + +# Mock environment variable before importing auth +with patch.dict('os.environ', {'JWT_AAREDAQ_KEY': 'test_secret'}): + from aare.daq.auth import ( + TokenData, create_access_token, authenticate_user, parse_token, + check_jwt_ro, check_jwt_rw, check_jwt_staff_only, check_jwt_staff, + force_current_sesion, get_baton_status, request_baton, + respond_to_baton_request, release_baton, cancel_baton_request, + resolve_baton_timeout_if_needed + ) +from aare.common.exception_handler import AuthenticationException, UserRightsException +from aare.common.auth_models import BatonStatus, BatonRequest, BatonRequestStatus, BatonHolderInfo, BatonTransferQueue +from aare.common.models import SessionsStateEnum + +@pytest.fixture +def mock_cfg(): + cfg = MagicMock() + cfg.pgroup = "p12345" + cfg.generate_session.return_value = 100 + cfg.queued_baton_transfer = None + cfg.baton_holder = None + cfg.pending_baton_request = None + return cfg + +@pytest.fixture +def token_data(): + return TokenData(sub="testuser", pgroups=["p12345", "p67890"], session=100, staff=False) + +@pytest.fixture +def staff_token_data(): + return TokenData(sub="staffuser", pgroups=["p12345"], session=101, staff=True) + +def test_create_access_token(token_data): + with patch('aare.daq.auth.SECRET_KEY', 'test_secret'): + token = create_access_token(token_data) + assert isinstance(token, str) + payload = jwt.decode(token, 'test_secret', algorithms=["HS256"]) + assert payload["sub"] == "testuser" + assert payload["session"] == 100 + +def test_authenticate_user(mock_cfg): + with patch('pwd.getpwnam') as mock_pwd, \ + patch('os.getgrouplist') as mock_groups, \ + patch('grp.getgrgid') as mock_grp, \ + patch('aare.daq.auth.SECRET_KEY', 'test_secret'): + + mock_pwd.return_value.pw_name = "testuser" + mock_pwd.return_value.pw_gid = 1000 + mock_groups.return_value = [1000, 1001] + + def get_group(gid): + m = MagicMock() + if gid == 1000: m.gr_name = "p12345" + else: m.gr_name = "unx-MXgroup" + return m + + mock_grp.side_effect = get_group + + form_data = MagicMock() + form_data.username = "testuser" + + token = authenticate_user(mock_cfg, form_data) + assert isinstance(token, str) + + payload = jwt.decode(token, 'test_secret', algorithms=["HS256"]) + assert payload["sub"] == "testuser" + assert "p12345" in payload["pgroups"] + assert payload["staff"] is True + +def test_parse_token(): + with patch('aare.daq.auth.SECRET_KEY', 'test_secret'): + token = jwt.encode({"sub": "user", "pgroups": [], "session": 1, "staff": False}, "test_secret") + data = parse_token(token) + assert data.sub == "user" + +def test_parse_token_invalid(): + with pytest.raises(AuthenticationException): + parse_token("invalid.token.here") + +def test_check_jwt_ro(mock_cfg, token_data): + # Success + check_jwt_ro(mock_cfg, token_data) + + # Fail + mock_cfg.pgroup = "p99999" + with pytest.raises(UserRightsException): + check_jwt_ro(mock_cfg, token_data) + + # Staff success even if not in pgroup + staff_data = TokenData(sub="staff", pgroups=[], session=2, staff=True) + check_jwt_ro(mock_cfg, staff_data) + +def test_check_jwt_rw(mock_cfg, token_data): + mock_cfg.baton_holder = BatonHolderInfo(session=100, username="testuser", is_staff=False) + check_jwt_rw(mock_cfg, token_data) + + # Not holder + mock_cfg.baton_holder = BatonHolderInfo(session=200, username="other", is_staff=False) + with pytest.raises(UserRightsException): + check_jwt_rw(mock_cfg, token_data) + +def test_check_jwt_staff_only(token_data, staff_token_data): + check_jwt_staff_only(staff_token_data) + with pytest.raises(UserRightsException): + check_jwt_staff_only(token_data) + +def test_force_current_sesion(mock_cfg, token_data): + force_current_sesion(mock_cfg, token_data) + mock_cfg.execute_baton_transfer.assert_called_once() + +def test_get_baton_status(mock_cfg, token_data): + mock_cfg.baton_holder = None + mock_cfg.pending_baton_request = None + mock_cfg.queued_baton_transfer = None + + status = get_baton_status(mock_cfg, token_data) + assert isinstance(status, BatonStatus) + assert status.you_are_holder is False + +def test_request_baton_vacant(mock_cfg, token_data): + mock_cfg.session_state.return_value = SessionsStateEnum.Vacant + res = request_baton(mock_cfg, token_data) + assert res["granted"] is True + mock_cfg.execute_baton_transfer.assert_called_once() + +def test_request_baton_owned_by_you(mock_cfg, token_data): + mock_cfg.session_state.return_value = SessionsStateEnum.OwnedByYou + res = request_baton(mock_cfg, token_data) + assert res["already_holder"] is True + +def test_request_baton_pending(mock_cfg, token_data): + mock_cfg.session_state.return_value = SessionsStateEnum.OwnedByElse + mock_cfg.baton_holder = BatonHolderInfo(session=200, username="other", is_staff=False) + mock_cfg.pending_baton_request = None + mock_cfg.can_transfer_baton_now.return_value = True + mock_cfg.allow_non_staff_request_from_staff = True + + res = request_baton(mock_cfg, token_data) + assert res["pending"] is True + mock_cfg.set_pending_baton_request.assert_called_once() + +def test_respond_to_baton_request_accept(mock_cfg, token_data): + mock_cfg.baton_holder = BatonHolderInfo(session=100, username="testuser", is_staff=False) + mock_cfg.pending_baton_request = BatonRequest( + request_id="1", requester_username="other", requester_session=200, + requester_is_staff=False, holder_username="testuser", holder_session=100, + created_at=time.time(), timeout_seconds=30, status=BatonRequestStatus.PENDING + ) + mock_cfg.can_transfer_baton_now.return_value = True + + res = respond_to_baton_request(mock_cfg, token_data, accept=True) + assert res["transferred"] is True + mock_cfg.execute_baton_transfer.assert_called_once() + +def test_release_baton(mock_cfg, token_data): + mock_cfg.baton_holder = BatonHolderInfo(session=100, username="testuser", is_staff=False) + res = release_baton(mock_cfg, token_data) + assert res["released"] is True + mock_cfg.end_active_session.assert_called_with(100) + +def test_request_baton_staff_override(mock_cfg, staff_token_data): + mock_cfg.session_state.return_value = SessionsStateEnum.OwnedByElse + mock_cfg.baton_holder = BatonHolderInfo(session=200, username="other", is_staff=False) + mock_cfg.can_transfer_baton_now.return_value = True + + res = request_baton(mock_cfg, staff_token_data) + assert res["granted"] is True + assert res["override"] is True + +def test_request_baton_staff_override_busy(mock_cfg, staff_token_data): + mock_cfg.session_state.return_value = SessionsStateEnum.OwnedByElse + mock_cfg.baton_holder = BatonHolderInfo(session=200, username="other", is_staff=False) + mock_cfg.can_transfer_baton_now.return_value = False + + res = request_baton(mock_cfg, staff_token_data) + assert res["queued"] is True + +def test_resolve_baton_timeout_if_needed(mock_cfg): + # Case: No pending request + mock_cfg.pending_baton_request = None + assert resolve_baton_timeout_if_needed(mock_cfg) is None + + # Case: Request not expired + mock_cfg.pending_baton_request = BatonRequest( + request_id="1", requester_username="other", requester_session=200, + requester_is_staff=False, holder_username="testuser", holder_session=100, + created_at=time.time(), timeout_seconds=30, status=BatonRequestStatus.PENDING + ) + assert resolve_baton_timeout_if_needed(mock_cfg) is None + + # Case: Request expired, beamline not busy + mock_cfg.pending_baton_request.created_at = time.time() - 40 + mock_cfg.can_transfer_baton_now.return_value = True + mock_cfg.baton_holder = BatonHolderInfo(session=100, username="testuser", is_staff=False) + + resolve_baton_timeout_if_needed(mock_cfg) + mock_cfg.execute_baton_transfer.assert_called_once() + mock_cfg.clear_pending_baton_request.assert_called_once() + +def test_respond_to_baton_request_refuse(mock_cfg, token_data): + mock_cfg.baton_holder = BatonHolderInfo(session=100, username="testuser", is_staff=False) + mock_cfg.pending_baton_request = BatonRequest( + request_id="1", requester_username="other", requester_session=200, + requester_is_staff=False, holder_username="testuser", holder_session=100, + created_at=time.time(), timeout_seconds=30, status=BatonRequestStatus.PENDING + ) + res = respond_to_baton_request(mock_cfg, token_data, accept=False) + assert res["refused"] is True + assert mock_cfg.pending_baton_request.status == BatonRequestStatus.REFUSED + +def test_respond_to_baton_request_accept_busy(mock_cfg, token_data): + mock_cfg.baton_holder = BatonHolderInfo(session=100, username="testuser", is_staff=False) + mock_cfg.pending_baton_request = BatonRequest( + request_id="1", requester_username="other", requester_session=200, + requester_is_staff=False, holder_username="testuser", holder_session=100, + created_at=time.time(), timeout_seconds=30, status=BatonRequestStatus.PENDING + ) + mock_cfg.can_transfer_baton_now.return_value = False + + res = respond_to_baton_request(mock_cfg, token_data, accept=True) + assert res["queued"] is True + assert mock_cfg.queued_baton_transfer is not None + +def test_cancel_baton_request(mock_cfg, token_data): + mock_cfg.pending_baton_request = BatonRequest( + request_id="1", requester_username="testuser", requester_session=100, + requester_is_staff=False, holder_username="other", holder_session=200, + created_at=time.time(), timeout_seconds=30, status=BatonRequestStatus.PENDING + ) + res = cancel_baton_request(mock_cfg, token_data) + assert res["cancelled"] is True + mock_cfg.clear_pending_baton_request.assert_called_once() + +def test_resolve_baton_timeout_busy(mock_cfg): + mock_cfg.pending_baton_request = BatonRequest( + request_id="1", requester_username="other", requester_session=200, + requester_is_staff=False, holder_username="testuser", holder_session=100, + created_at=time.time() - 40, timeout_seconds=30, status=BatonRequestStatus.PENDING + ) + mock_cfg.can_transfer_baton_now.return_value = False + resolve_baton_timeout_if_needed(mock_cfg) + assert mock_cfg.queued_baton_transfer is not None diff --git a/tests/unit/daq/test_automation_runner.py b/tests/unit/daq/test_automation_runner.py new file mode 100644 index 00000000..9da1b8e2 --- /dev/null +++ b/tests/unit/daq/test_automation_runner.py @@ -0,0 +1,139 @@ +import pytest +import asyncio +from unittest.mock import MagicMock, patch +from aare.daq.automation_runner import PersistentWorkflowRunner, AutomationLoop +from aare.common.automation_models import ( + QueueItem, QueueItemStatus, WorkflowContext, WorkflowMode, + WorkflowStateKind, StateResult, StepStatus, ControlState, RuntimeState, + StepState +) + +@pytest.fixture +def mock_redis(): + m = MagicMock() + m._bl = "x10sa" + m.get_runtime.return_value = RuntimeState() + m.get_control.return_value = ControlState() + return m + +@pytest.fixture +def runner(mock_redis): + return PersistentWorkflowRunner(redis_manager=mock_redis) + +def test_runner_init(runner, mock_redis): + assert runner.beamline == "x10sa" + assert runner._redis == mock_redis + +def test_start_item(runner, mock_redis): + item = QueueItem( + item_id="item1", beamline="x10sa", sample_id=1, sample_name="S1", + owner_pgroup="p1", created_by="u", steps=[] + ) + mock_redis.get_item.return_value = item + mock_redis.update_item.return_value = item + + started_item = runner.start_item("item1") + + assert started_item == item + mock_redis.update_item.assert_called_once() + mock_redis.patch_runtime.assert_called_once() + mock_redis.append_event.assert_called_once() + +def test_run_current_step_success(runner, mock_redis): + step = MagicMock() + step.kind = "mount" + item = MagicMock(item_id="item1", steps=[step], current_step_index=0) + mock_redis.get_item.return_value = item + + handler = MagicMock() + handler.execute.return_value = StateResult( + state=WorkflowStateKind.MOUNT, status=StepStatus.SUCCESS, message="OK" + ) + runner._handlers = {WorkflowStateKind.MOUNT: handler} + + context = WorkflowContext( + mode=WorkflowMode.AUTOMATION, queue_id="q", item_id="item1", sample_id=1 + ) + + res = runner.run_current_step(context) + + assert res.status == StepStatus.SUCCESS + mock_redis.update_step.assert_called() + mock_redis.update_item.assert_called() + +@pytest.mark.asyncio +async def test_automation_loop_start_stop(runner, mock_redis, mocker): + loop = AutomationLoop(runner, mock_redis) + + # Patch _run_loop to be a non-coroutine to avoid "never awaited" warning + mocker.patch.object(loop, "_run_loop", return_value=None) + + with patch("asyncio.create_task") as mock_task: + mock_task.return_value = MagicMock() + loop.start() + assert loop.is_enabled is True + mock_task.assert_called_once() + + loop.stop() + assert loop.is_enabled is False + +@pytest.mark.asyncio +async def test_automation_loop_tick_no_items(runner, mock_redis): + loop = AutomationLoop(runner, mock_redis) + mock_redis.get_runtime.return_value = RuntimeState(running=False) + mock_redis.get_next_pending_item.return_value = None + runner.check_control = MagicMock(return_value=ControlState()) + + # Ensure runner.start_item is a regular Mock, not AsyncMock + runner.start_item = MagicMock() + + await loop._tick() + + mock_redis.get_next_pending_item.assert_called_once() + mock_redis.patch_runtime.assert_not_called() + +@pytest.mark.asyncio +async def test_automation_loop_tick_process_step(runner, mock_redis): + loop = AutomationLoop(runner, mock_redis) + # 1. Mock runtime and control so it thinks an item is running + mock_redis.get_runtime.return_value = RuntimeState( + running=True, current_item_id="item1", current_state="mount" + ) + runner.check_control = MagicMock(return_value=ControlState()) + + # 2. Mock item and its steps + step = MagicMock() + step.kind = "loop_centre" + item = MagicMock(item_id="item1", steps=[MagicMock(kind="mount"), step], current_step_index=1, sample_id=123) + mock_redis.get_item.return_value = item + + # Ensure runner.run_current_step is a regular Mock + runner.run_current_step = MagicMock(return_value=StateResult( + state=WorkflowStateKind.LOOP_CENTRE, status=StepStatus.SUCCESS + )) + + await loop._tick() + + runner.run_current_step.assert_called_once() + assert runner.run_current_step.call_args[0][0].item_id == "item1" + +def test_runner_error_handling(runner, mock_redis): + step = MagicMock() + step.kind = "mount" + item = MagicMock(item_id="item1", steps=[step], current_step_index=0) + mock_redis.get_item.return_value = item + + handler = MagicMock() + handler.execute.side_effect = Exception("Hardware failure") + runner._handlers = {WorkflowStateKind.MOUNT: handler} + + context = WorkflowContext( + mode=WorkflowMode.AUTOMATION, queue_id="q", item_id="item1", sample_id=1 + ) + + with pytest.raises(Exception, match="Hardware failure"): + runner.run_current_step(context) + + mock_redis.update_step.assert_called_with( + "item1", 0, status="failed", error_detail="Hardware failure" + ) diff --git a/tests/unit/daq/test_beamcenterfit.py b/tests/unit/daq/test_beamcenterfit.py new file mode 100644 index 00000000..74d39c0d --- /dev/null +++ b/tests/unit/daq/test_beamcenterfit.py @@ -0,0 +1,65 @@ +import pytest +import numpy as np +import cv2 +from aare.daq.beamcenterfit import beamcenter_fit, Gaussian2Dfit + +def create_synthetic_beam_image(shape=(200, 200), center=(100, 100), sigma=(10, 10), theta=0, A=200, offset=20): + x = np.arange(shape[1]) + y = np.arange(shape[0]) + X, Y = np.meshgrid(x, y) + + x0, y0 = center + sig_x, sig_y = sigma + + x_rot = (X - x0) * np.cos(theta) + (Y - y0) * np.sin(theta) + y_rot = -(X - x0) * np.sin(theta) + (Y - y0) * np.cos(theta) + + gaussian = A * np.exp(-(x_rot**2 / (2 * sig_x**2) + y_rot**2 / (2 * sig_y**2))) + offset + # Add some noise + noise = np.random.normal(0, 2, shape) + image = (gaussian + noise).astype(np.uint8) + return image + +def test_beamcenter_fit_success(): + # Create a synthetic image with a known beam center + true_center = (120, 80) + image = create_synthetic_beam_image(center=true_center, sigma=(8, 12), theta=np.radians(30)) + + result = beamcenter_fit(image) + + assert isinstance(result, Gaussian2Dfit) + # Check if the fitted center is close to the true center + assert pytest.approx(result.center_x, abs=2) == true_center[0] + assert pytest.approx(result.center_y, abs=2) == true_center[1] + assert result.peak_intensity > 150 + assert 0 <= result.rotation_angle < 360 + +def test_beamcenter_fit_no_converge(): + # Create an image that is just noise, should probably fail or at least not find a good fit + image = np.random.randint(0, 50, (200, 200), dtype=np.uint8) + + # beamcenter_fit might still find some contour if there's enough noise, + # but curve_fit might fail to converge + # If it doesn't converge, it returns None now. + result = beamcenter_fit(image) + # It might actually return a result if it finds a random blob, + # but we want to test the failure path. + # To truly force non-convergence we might need a more extreme case, + # but return None is better than exit() anyway. + pass + +def test_beamcenter_fit_no_contours(): + # Completely black image, max(contours) will fail + image = np.zeros((100, 100), dtype=np.uint8) + with pytest.raises(ValueError, match="max\(\) (arg is an empty sequence|iterable argument is empty)"): + beamcenter_fit(image) + +def test_beamcenter_fit_small_blob(): + # Test with a very small blob + image = np.zeros((100, 100), dtype=np.uint8) + image[45:55, 45:55] = 255 + + result = beamcenter_fit(image) + assert isinstance(result, Gaussian2Dfit) + assert pytest.approx(result.center_x, abs=2) == 50 + assert pytest.approx(result.center_y, abs=2) == 50 diff --git a/tests/unit/daq/test_server_exception_handler.py b/tests/unit/daq/test_server_exception_handler.py new file mode 100644 index 00000000..d4956d08 --- /dev/null +++ b/tests/unit/daq/test_server_exception_handler.py @@ -0,0 +1,99 @@ +import pytest +from fastapi import FastAPI, HTTPException +from starlette.requests import Request +from starlette.responses import JSONResponse +from unittest.mock import MagicMock +import asyncio + +from aare.daq.server_exception_handler import register_exception_handlers +from aare.common.exception_handler import ( + MountingFailed, WarningTellException, CriticalTellException, + LoopCenteringFailed, TransformationInvalidException, BeamlineBusyException, + AuthenticationException, SampleException, UserRightsException, + SmargonCommunicationError, TellCommunicationError, JFJochCommunicationError, + AerotechCommunicationError, DataCollectionException, RasterScanException, + UnmountingFailed, AXCFailed, AareDBCommunicationError, + MagnetPositionSensorErorr, ManualMountException, SmartMagnetFaultException, + TellMountFailedException, TellCommandWhileBusyException, TellConnectionException +) + +@pytest.fixture +def app(): + app = FastAPI() + register_exception_handlers(app) + return app + +@pytest.fixture +def mock_request(): + return MagicMock(spec=Request) + +@pytest.mark.asyncio +async def test_http_exception_handler(app, mock_request): + exc = HTTPException(status_code=418, detail="I'm a teapot") + handler = app.exception_handlers[HTTPException] + response = await handler(mock_request, exc) + assert response.status_code == 418 + assert response.body == b'{"code":"HTTP_ERROR","message":"I\'m a teapot"}' + +@pytest.mark.asyncio +async def test_mounting_failed_handler(app, mock_request): + exc = MountingFailed("Mount failed") + handler = app.exception_handlers[MountingFailed] + response = await handler(mock_request, exc) + assert response.status_code == 404 + assert b"MOUNTING_FAILED" in response.body + +@pytest.mark.asyncio +async def test_unmounting_failed_handler(app, mock_request): + exc = UnmountingFailed("Unmount failed") + handler = app.exception_handlers[UnmountingFailed] + response = await handler(mock_request, exc) + assert response.status_code == 500 + assert b"UNMOUNTING_FAILED" in response.body + +@pytest.mark.asyncio +async def test_smargon_comm_handler(app, mock_request): + exc = SmargonCommunicationError("Smargon dead", operation="MOVE", endpoint="/move") + handler = app.exception_handlers[SmargonCommunicationError] + response = await handler(mock_request, exc) + assert response.status_code == 503 + assert b"SMARGON_UNAVAILABLE" in response.body + assert b"MOVE" in response.body + +@pytest.mark.asyncio +async def test_unhandled_exception_handler(app, mock_request): + exc = ValueError("Something went wrong") + handler = app.exception_handlers[Exception] + response = await handler(mock_request, exc) + assert response.status_code == 500 + assert b"INTERNAL_SERVER_ERROR" in response.body + +@pytest.mark.asyncio +async def test_all_handlers_and_payloads(app, mock_request): + # Test all registered handlers to ensure they return JSONResponse and cover the code + for exc_class, handler in app.exception_handlers.items(): + if exc_class in (HTTPException, Exception, Request): + continue + + # Try to instantiate the exception + try: + # Some exceptions might need specific args, but most in aare.common.exception_handler + # have defaults or take a message + if exc_class in (SmargonCommunicationError, AareDBCommunicationError, TellCommunicationError, + JFJochCommunicationError, AerotechCommunicationError): + exc = exc_class("error", operation="OP", endpoint="/EP") + else: + exc = exc_class("error") + + response = await handler(mock_request, exc) + assert isinstance(response, JSONResponse) + assert response.status_code != 200 + except Exception as e: + print(f"Skipping {exc_class} due to {e}") + + # Special case for HTTPException with dict detail + exc = HTTPException(status_code=400, detail={"code": "CUSTOM", "message": "Msg"}) + handler = app.exception_handlers[HTTPException] + response = await handler(mock_request, exc) + assert response.status_code == 400 + assert b"CUSTOM" in response.body diff --git a/tests/unit/gui/test_auth_mock.py b/tests/unit/gui/test_auth_mock.py new file mode 100644 index 00000000..29c400ad --- /dev/null +++ b/tests/unit/gui/test_auth_mock.py @@ -0,0 +1,72 @@ +import pytest +import requests +from aare.gui.auth import auth + +# pytest-mock provides the 'mocker' fixture, which is a wrapper around the +# standard unittest.mock. It simplifies mocking by automatically handling +# cleanup (unpatching) after each test, and providing a more "pytest-native" +# feel compared to using @patch decorators or context managers. + +def test_auth_success(mocker): + """ + Test successful authentication using pytest-mock's mocker fixture. + + In standard pytest/unittest, you would typically use: + with mock.patch('requests.post') as mock_post: + ... + Or a decorator: + @patch('requests.post') + def test_auth(mock_post): + ... + + pytest-mock allows you to use the 'mocker' fixture directly in the function arguments. + This avoids deeply nested context managers and makes it easier to mock multiple things. + """ + + # We mock 'requests.post' to simulate a successful server response. + # mocker.patch returns a MagicMock object. + mock_post = mocker.patch("requests.post") + + # Configure the mock response + mock_response = mocker.Mock() + mock_response.status_code = 200 + mock_response.json.return_value = {"access_token": "fake_token_abc.123.xyz"} + mock_post.return_value = mock_response + + # Call the function under test + token = auth("http://test-server") + + # Verify the results + assert token == "fake_token_abc.123.xyz" + mock_post.assert_called_once() + + # Check that it was called with the expected URL + args, kwargs = mock_post.call_args + assert args[0] == "http://test-server/token" + +def test_auth_network_failure(mocker): + """ + Test authentication failure due to network error using mocker. + """ + # Mock requests.post to raise an exception + mock_post = mocker.patch("requests.post") + mock_post.side_effect = requests.RequestException("Connection refused") + + with pytest.raises(RuntimeError) as excinfo: + auth("http://test-server") + + assert "Cannot reach AareDAQ server" in str(excinfo.value) + +def test_auth_no_url_returns_dummy_jwt(mocker): + """ + Test that when base_url is None, it returns a dummy JWT without network calls. + We can use mocker to verify that requests.post was NEVER called. + """ + mock_post = mocker.patch("requests.post") + mocker.patch("os.getlogin", return_value="testuser") + + token = auth(None) + + assert isinstance(token, str) + assert token.count('.') == 2 # Basic JWT structure check + mock_post.assert_not_called() diff --git a/tests/unit/gui/test_main_window.py b/tests/unit/gui/test_main_window.py new file mode 100644 index 00000000..d773d53b --- /dev/null +++ b/tests/unit/gui/test_main_window.py @@ -0,0 +1,74 @@ +import pytest +from unittest.mock import MagicMock, patch +from PySide6.QtCore import Qt +from aare.gui.main_window import MainWindow + +@pytest.fixture +def mock_ui_state(): + with patch("aare.gui.main_window.UIStateManager") as mock: + yield mock + +def test_main_window_init(qtbot, mock_ui_state): + # Mocking many things that MainWindow __init__ tries to do + # Especially things that hit the network or expected files + with patch("requests.get") as mock_get, \ + patch("aare.gui.main_window.DAQWorker"), \ + patch("aare.gui.main_window.SampleCameraThread"), \ + patch("aare.gui.main_window.PredictionSubscriber"), \ + patch("aare.gui.main_window.VideoThread"), \ + patch("aare.gui.main_window.JFJochDBusClient"), \ + patch("aare.gui.main_window.WorkflowSSEClient"), \ + patch("aare.gui.main_window.jwt.decode") as mock_jwt: + + mock_get.return_value.status_code = 200 + mock_get.return_value.json.return_value = {"status": "ok"} + mock_jwt.return_value = {"sub": "testuser", "staff": True, "pgroups": ["p123"], "session": 15} + + # Create a valid-looking fake JWT + fake_token = "header.payload.signature" + + # We need to provide all arguments to MainWindow + win = MainWindow( + base_url="http://localhost:5210", + token=fake_token, + default_image=None, + zmq_addr=None, + pred_zmq_addr=None, + beamline_cam_addr="localhost", + gonio_cam_addr="localhost", + gonio_cam_id=1 + ) + win.show() + qtbot.addWidget(win) + + assert win.windowTitle() == "AareGUI" + assert win.isVisible() + +def test_main_window_mount_view(qtbot, mock_ui_state): + with patch("requests.get"), \ + patch("aare.gui.main_window.DAQWorker"), \ + patch("aare.gui.main_window.SampleCameraThread"), \ + patch("aare.gui.main_window.PredictionSubscriber"), \ + patch("aare.gui.main_window.VideoThread"), \ + patch("aare.gui.main_window.JFJochDBusClient"), \ + patch("aare.gui.main_window.WorkflowSSEClient"), \ + patch("aare.gui.main_window.jwt.decode") as mock_jwt: + + mock_jwt.return_value = {"sub": "testuser", "staff": True, "pgroups": ["p123"], "session": 15} + fake_token = "header.payload.signature" + + win = MainWindow( + base_url=None, + token=fake_token, + default_image=None, + zmq_addr=None, + pred_zmq_addr=None, + beamline_cam_addr=None, + gonio_cam_addr=None, + gonio_cam_id=None + ) + qtbot.addWidget(win) + + win.mount_view() + # This just changes some panel visibility, hard to assert without deep inspection + # but we check it doesn't crash diff --git a/tests/unit/gui/test_models.py b/tests/unit/gui/test_models.py new file mode 100644 index 00000000..743d10b7 --- /dev/null +++ b/tests/unit/gui/test_models.py @@ -0,0 +1,84 @@ +import pytest +from PySide6.QtCore import Qt +from aare.gui.models.sample_queue_model import SampleQueueSpreadsheet +from aare.gui.models.user_sample_model import UserSampleSpreadsheet +from aare.common.models import SampleShortInfo, DewarAddress + +@pytest.fixture +def sample_list(): + return [ + SampleShortInfo(db_id=1, puck_name="P1", dewar_name="D1", sample_name="S1", run_number=1, user="U1", pin=1, location=DewarAddress(segment="A", pos=1)), + SampleShortInfo(db_id=2, puck_name="P2", dewar_name="D2", sample_name="S2", run_number=2, user="U2", pin=2, location=DewarAddress(segment="A", pos=2)), + SampleShortInfo(db_id=3, puck_name="P1", dewar_name="D1", sample_name="S3", run_number=3, user="U1", pin=3, location=DewarAddress(segment="A", pos=1)), + ] + +def test_user_sample_model_init(sample_list): + model = UserSampleSpreadsheet(samples=sample_list) + model.set_show_all_pgroups(True) # Ensure all pgroups are shown for testing + assert model.rowCount() == 3 + # Check if filtering works + # "User" is column 5 + model.set_filter("User", "U2") + assert model.rowCount() == 1 + model.clear_filter() + assert model.rowCount() == 3 + +def test_user_sample_model_column_filter(sample_list): + model = UserSampleSpreadsheet(samples=sample_list) + model.set_show_all_pgroups(True) + # Column 5 is user + model.set_column_filter(5, "U1") + assert model.rowCount() == 2 + model.clear_all_column_filters() + assert model.rowCount() == 3 + +def test_user_sample_model_unique_values(sample_list): + model = UserSampleSpreadsheet(samples=sample_list) + model.set_show_all_pgroups(True) + # Column 5 is User + users = model.unique_values_for_column(5) + assert "U1" in users + assert "U2" in users + assert len(users) == 2 + +def test_sample_queue_model_init(sample_list): + model = SampleQueueSpreadsheet(samples=sample_list[:2]) + assert model.rowCount() == 2 + assert model.columnCount() == 3 + assert model.data(model.index(0, 0), Qt.ItemDataRole.DisplayRole) == "U1" + assert model.data(model.index(1, 2), Qt.ItemDataRole.DisplayRole) == "S2" + +def test_sample_queue_model_update(sample_list): + model = SampleQueueSpreadsheet() + assert model.rowCount() == 0 + model.updateData(sample_list[:2]) + assert model.rowCount() == 2 + +def test_sample_queue_model_remove(sample_list): + model = SampleQueueSpreadsheet(samples=sample_list[:2]) + model.remove_sample(1) + assert model.rowCount() == 1 + assert model.samples[0].db_id == 2 + +def test_sample_queue_model_clear(sample_list): + model = SampleQueueSpreadsheet(samples=sample_list[:2]) + model.clearSamples() + assert model.rowCount() == 0 + +def test_sample_queue_model_set_running(sample_list): + model = SampleQueueSpreadsheet(samples=sample_list[:2]) + # Background color role for first row + color_not_running = model.data(model.index(0, 0), Qt.ItemDataRole.BackgroundRole) + model.set_running(True) + color_running = model.data(model.index(0, 0), Qt.ItemDataRole.BackgroundRole) + assert color_not_running != color_running + +def test_sample_queue_model_header(sample_list): + model = SampleQueueSpreadsheet(samples=sample_list[:2]) + assert model.headerData(0, Qt.Orientation.Horizontal, Qt.ItemDataRole.DisplayRole) == "User" + assert model.headerData(1, Qt.Orientation.Vertical, Qt.ItemDataRole.DisplayRole) == "2" + +def test_sample_queue_model_flags(sample_list): + model = SampleQueueSpreadsheet(samples=sample_list[:2]) + flags = model.flags(model.index(0, 0)) + assert flags & Qt.ItemFlag.ItemIsDropEnabled diff --git a/tests/unit/gui/test_panels.py b/tests/unit/gui/test_panels.py new file mode 100644 index 00000000..41bf5495 --- /dev/null +++ b/tests/unit/gui/test_panels.py @@ -0,0 +1,80 @@ +import pytest +from PySide6.QtCore import Qt +from aare.gui.panels.status_panel import StatusPanel +from aare.common.models import DAQStatusModel, BeamlineStatus, SessionStatus, SampleCameraSettings, BeamlineStateEnum +from aare.common.diffraction_geometry import DiffractionGeometry +from aare.common.sample_geometry import SampleGeometryModel +from aare.common.coordinate import Coordinate, SmargonCoordinate +from aare.common.beamline import MXBeamline + +@pytest.fixture +def mock_daq_status(): + status = DAQStatusModel( + bl=BeamlineStatus( + name="X10SA", + ring_current_mA=400.0, + flux_ph_s=1e12, + transmission=1.0, + cryojet_K=100.0, + front_light=50.0, + back_light=50.0, + shutter_open=False, + exp_shutter_open=False, + sample_camera=SampleCameraSettings(gain=1.0, exposure=0.02), + zoom=1.0, + commissioning_mode=False, + dtz_min=120.0, + dtz_max=1600.0 + ), + diffraction=DiffractionGeometry( + detector_description="EIGER", + detector_serial_number="123", + dtz_mm=200.0, + pixel_size_mm=0.075, + energy_keV=12.658, + beam_center_pxl=(1000.0, 1000.0), + detector_size_pxl=(4000, 4000), + poni_rot1_rad=0.0, + poni_rot2_rad=0.0 + ), + geom=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=10.5, + beam_size_mm=Coordinate(x=0.02, y=0.01) + ), + session=SessionStatus(), + state=BeamlineStateEnum.Maintenance, + busy=False + ) + return status + +def test_status_panel_update(qtbot, mock_daq_status): + panel = StatusPanel() + qtbot.addWidget(panel) + + # Initial state check (some values from __init__ defaults) + assert "400" in panel.ring_current.text() + + # Update with mock status + panel.update_daq_status(mock_daq_status) + + assert panel.omega.text() == "10.50" + assert "20.0" in panel.beam_size.text() # 0.02 * 1000 + assert "10.0" in panel.beam_size.text() # 0.01 * 1000 + assert "400.0" in panel.ring_current.text() + assert "1000.00" in panel.flux.text() # 1e12 / 1e9 = 1000 + assert "0.979" in panel.wavelength.text() + +def test_status_panel_low_current(qtbot, mock_daq_status): + panel = StatusPanel() + qtbot.addWidget(panel) + + mock_daq_status.bl.ring_current_mA = 300.0 + panel.update_daq_status(mock_daq_status) + + assert "color: red" in panel.ring_current.text() + assert "300.0" in panel.ring_current.text() diff --git a/tests/unit/gui/test_threads_logic.py b/tests/unit/gui/test_threads_logic.py new file mode 100644 index 00000000..1eb80707 --- /dev/null +++ b/tests/unit/gui/test_threads_logic.py @@ -0,0 +1,47 @@ +import pytest +from unittest.mock import MagicMock +from aare.gui.scan_logic.sample_mount_logic import SampleMountLogic +from aare.common.models import DAQStatusModel, SampleShortInfo, DewarAddress + +def test_sample_mount_logic_emits_on_change(): + logic = SampleMountLogic() + mock_slot = MagicMock() + logic.sample_changed.connect(mock_slot) + + sample1 = SampleShortInfo(db_id=1, puck_name="P1", dewar_name="D1", sample_name="S1", run_number=1, user="U1", pin=1, location=DewarAddress(segment="A", pos=1)) + sample2 = SampleShortInfo(db_id=2, puck_name="P1", dewar_name="D1", sample_name="S2", run_number=1, user="U1", pin=2, location=DewarAddress(segment="A", pos=1)) + + status = MagicMock(spec=DAQStatusModel) + status.sample = sample1 + + # First update, should emit + logic.update_daq_status(status) + mock_slot.assert_called_once_with(sample1) + mock_slot.reset_mock() + + # Second update same sample, should NOT emit + logic.update_daq_status(status) + mock_slot.assert_not_called() + + # Third update different sample, should emit + status.sample = sample2 + logic.update_daq_status(status) + mock_slot.assert_called_once_with(sample2) + +def test_sample_mount_logic_none_to_sample(): + logic = SampleMountLogic() + mock_slot = MagicMock() + logic.sample_changed.connect(mock_slot) + + sample1 = SampleShortInfo(db_id=1, puck_name="P1", dewar_name="D1", sample_name="S1", run_number=1, user="U1", pin=1, location=DewarAddress(segment="A", pos=1)) + + status = MagicMock(spec=DAQStatusModel) + status.sample = None + + logic.update_daq_status(status) + # Initial __sample is None, so if status.sample is also None, it won't emit + mock_slot.assert_not_called() + + status.sample = sample1 + logic.update_daq_status(status) + mock_slot.assert_called_once_with(sample1) diff --git a/tests/unit/gui/test_tutorials.py b/tests/unit/gui/test_tutorials.py new file mode 100644 index 00000000..f9d6a60b --- /dev/null +++ b/tests/unit/gui/test_tutorials.py @@ -0,0 +1,61 @@ +import pytest +from unittest.mock import MagicMock +from PySide6.QtWidgets import QWidget +from aare.gui.tutorials.tutorial_manager import TutorialManager +from aare.gui.tutorials.tutorial_models import ( + TutorialScenario, TutorialStepDefinition, TutorialMode, StepKind, TutorialTextRef, TutorialTarget, TargetKind +) +from aare.gui.tutorials.tutorial_runtime import ( + DictionaryTextResolver, NoOpActionExecutor, DictTargetResolver, TutorialEventBus +) + +@pytest.fixture +def mock_window(qapp): + return QWidget() + +@pytest.fixture +def tutorial_scenario(): + step1 = TutorialStepDefinition( + id="step1", + kind=StepKind.INFO, + title=TutorialTextRef(key="step1_title"), + body=TutorialTextRef(key="step1_body"), + target=TutorialTarget(kind=TargetKind.NONE, target_id="") + ) + scenario = TutorialScenario( + id="test_scenario", + title=TutorialTextRef(key="scenario_title"), + description=TutorialTextRef(key="scenario_desc"), + mode=TutorialMode.LINEAR, + steps=[step1] + ) + return scenario + +def test_tutorial_manager_start_stop(qtbot, mock_window, tutorial_scenario): + text_resolver = DictionaryTextResolver({ + "en": { + "step1_title": "Step 1", + "step1_body": "This is step 1", + "scenario_title": "Test", + "scenario_desc": "Test Desc" + } + }) + target_resolver = DictTargetResolver({}) + action_executor = NoOpActionExecutor() + event_bus = TutorialEventBus() + + manager = TutorialManager( + parent_window=mock_window, + target_resolver=target_resolver, + text_resolver=text_resolver, + action_executor=action_executor, + event_bus=event_bus + ) + manager.add_scenario(tutorial_scenario) + + manager.start("test_scenario") + assert manager.get_current_step() is not None + assert manager.get_current_step().id == "step1" + + manager.stop() + assert manager.get_current_step() is None diff --git a/tests/unit/gui/test_widgets.py b/tests/unit/gui/test_widgets.py new file mode 100644 index 00000000..ed4c0160 --- /dev/null +++ b/tests/unit/gui/test_widgets.py @@ -0,0 +1,46 @@ +import pytest +from PySide6.QtCore import Qt +from aare.gui.widgets.alert_banner import AlertBanner +from aare.gui.widgets.status_label import StatusLabel + +def test_alert_banner_show_message(qtbot): + banner = AlertBanner() + qtbot.addWidget(banner) + + assert not banner.isVisible() + + banner.show_message("Test Error", is_error=True) + assert banner.isVisible() + assert "Test Error" in banner._label.text() + assert "🛑" in banner._label.text() + + banner.show_message("Test Success", is_error=False) + assert banner.isVisible() + assert "Test Success" in banner._label.text() + assert "✅" in banner._label.text() + + banner.clear_message() + assert not banner.isVisible() + +def test_alert_banner_waiting(qtbot): + banner = AlertBanner() + qtbot.addWidget(banner) + + banner.show_waiting("Working", countdown_seconds=10) + assert banner.isVisible() + assert "Working" in banner._label.text() + assert "⏳" in banner._label.text() + assert "(10s)" in banner._label.text() + + # Tick manually if we wanted to test timer, but usually we just test state + banner._tick_countdown() + assert "(9s)" in banner._label.text() + +def test_status_label(qtbot): + label = StatusLabel(val=1.234, decimals=2) + qtbot.addWidget(label) + + assert label.text() == "1.23" + + label.new_value(5.6) + assert label.text() == "5.60"