"""Unit tests for src.sqs_helpers.""" import json from unittest.mock import AsyncMock, MagicMock, patch import pytest from src import config, sqs_helpers from vector_utils.aws_utils import sqs @pytest.mark.parametrize( ("encoder_id", "priority", "expected_dms_priority"), [ (18, 15, 10000), (18, 14, 14), (23, 15, 15), ], ) @patch("src.sqs_helpers.sqs.format_queue_name") def test_get_queue_name_priority_normalization( mock_format_queue_name: MagicMock, encoder_id: int, priority: int, expected_dms_priority: int, ) -> None: """Non-physical encoders above top-priority threshold use a shared queue slot.""" mock_format_queue_name.return_value = "test-queue" sqs_helpers.get_queue_name_for_delivery_job( encoder_id=encoder_id, dms_id=1, priority=priority, encoding_order_priority=1 ) mock_format_queue_name.assert_called_once_with( sqs.ENCODING_QUEUE_NAME_PATTERN, encoder_id=encoder_id, dms_priority=expected_dms_priority, priority=1, env=config.ENVIRONMENT, ) @patch("src.sqs_helpers.config.GRAS_DELIVERY_STORE_IDS", [42]) def test_get_queue_name_gras_store_physical_encoder() -> None: """GRAS stores with PHYSICAL_ENCODER use the GRAS delivery queue.""" result = sqs_helpers.get_queue_name_for_delivery_job( encoder_id=23, dms_id=42, priority=5, encoding_order_priority=1 ) expected = sqs.GRAS_DELIVERY_QUEUE_NAME_PATTERN.format(env=config.ENVIRONMENT) assert result == expected @patch("src.sqs_helpers.sqs.format_queue_name") @patch("src.sqs_helpers.config.GRAS_DELIVERY_STORE_IDS", [42]) def test_get_queue_name_gras_store_non_physical_encoder( mock_format_queue_name: MagicMock, ) -> None: """GRAS stores with non-PHYSICAL_ENCODER use the encoding queue.""" mock_format_queue_name.return_value = "encoding-queue" result = sqs_helpers.get_queue_name_for_delivery_job( encoder_id=18, dms_id=42, priority=5, encoding_order_priority=1 ) mock_format_queue_name.assert_called_once_with( sqs.ENCODING_QUEUE_NAME_PATTERN, encoder_id=18, dms_priority=5, priority=1, env=config.ENVIRONMENT, ) assert result == "encoding-queue" @pytest.fixture(autouse=True) def _clear_queue_url_cache() -> None: """Reset the module-level queue URL cache between tests.""" sqs_helpers._queue_url_cache.clear() async def test_get_or_create_queue_url_returns_existing() -> None: """When the queue exists, return its URL without calling create_queue.""" mock_sqs_client = MagicMock() mock_sqs_client.get_queue_url = AsyncMock( return_value={"QueueUrl": "https://existing"} ) mock_sqs_client.create_queue = AsyncMock() result = await sqs_helpers._get_or_create_queue_url(mock_sqs_client, "q") assert result == "https://existing" mock_sqs_client.create_queue.assert_not_called() async def test_get_or_create_queue_url_creates_when_missing() -> None: """If get_queue_url raises QueueDoesNotExist, fall back to create_queue.""" class QueueDoesNotExist(Exception): pass mock_sqs_client = MagicMock() mock_sqs_client.exceptions.QueueDoesNotExist = QueueDoesNotExist mock_sqs_client.get_queue_url = AsyncMock(side_effect=QueueDoesNotExist()) mock_sqs_client.create_queue = AsyncMock(return_value={"QueueUrl": "https://new"}) result = await sqs_helpers._get_or_create_queue_url(mock_sqs_client, "q") assert result == "https://new" async def test_get_or_create_queue_url_uses_cache_on_second_call() -> None: """A resolved queue URL is cached and not looked up again.""" mock_sqs_client = MagicMock() mock_sqs_client.get_queue_url = AsyncMock( return_value={"QueueUrl": "https://cached"} ) first = await sqs_helpers._get_or_create_queue_url(mock_sqs_client, "q") second = await sqs_helpers._get_or_create_queue_url(mock_sqs_client, "q") assert first == second == "https://cached" mock_sqs_client.get_queue_url.assert_awaited_once() def _mock_aioboto3_session( monkeypatch: pytest.MonkeyPatch, sqs_client: MagicMock ) -> MagicMock: """Patch aioboto3.Session so its async client() context yields sqs_client.""" session = MagicMock() client_cm = MagicMock() client_cm.__aenter__ = AsyncMock(return_value=sqs_client) client_cm.__aexit__ = AsyncMock(return_value=False) session.client = MagicMock(return_value=client_cm) monkeypatch.setattr( "src.sqs_helpers.aioboto3.Session", MagicMock(return_value=session) ) return session def _build_message(eqd_id: int) -> dict[str, int]: return {"encoding_queue_detail_id": eqd_id} async def test_fan_out_to_sqs_sends_all_batches( monkeypatch: pytest.MonkeyPatch, ) -> None: """fan_out_to_sqs resolves URLs, batches, and sends every message.""" monkeypatch.setattr(config, "SQS_BATCH_SIZE", 2) monkeypatch.setattr(config, "SQS_WORKER_COUNT", 2) sqs_client = MagicMock() sqs_client.get_queue_url = AsyncMock( side_effect=lambda QueueName: {"QueueUrl": f"https://sqs/{QueueName}"} ) sqs_client.send_message_batch = AsyncMock(return_value={}) _mock_aioboto3_session(monkeypatch, sqs_client) messages_by_queue = { "q1": [_build_message(1), _build_message(2), _build_message(3)], "q2": [_build_message(4)], } failed = await sqs_helpers.fan_out_to_sqs( messages_by_queue, [1, 2, 3, 4], MagicMock() ) assert failed == [] # 3 batches total: q1 → [1,2] + [3], q2 → [4] assert sqs_client.send_message_batch.await_count == 3 sent_ids: list[int] = [] for call in sqs_client.send_message_batch.await_args_list: for entry in call.kwargs["Entries"]: sent_ids.append(int(entry["Id"])) assert json.loads(entry["MessageBody"])["encoding_queue_detail_id"] == int( entry["Id"] ) assert sorted(sent_ids) == [1, 2, 3, 4] async def test_fan_out_to_sqs_reports_partial_failures( monkeypatch: pytest.MonkeyPatch, ) -> None: """SQS partial-failure responses are aggregated into the returned list.""" monkeypatch.setattr(config, "SQS_BATCH_SIZE", 10) monkeypatch.setattr(config, "SQS_WORKER_COUNT", 1) sqs_client = MagicMock() sqs_client.get_queue_url = AsyncMock(return_value={"QueueUrl": "https://sqs/q"}) sqs_client.send_message_batch = AsyncMock( return_value={ "Failed": [{"Id": "7", "Code": "X", "Message": "m", "SenderFault": True}] } ) _mock_aioboto3_session(monkeypatch, sqs_client) failed = await sqs_helpers.fan_out_to_sqs( {"q": [_build_message(7), _build_message(8)]}, [7, 8], MagicMock() ) assert failed == [7] async def test_fan_out_to_sqs_reports_send_exceptions( monkeypatch: pytest.MonkeyPatch, ) -> None: """When send_message_batch raises, every id in that batch is marked failed.""" monkeypatch.setattr(config, "SQS_BATCH_SIZE", 10) monkeypatch.setattr(config, "SQS_WORKER_COUNT", 1) sqs_client = MagicMock() sqs_client.get_queue_url = AsyncMock(return_value={"QueueUrl": "https://sqs/q"}) sqs_client.send_message_batch = AsyncMock(side_effect=RuntimeError("boom")) _mock_aioboto3_session(monkeypatch, sqs_client) failed = await sqs_helpers.fan_out_to_sqs( {"q": [_build_message(11), _build_message(12)]}, [11, 12], MagicMock() ) assert sorted(failed) == [11, 12] async def test_fan_out_to_sqs_marks_unresolved_queue_messages_failed( monkeypatch: pytest.MonkeyPatch, ) -> None: """If a queue's URL cannot be resolved, all its messages are marked failed.""" monkeypatch.setattr(config, "SQS_BATCH_SIZE", 10) monkeypatch.setattr(config, "SQS_WORKER_COUNT", 1) sqs_client = MagicMock() sqs_client.get_queue_url = AsyncMock(side_effect=RuntimeError("boom")) sqs_client.send_message_batch = AsyncMock() _mock_aioboto3_session(monkeypatch, sqs_client) failed = await sqs_helpers.fan_out_to_sqs( {"q": [_build_message(1), _build_message(2)]}, [1, 2], MagicMock() ) assert sorted(failed) == [1, 2] sqs_client.send_message_batch.assert_not_called() async def test_fan_out_to_sqs_resets_all_on_client_failure( monkeypatch: pytest.MonkeyPatch, ) -> None: """If the SQS client cannot be created, every promoted id is returned for reset.""" session = MagicMock() session.client = MagicMock(side_effect=RuntimeError("no client")) monkeypatch.setattr( "src.sqs_helpers.aioboto3.Session", MagicMock(return_value=session) ) failed = await sqs_helpers.fan_out_to_sqs( {"q": [_build_message(1)]}, [1, 2, 3], MagicMock() ) assert sorted(failed) == [1, 2, 3] async def test_fan_out_to_sqs_partially_resolves_queues( monkeypatch: pytest.MonkeyPatch, ) -> None: """Resolvable queues are sent; messages in unresolved queues are marked failed.""" monkeypatch.setattr(config, "SQS_BATCH_SIZE", 10) monkeypatch.setattr(config, "SQS_WORKER_COUNT", 2) def get_queue_url(QueueName: str) -> dict[str, str]: if QueueName == "bad": raise RuntimeError("boom") return {"QueueUrl": f"https://sqs/{QueueName}"} sqs_client = MagicMock() sqs_client.get_queue_url = AsyncMock(side_effect=get_queue_url) sqs_client.send_message_batch = AsyncMock(return_value={}) _mock_aioboto3_session(monkeypatch, sqs_client) failed = await sqs_helpers.fan_out_to_sqs( {"good": [_build_message(1)], "bad": [_build_message(2)]}, [1, 2], MagicMock(), ) assert failed == [2] sqs_client.send_message_batch.assert_awaited_once()