"""Tests for lambda main handler.""" import json import exponent_server_sdk as expo from send_push_messages.logger import get_logger from send_push_messages.handler import send_and_save_receipts from send_push_messages.push_client import PushClient from send_push_messages.sqs_client import SQSClient from tests.expo_responses import expo_ok_response, expo_unregistered_device_response def test_handler_common_case(mocker): logger = get_logger('test') expo_provider = expo.PushClient() mocker.patch.object(expo_provider, 'publish_multiple') push_client = PushClient(provider=expo_provider, logger=logger) attrs = {"send_message_batch.return_value": "ok"} sqs_provider = mocker.Mock(**attrs) sqs_client = SQSClient(provider=sqs_provider, logger=logger) ok_messages_count = 2 inactive_device_messages_count = 2 invalid_token_messages_count = 2 invalid_data_messages_count = 1 token_prefix = "ExponentPushToken P2WrTIIEbSdGH6xBMz8iK" ok_messages = [ { 'id': i, 'to': f'{token_prefix}{i}', 'body': f'Body{i}', 'title': 'Ok Message', 'data': {}, 'device_id': i + 100 } for i in range(ok_messages_count)] inactive_devices_messages = [ { 'id': i, 'to': f'{token_prefix}{i}', 'body': f'Body{i}', 'title': 'Inactive Device Message', 'data': {}, 'device_id': i + 100 } for i in range(ok_messages_count, ok_messages_count + inactive_device_messages_count)] invalid_data_messages = [ { 'id': i, 'body': f'Body {i}', 'title': 'Invalid Data Message', 'data': {} } for i in range( ok_messages_count + inactive_device_messages_count, ok_messages_count + inactive_device_messages_count + invalid_token_messages_count + invalid_data_messages_count)] def _expo_bulk_response(_messages): _responses = [] for _m in _messages: _title = _m.title if _title == "Ok Message": _responses.append(expo_ok_response(_m)) elif _title == "Inactive Device Message": _responses.append(expo_unregistered_device_response(_m)) else: raise AssertionError(f"Unexpected message type: {_title}") return _responses expo_provider.publish_multiple.side_effect = _expo_bulk_response validated_messages = ok_messages + inactive_devices_messages messages = ok_messages + inactive_devices_messages + invalid_data_messages records = [{"id": i, "body": json.dumps({"Message": json.dumps(m)})} for i, m in enumerate(messages)] records.append({"body": json.dumps({"Message": json.dumps(({"title": "Message without id."}))})}) send_and_save_receipts(logger, records, len(records), push_client, sqs_client) sent_messages = expo_provider.publish_multiple.call_args.args[0] assert len(sent_messages) == len(validated_messages) for i, m in enumerate(sent_messages): original_message = validated_messages[i] assert m.to == original_message["to"] assert m.body == original_message["body"] assert m.title == original_message["title"] sent_receipts = sqs_provider.send_message_batch.call_args.kwargs["Entries"] assert len(sent_receipts) == len(messages) for r in sent_receipts: receipt_message_id = int(r["Id"]) receipt = json.loads(r["MessageBody"]) receipt_status = receipt["status"] if receipt_message_id < ok_messages_count: assert receipt_status == "delivered" else: assert receipt_status == "failed"