"""Test the SQS connector of the reverse payout flow.""" from unittest import mock from botocore import exceptions import pytest from accounting.flows.reserve_payouts.connectors import sqs @pytest.mark.parametrize('exception, expected_result, should_raise', [ (exceptions.BotoCoreError(), False, False), (exceptions.ClientError({}, 1), False, False), (Exception, False, True), (None, True, False), ]) @pytest.mark.parametrize('expected_queue_name', [ 'test-queue_name', 'dev-queue-name2' ]) @mock.patch('accounting.flows.reserve_payouts.connectors.sqs.setting') @mock.patch('accounting.flows.reserve_payouts.connectors.sqs.sqs') def test_health_check( mock_util_sqs, mock_setting, expected_queue_name, exception, expected_result, should_raise): """Test health_check function.""" mock_queue = mock.MagicMock() mock_util_sqs.get_queue_by_name.return_value = mock_queue if exception: mock_util_sqs.get_queue_number_of_messages.side_effect = exception mock_setting.SQS_QUEUE = expected_queue_name if should_raise: with pytest.raises(exception): sqs.health_check() else: result = sqs.health_check() assert result.bool == expected_result mock_util_sqs.get_queue_by_name.assert_called_with(expected_queue_name) mock_util_sqs.get_queue_number_of_messages.assert_called_with(mock_queue) @pytest.mark.parametrize('message_count', [0, 42, 4242]) @pytest.mark.parametrize('expected_queue_name', [ 'test-queue_name', 'dev-queue-name2' ]) @mock.patch('accounting.flows.reserve_payouts.connectors.sqs.setting') @mock.patch('accounting.flows.reserve_payouts.connectors.sqs.sqs') def test_get_queue_message_count( mock_util_sqs, mock_setting, expected_queue_name, message_count): """Test get_queue_message_count function.""" mock_queue = mock.MagicMock() mock_util_sqs.get_queue_number_of_messages.return_value = message_count mock_util_sqs.get_queue_by_name.return_value = mock_queue mock_setting.SQS_QUEUE = expected_queue_name result = sqs.get_queue_message_count() assert result == message_count mock_util_sqs.get_queue_by_name.assert_called_with(expected_queue_name) mock_util_sqs.get_queue_number_of_messages.assert_called_with(mock_queue) @pytest.mark.parametrize('message_id_const, message_body_const', [ ('Id', 'MessageBody'), ('different_id', 'different_message_body'), ]) @pytest.mark.parametrize('vendor_data', [ {}, {'vendor_id': 1, 'some_field': 'any_value'}, ]) @mock.patch('accounting.flows.reserve_payouts.connectors.sqs.json') @mock.patch('accounting.flows.reserve_payouts.connectors.sqs.uuid') @mock.patch('accounting.flows.reserve_payouts.connectors.sqs.sqs_constants') def test_construct_vendor_message( sqs_constants, uuid, json, vendor_data, message_id_const, message_body_const): """Test construct_vendor_message function.""" expected_uuid = 'test_uuid' uuid.uuid4.return_value = expected_uuid mock_dumps = mock.MagicMock() json.dumps.return_value = mock_dumps sqs_constants.SQS_MESSAGE_ID = message_id_const sqs_constants.SQS_MESSAGE_BODY = message_body_const expected_message = { message_id_const: expected_uuid, message_body_const: mock_dumps, } actual_message = sqs.construct_vendor_message(vendor_data) assert actual_message == expected_message json.dumps.assert_called_with(vendor_data) uuid.uuid4.assert_called_once() @mock.patch('accounting.flows.reserve_payouts.connectors.sqs.setting') @mock.patch('accounting.flows.reserve_payouts.connectors.sqs.sqs') def test_batch_vendor_send_messages(utils_sqs, setting): """Test batch_vendor_send_messages function.""" queue_name = 'test_queue' max_batch = 99 message_list = [1, 2] setting.SQS_QUEUE = queue_name setting.SQS_MAX_BATCH_SIZE = max_batch mock_queue = mock.MagicMock() utils_sqs.get_queue_by_name.return_value = mock_queue sublist_generator = mock.MagicMock() sublist_1 = mock.MagicMock() sublist_2 = mock.MagicMock() sublist_generator.__iter__.return_value = [sublist_1, sublist_2] utils_sqs.yield_sublists.return_value = sublist_generator sqs.batch_vendor_send_messages(message_list) utils_sqs.get_queue_by_name.assert_called_with(queue_name) utils_sqs.yield_sublists.assert_called_with(message_list, max_batch) expected_send_messages_calls = [ mock.call(Entries=sublist_1), mock.call(Entries=sublist_2), ] mock_queue.send_messages.assert_has_calls( expected_send_messages_calls, any_order=True) @mock.patch('accounting.flows.reserve_payouts.connectors.sqs.json') @mock.patch('accounting.flows.reserve_payouts.connectors.sqs.setting') @mock.patch('accounting.flows.reserve_payouts.connectors.sqs.sqs') def test_batch_vendor_send_messages_exception(utils_sqs, setting, json): """Test batch_vendor_send_messages function.""" queue_name = 'test_queue' max_batch = 99 message_list = [1, 2] setting.SQS_QUEUE = queue_name setting.SQS_MAX_BATCH_SIZE = max_batch mock_queue = mock.MagicMock() failed_response = { 'Failed': [{'message_id': '12345'}] } mock_queue.send_messages.return_value = failed_response utils_sqs.get_queue_by_name.return_value = mock_queue json.dumps.return_value = 'failed items json' sublist_generator = mock.MagicMock() sublist_1 = mock.MagicMock() sublist_2 = mock.MagicMock() sublist_generator.__iter__.return_value = [sublist_1, sublist_2] utils_sqs.yield_sublists.return_value = sublist_generator with pytest.raises(Exception): sqs.batch_vendor_send_messages(message_list) expected_send_messages_calls = [ mock.call(Entries=sublist_1), ] mock_queue.send_messages.assert_has_calls( expected_send_messages_calls, any_order=True) json.dumps.assert_called_with(failed_response) @mock.patch('accounting.flows.reserve_payouts.connectors.sqs.setting') @mock.patch('accounting.flows.reserve_payouts.connectors.sqs.sqs') def test_get_message(utils_sqs, setting): """Test get_message function.""" queue_name = 'test_queue' sqs_visibility_timeout = '99' setting.SQS_QUEUE = queue_name setting.SQS_VISIBILITY_TIMEOUT = sqs_visibility_timeout mock_message = mock.MagicMock() mock_queue = mock.MagicMock() utils_sqs.get_message.return_value = mock_message utils_sqs.get_queue_by_name.return_value = mock_queue result = sqs.get_message() assert result == mock_message utils_sqs.get_queue_by_name.assert_called_with(queue_name) utils_sqs.get_message.assert_called_with(mock_queue) mock_queue.set_attributes.assert_called_with( Attributes={'VisibilityTimeout': sqs_visibility_timeout}) @mock.patch('accounting.flows.reserve_payouts.connectors.sqs.setting') @mock.patch('accounting.flows.reserve_payouts.connectors.sqs.sqs') def test_get_message_exception_is_raised(utils_sqs, setting): """Test get_message function.""" mock_message = mock.MagicMock() mock_queue = mock.MagicMock() utils_sqs.get_message.return_value = mock_message utils_sqs.get_message.side_effect = Exception() utils_sqs.get_queue_by_name.return_value = mock_queue with pytest.raises(Exception): sqs.get_message()