import json from datetime import date, timedelta from typing import Any, Dict, Iterable, Optional, Sequence from apollo_notifications.constants import SPOTIFY_CATEGORY_ID, SPOTIFY_PLAYLIST_IMAGE_URL_MASK from apollo_notifications.playlists.utils import get_playlist_update_push_key from apollo_notifications.utils import load_datetime from tests import factories from apollo_main_db.push_notifications.models import TopicEnum, VendorEnum NO_VALUE_KEY = "__NO_VALUE" def get_formatted_today_yesterday(): today = date.today() yesterday = today - timedelta(days=1) return today.strftime("%Y-%m-%d"), yesterday.strftime("%Y-%m-%d") def playlist_id(i): return f"playlist_id_{i}" def playlist_name(i): return f"playlist_name_{i}" def playlist_uri(i): return f"playlist_uri_{i}" def user_id(i): return f"user_id_{i}" def account_id(i): return f"account_id_{i}" def create_basic_structures( playlists_n: int, # creates batch of playlists from 1 to integer passed users_n: int, # creates batch of users from 1 to integer passed playlist_to_updates_dates: Dict[int, Iterable[str]], category_id_to_playlists_map: Dict[int, Iterable[int]] = None, # by default spotify category is used blacklisted_playlists: Iterable[int] = None, market: str = "us", market_to_playlists_map: Dict[str, Iterable[int]] = None, active_devices: int = 2, last_update_str: Optional[str] = None ): if last_update_str: # log job last running factories.ApolloKeyValueStorageFactory.create( key="spotify_pl_upd_push", value=last_update_str ) # create playlists playlist_to_category_id_map = {} for category_id, playlist_idxs in (category_id_to_playlists_map or {}).items(): for pl_idx in playlist_idxs: playlist_to_category_id_map[pl_idx] = category_id playlist_to_market_map = {} for _market, playlist_idxs in (market_to_playlists_map or {}).items(): for pl_idx in playlist_idxs: playlist_to_market_map[pl_idx] = _market playlists = [factories.SpotifyPlaylistFactory.create( id=playlist_id(i), uri=playlist_uri(i), name=playlist_name(i), country_code=playlist_to_market_map.get(i, market), buzz_category_id=playlist_to_category_id_map.get(i, SPOTIFY_CATEGORY_ID) ) for i in range(1, playlists_n + 1)] for pl_idx in blacklisted_playlists or []: factories.PlaylistBlacklistFactory.create(vendor="spotify", playlist_id=playlists[pl_idx - 1].id) # create updates for pl_idx, update_dates in playlist_to_updates_dates.items(): for _datetime_str in update_dates: factories.SpotifyPlaylistStatistics.create( playlist_id=playlists[pl_idx - 1].id, market=playlist_to_market_map.get(pl_idx, market), tracklist_updated_at=_datetime_str ) # create user data for user_idx in range(1, users_n + 1): [factories.UserDeviceFactory.create( token=f"inactive_{user_idx}{i}", user_id=user_id(user_idx), is_active=False) for i in range(1, 3)] [factories.UserDeviceFactory.create( token=f"token_{user_idx}{i}", user_id=user_id(user_idx), is_active=True) for i in range(1, active_devices + 1)] def create_push_message_data( playlist_idx: int, user_idx: int, topic: str, country_code: Optional[str], title: str, date: str, vendor: str = "spotify", add_playlist_image: bool = False ) -> Dict[str, Any]: _playlist_id, _playlist_name = playlist_id(playlist_idx), playlist_name(playlist_idx) data = { "id": str(get_playlist_update_push_key( date, _playlist_id, user_id(user_idx), topic, country_code, vendor)), "topic": topic, "target": _playlist_name, "vendor": vendor, "country_code": country_code or '', "date": date, "playlist_id": _playlist_id, "playlist_name": _playlist_name, "playlist_image_url": SPOTIFY_PLAYLIST_IMAGE_URL_MASK.format( playlist_id=_playlist_id) if add_playlist_image else None } return { "id": data["id"], "user_id": user_id(user_idx), "title": title, "topic": TopicEnum.__members__[topic.upper()], "vendor": VendorEnum.__members__[vendor.upper()], "date": load_datetime(date), "data": data } def verify_push( expected_data, push, json_fields=("data",), push_attr_getter=lambda p, attr: getattr(p, attr, NO_VALUE_KEY)): for name, expected_value in expected_data.items(): value = push_attr_getter(push, name) assert value != NO_VALUE_KEY if name in json_fields: verify_push(expected_value, json.loads(value), push_attr_getter=lambda p, attr: p.get(attr, NO_VALUE_KEY)) else: assert value == expected_value def verify_push_messages(expected_push_id_to_data_map: Dict[str, Dict[str, Any]], push_messages: Sequence): assert len(expected_push_id_to_data_map) == len(push_messages) for push in push_messages: expected_push = expected_push_id_to_data_map.get(str(push.id)) assert expected_push is not None verify_push(expected_push, push)