import json import os import uuid from dataclasses import asdict from logging import Logger from typing import Callable, Dict, Iterable, List, Optional import requests from botocore.client import BaseClient from botocore.exceptions import ClientError from sqlalchemy.orm import Session from apollo_notifications.push_client import constants from apollo_notifications.push_client.config import PushClientConfig from apollo_notifications.push_client.data_classes import Push from apollo_notifications.push_client.exceptions import NoSessionError, PushClientError from apollo_notifications.user_data.base_client import BaseUserDataClient def chunks(array, chuck_size: int) -> Iterable: """Yield chunks from array""" for i in range(0, len(array), chuck_size): yield array[i:i + chuck_size] class PushClient(BaseUserDataClient): """Class for pushing PushMessages to SQS Queue and saving them to DB""" def __init__(self, logger: Logger, sqs_client: BaseClient, config: PushClientConfig, db_session: Session = None): super().__init__(config, logger) self._sqs_client = sqs_client self._queue = os.environ.get("PUSH_NOTIFICATIONS_QUEUE", config.PUSH_NOTIFICATIONS_QUEUE) self._chunk_size = int(os.environ.get("PUSH_CHUNK_SIZE", config.PUSH_CHUNK_SIZE)) self._sqs_messages = os.environ.get("PUSH_SQS_MESSAGES", config.PUSH_SQS_MESSAGES) self._sqs_messages_raise_exc = os.environ.get( "PUSH_SQS_MESSAGES_RAISE_EXC", config.PUSH_SQS_MESSAGES_RAISE_EXC ) self._save_to_db = os.environ.get("PUSH_SAVE_TO_DB", config.PUSH_SAVE_TO_DB) self._user_data_messages_feed = os.environ.get( "PUSH_USER_DATA_MESSAGES_FEED", config.PUSH_USER_DATA_MESSAGES_FEED ) self._user_data_messages_feed_public = os.environ.get( "PUSH_USER_DATA_MESSAGES_FEED_PUBLIC", config.PUSH_USER_DATA_MESSAGES_FEED_PUBLIC ) self._user_data_messages_feed_ttl = int( os.environ.get("PUSH_USER_DATA_MESSAGES_FEED_TTL", config.PUSH_USER_DATA_MESSAGES_FEED_TTL) ) self._user_data_messages_push = os.environ.get( "PUSH_USER_DATA_MESSAGES_PUSH", config.PUSH_USER_DATA_MESSAGES_PUSH ) self._user_data_messages_push_public = os.environ.get( "PUSH_USER_DATA_MESSAGES_PUSH_PUBLIC", config.PUSH_USER_DATA_MESSAGES_PUSH_PUBLIC ) self._user_data_messages_push_ttl = int( os.environ.get("PUSH_USER_DATA_MESSAGES_PUSH_TTL", config.PUSH_USER_DATA_MESSAGES_PUSH_TTL) ) self._push_save_mode = constants.PushSaveMode( os.environ.get("PUSH_SAVE_MODE", getattr(config, "PUSH_SAVE_MODE", None) or constants.PushSaveMode.REDUCED) ) self.set_db_session(db_session) def set_db_session(self, db_session): self._db_session = db_session @property def db_session(self): if self._db_session is None: raise NoSessionError(f"{self} session has not been set. Use set_session method to set one.") return self._db_session def _save_messages_to_db( self, messages_chunk: List[Push], push_date_getter: Optional[Callable], track_id_getter: Optional[Callable] ) -> Dict[str, str]: models = constants.PUSH_SAVE_MODE_TO_MODELS[self._push_save_mode] objects = [] for message in messages_chunk: for model in models: objects.append( model(**dict( id=uuid.UUID(message.id), user_id=message.user_id, title=message.title, message=message.message.encode(), data=json.dumps(asdict(message.data)), track_id=track_id_getter and track_id_getter(message), topic=message.data.topic, vendor=message.data.vendor, date=push_date_getter and push_date_getter(message), )) ) self.db_session.bulk_save_objects(objects, return_defaults=True) self.db_session.commit() return {str(i.id): i.inner_id for i in objects} def _send_sqs_messages(self, messages_chunk: List[Push], raise_exc: bool = False): try: self._sqs_client.send_message_batch( QueueUrl=self._queue, Entries=[ { "Id": m.id, "MessageBody": json.dumps(asdict(m), default=str) } for m in messages_chunk], ) return messages_chunk except ClientError as ex: if raise_exc: raise PushClientError from ex self._logger.error(ex) self._logger.info(messages_chunk) return [] def _send_post_request(self, url: str, body: dict): """Send post request to the user data API service. Args: url: Relative URL path. body: Post request body. """ request = requests.Request( "POST", f"{self._base_url}/api/{url}", json=body, headers={"X-App-Slug": self._app_slug} ) return self._send_request(request) @staticmethod def __get_message_type(topic: str) -> str: """Get message type by topic. Args: topic: Message topic. Returns: Message type. """ if topic == "playlist_additions": return "starred_track_top_playlist_entry" elif topic == "starred_playlist_additions": return "starred_track_starred_playlist_entry" elif topic == "playlist_update": return "starred_playlist_tracklist_update" else: raise NotImplementedError("Incorrect topic.") @staticmethod def __get_feed_content(message: Push) -> dict: """Get data content for feed endpoint. Args: message: Push message data. Returns: Content data. """ result = { "playlist": { "id": message.data.playlist_id, "name": message.data.playlist_name, "image_url": message.data.playlist_image_url, **({"updated_at": message.data.date} if getattr(message.data, "date", None) is not None else {}), }, "title": message.title, "body": message.message, } if getattr(message.data, "track_id", None) is not None: result.update( { "track": { "id": message.data.track_id, "isrc": message.data.isrc, "name": message.data.track_name, "artists": [{"name": message.data.artist_name}], }, "current_position": message.data.position, "previous_position": ( (message.data.position - message.data.change) if message.data.change is not None else None ), } ) return result def _send_messages_feed(self, messages_chunk: List[Push], id_mapping: Dict[str, str]): """Send info to the user data API service/messages/feed/ endpoint. Args: messages_chunk: Message data list. id_mapping: Message ID mapping. """ request_data = [] for message in messages_chunk: request_data.append( { "ttl": self._user_data_messages_feed_ttl, "message_id": id_mapping.get(message.id, message.id), "account_id": message.account_id, "meta": { "dsp": message.data.vendor, "country_code": message.data.country_code, "type": self.__get_message_type(message.data.topic), }, "data": { "recipient": {"devices": [{"expo_token": token} for token in message.tokens]}, "content": self.__get_feed_content(message), }, } ) if request_data: return self._send_post_request( url="service/messages/feed/", body={"public": self._user_data_messages_feed_public, "data": request_data}, ) def _send_messages_push(self, messages_chunk: List[Push], id_mapping: Dict[str, str], track_id_getter: Callable): """Send info to the user data API service/messages/push/ endpoint. Args: messages_chunk: Message data list. id_mapping: Message ID mapping. track_id_getter: Track ID getter. """ request_data = [] for message in messages_chunk: for message_token in message.tokens: request_data.append( { "ttl": self._user_data_messages_push_ttl, "message_id": id_mapping.get(message.id, message.id), "account_id": message.account_id, "meta": { "dsp": message.data.vendor, "country_code": message.data.country_code, "type": self.__get_message_type(message.data.topic), }, "to": message_token, "title": message.title, "body": message.message, "sound": message.sound, "channel_id": message.channel_id, "data": asdict(message.data), } ) if request_data: return self._send_post_request( url="service/messages/push/", body={"public": self._user_data_messages_push_public, "data": request_data}, ) def process_messages( self, messages: List[Push], push_date_getter: Callable = lambda item: item.data.date, track_id_getter: Callable = lambda item: item.data.track_id ): """ Process push-messages: send them to sqs and save to database if needed. Args: messages: list of push objects to send. push_date_getter: function to get date from passed message. track_id_getter: function to get track id from passed message. """ for messages_chunk in chunks(messages, self._chunk_size): sqs_messages = ( self._send_sqs_messages(messages_chunk, self._sqs_messages_raise_exc) if self._sqs_messages else None ) id_mapping = ( self._save_messages_to_db(messages_chunk, push_date_getter, track_id_getter) if sqs_messages and self._save_to_db else {} ) if self._user_data_messages_feed: self._send_messages_feed(messages_chunk, id_mapping) if self._user_data_messages_push: self._send_messages_push(messages_chunk, id_mapping, track_id_getter)