"""Tests for PushClient.""" from collections import defaultdict from typing import List import exponent_server_sdk as expo import pytest from send_push_messages.config import Config as config from send_push_messages.logger import get_logger from send_push_messages.push_client import PushClient from tests.expo_responses import expo_ok_response @pytest.mark.parametrize( 'message_numbers_per_projects', ( [3], [2, 3, 4], [1, 3] ) ) def test_push_client_send_for_multiple_projects(mocker, message_numbers_per_projects): logger = get_logger('test') provider = expo.PushClient() max_message_id_and_project = [(sum(message_numbers_per_projects[: i + 1]), i) for i, _ in enumerate(message_numbers_per_projects)] total_projects_number = len(message_numbers_per_projects) total_messages_number = sum(message_numbers_per_projects) token_prefix = "ExponentPushToken P2WrTIIEbSdGH6xBMz8iK" messages = [expo.PushMessage( to=f"{token_prefix}{i}", body=f'Message {i}', title='Title', data={} ) for i in range(total_messages_number)] def _expo_response(_messages) -> List[expo.PushTicket]: project_to_tokens = defaultdict(list) for _m in _messages: message_id = int(_m.to[-1]) project_id = next(p_id for m_id, p_id in max_message_id_and_project if message_id < m_id) project_to_tokens[project_id].append(_m.to) if len(project_to_tokens) > 1: errors = [{ 'code': 'PUSH_TOO_MANY_EXPERIENCE_IDS', 'message': 'All push notification messages in the same request must be for the same project; ' 'check the details field to investigate conflicting tokens.', 'details': project_to_tokens, 'isTransient': False}] raise expo.PushServerError( message="Request failed", response="", errors=errors, response_data={'errors': errors} ) return [expo_ok_response(m) for m in _messages] mocker.patch.object(provider, 'publish_multiple', side_effect=_expo_response) client = PushClient(provider=provider, logger=logger) receipts = client.send(messages) assert len(receipts) == len(messages) assert set([id(r.push_message) for r in receipts]) == set([id(m) for m in messages]) call_count = total_projects_number if total_projects_number > 1: call_count += 1 assert provider.publish_multiple.call_count == call_count def test_push_client_retry(mocker): logger = get_logger('test') provider = expo.PushClient() number_of_messages = 2 token_prefix = "ExponentPushToken P2WrTIIEbSdGH6xBMz8iK" messages = [expo.PushMessage( to=f"{token_prefix}{i}", body=f'Message {i}', title='Title', sound="boo", data={} ) for i in range(number_of_messages)] def _expo_response(_messages): errors = [ { 'code': 'VALIDATION_ERROR', 'message': '"[0].sound" must be one of [default, null, object].', 'isTransient': False } ] raise expo.PushServerError( message="Request failed", response="", errors=errors, response_data={'errors': errors} ) mocker.patch.object(provider, 'publish_multiple', side_effect=_expo_response) client = PushClient(provider=provider, logger=logger) with pytest.raises(expo.PushServerError): client.send(messages) assert provider.publish_multiple.call_count == config.EXPO_RETRY_COUNT + 1