import json import os import random import uuid from dataclasses import asdict from datetime import datetime from typing import List, Type, Union import boto3 import pytest import requests_mock from apollo_main_db.push_notifications.models import PushMessage, PushMessageReduced from moto import mock_sqs from sqlalchemy.orm import Session from apollo_notifications.logger import logger from apollo_notifications.push_client.client import PushClient from apollo_notifications.push_client.config import PushClientConfig from apollo_notifications.push_client.data_classes import PlaylistUpdatePushData, Push, PushData class PushSchemaConfig(PushClientConfig): USER_DATA_SCHEMA: str = "http" USER_DATA_HOST: str = "localhost" USER_DATA_APIKEY: str = "test" PUSH_NOTIFICATIONS_QUEUE: str = "test_sqs" PUSH_SAVE_MODE: str = "ALL" def get_id_list(count: int = 4) -> List[str]: rnd = random.Random() rnd.seed(123) return [str(uuid.UUID(int=rnd.getrandbits(128), version=4)) for _ in range(count)] def get_messages(is_change: bool = True, count: int = 4) -> List[Push]: inner_class = PushData if is_change else PlaylistUpdatePushData id_list = get_id_list(count) return [ Push( **{ "id": id_list[i - 1], "title": f"title {i}", "message": f"message {i}", "tokens": [f"token_{i}_{j}" for j in range(1, i)], "user_id": f"user_{i}", "account_id": f"acc_{i}", "sound": "default" if i % 3 else "clock", "channel_id": "push" if i % 2 else "all", "data": inner_class( **{ "id": id_list[i - 1], "topic": "playlist_additions" if i % 2 else "playlist_update", "target": f"target {i}" if i % 3 else None, "vendor": "apple" if i % 2 else "spotify", "country_code": f"c{i}", "playlist_id": f"pl_id_{i}" if i % 3 else None, "playlist_name": f"pl_name {i}" if i % 3 else None, "playlist_image_url": f"pl_url_{i}" if i % 3 else None, **( { "artist_name": f"artist_{i}", "isrc": f"isrc_{i}", "track_id": f"tr_id_{i}", "track_name": f"track {i}", "position": i, "change": 3 - i, } if is_change else {"date": "2022-04-01"} ), }, ), }, ) for i in range(1, count + 1) ] def get_db_messages(message_list: List[Push], is_change: bool): return [ dict( id=message.id, user_id=message.user_id, title=message.title, message=message.message, data=asdict(message.data), track_id=message.data.track_id if is_change else None, topic=message.data.topic, vendor=message.data.vendor, date=None if is_change else datetime.fromisoformat(message.data.date), ) for message in message_list ] def load_from_db(db_session: Session, model: Type[Union[PushMessageReduced, PushMessage]]) -> List[dict]: result = [ i._asdict() for i in db_session.query( model.id, model.user_id, model.title, model.message, model.data, model.track_id, model.topic, model.vendor, model.date, ).all() ] for item in result: item["id"] = str(item["id"]) item["vendor"] = item["vendor"].name.lower() item["topic"] = item["topic"].name.lower() item["data"] = json.loads(item["data"]) return result def get_feed_messages(save_db: bool = True, is_change: bool = True, count: int = 4) -> dict: id_list = get_id_list(count) return { "public": False, "data": [ { "ttl": 3600, "message_id": i if save_db else id_list[i - 1], "account_id": f"acc_{i}", "meta": { "dsp": "apple" if i % 2 else "spotify", "country_code": f"c{i}", "type": "starred_track_top_playlist_entry" if i % 2 else "starred_playlist_tracklist_update", }, "data": { "recipient": {"devices": [{"expo_token": f"token_{i}_{j}"} for j in range(1, i)]}, "content": { "playlist": { "id": f"pl_id_{i}" if i % 3 else None, "name": f"pl_name {i}" if i % 3 else None, "image_url": f"pl_url_{i}" if i % 3 else None, **({} if is_change else {"updated_at": "2022-04-01"}), }, "title": f"title {i}", "body": f"message {i}", **( { "track": { "id": f"tr_id_{i}", "isrc": f"isrc_{i}", "name": f"track {i}", "artists": [{"name": f"artist_{i}"}], }, "current_position": i, "previous_position": (2 * i - 3), } if is_change else {} ), }, }, } for i in range(1, count + 1) ] } def get_push_messages(save_db: bool = True, is_change: bool = True, count: int = 4) -> dict: id_list = get_id_list(count) return { "public": True, "data": [ { "ttl": 3600, "message_id": i if save_db else id_list[i - 1], "account_id": f"acc_{i}", "meta": { "dsp": "apple" if i % 2 else "spotify", "country_code": f"c{i}", "type": "starred_track_top_playlist_entry" if i % 2 else "starred_playlist_tracklist_update", }, "to": f"token_{i}_{j}", "title": f"title {i}", "body": f"message {i}", "sound": "default" if i % 3 else "clock", "channel_id": "push" if i % 2 else "all", "data": { "id": id_list[i - 1], "topic": "playlist_additions" if i % 2 else "playlist_update", "target": f"target {i}" if i % 3 else None, "vendor": "apple" if i % 2 else "spotify", "country_code": f"c{i}", "playlist_id": f"pl_id_{i}" if i % 3 else None, "playlist_name": f"pl_name {i}" if i % 3 else None, "playlist_image_url": f"pl_url_{i}" if i % 3 else None, "date": None, **( { "artist_name": f"artist_{i}", "isrc": f"isrc_{i}", "track_id": f"tr_id_{i}", "track_name": f"track {i}", "position": i, "change": 3 - i, "date": None, } if is_change else {"date": "2022-04-01"} ) } } for i in range(1, count + 1) for j in range(1, i) ] } @pytest.mark.parametrize( "is_change,send_sqs,save_db,send_feed,send_push", ( (True, True, True, True, True), (False, True, True, True, True), (True, False, True, True, True), (True, False, False, True, True), (True, False, False, False, True), (True, True, True, False, False), ), ) @requests_mock.Mocker(kw="requests_mocker") @mock_sqs def test_process_messages(is_change, send_sqs, save_db, send_feed, send_push, db_session, **kwargs): os.environ["PUSH_SQS_MESSAGES"] = "1" if send_sqs else "" os.environ["PUSH_SAVE_TO_DB"] = "1" if save_db else "" os.environ["PUSH_USER_DATA_MESSAGES_FEED"] = "1" if send_feed else "" os.environ["PUSH_USER_DATA_MESSAGES_PUSH"] = "1" if send_push else "" requests_mocker = kwargs["requests_mocker"] mocked_feed = requests_mocker.register_uri( "POST", "http://localhost/api/service/messages/feed/", json={"status": "OK"} ) mocked_push = requests_mocker.register_uri( "POST", "http://localhost/api/service/messages/push/", json={"status": "OK"} ) config = PushSchemaConfig() sqs_client = boto3.client("sqs", region_name="us-east-1") sqs_client.create_queue(QueueName="test_sqs") client = PushClient(logger, sqs_client, config, db_session=db_session) message_list = get_messages(is_change) client.process_messages(message_list, **({"push_date_getter": None} if is_change else {"track_id_getter": None})) sqs_message = sqs_client.receive_message(QueueUrl="test_sqs", VisibilityTimeout=900) sqs_message_list = [] while sqs_message.get("Messages"): sqs_message_list.append(json.loads(sqs_message["Messages"][0]["Body"])) sqs_message = sqs_client.receive_message(QueueUrl="test_sqs", VisibilityTimeout=900) assert sqs_message_list == ([asdict(i) for i in message_list] if send_sqs else []) assert load_from_db(db_session, PushMessageReduced) == ( get_db_messages(message_list, is_change) if save_db and send_sqs else [] ) assert load_from_db(db_session, PushMessage) == ( get_db_messages(message_list, is_change) if save_db and send_sqs else [] ) assert mocked_feed.call_count == int(send_feed) if send_feed: assert mocked_feed.last_request.json() == get_feed_messages(save_db and send_sqs, is_change) assert mocked_push.call_count == int(send_push) if send_push: assert mocked_push.last_request.json() == get_push_messages(save_db and send_sqs, is_change)