import json import uuid from dataclasses import asdict, dataclass from datetime import datetime from typing import Callable, Iterable, List, Optional from botocore.exceptions import ClientError from apollo_push_client import constants from apollo_push_client.exceptions import NoSessionError, PushClientError def chunks(array, chuck_size: int = constants.DEFAULT_CHUNK_SIZE) -> Iterable: """Yield chunks from array""" for i in range(0, len(array), chuck_size): yield array[i:i + chuck_size] @dataclass class PushData: id: str artist_name: str isrc: str track_id: str track_name: str topic: str position: int target: str vendor: str country_code: Optional[str] = None change: Optional[int] = None playlist_id: Optional[str] = None playlist_name: Optional[str] = None playlist_image_url: Optional[str] = None @dataclass class PlaylistUpdatePushData: id: str topic: str target: str vendor: str date: str playlist_id: str country_code: Optional[str] = None playlist_name: Optional[str] = None playlist_image_url: Optional[str] = None @dataclass class Push: id: str title: str message: str data: PushData tokens: List[str] user_id: str sound: str = 'default' channel_id: str = 'all' class PushClient: """Class for pushing PushMessages to SQS Queue and saving them to DB""" def __init__(self, logger_client, sqs_client, config, session=None): self.logger_client = logger_client self.sqs_client = sqs_client self.queue = config.PUSH_NOTIFICATIONS_QUEUE self.push_save_mode = constants.PushSaveMode( getattr(config, "PUSH_SAVE_MODE", None) or constants.PushSaveMode.ALL) self.set_session(session) def set_session(self, session): self._session = session @property def session(self): if self._session is None: raise NoSessionError(f"{self} session has not been set. Use set_session method to set one.") return self._session def save_messages_to_db( self, messages_chunk: List[Push], push_date_getter: Optional[Callable], track_id_getter: Optional[Callable] ): 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.session.bulk_save_objects(objects) self.session.commit() def send_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_client.error(ex) self.logger_client.info(messages_chunk) return [] def process_messages( self, messages: List[Push], chunk_size: int = constants.DEFAULT_CHUNK_SIZE, save_to_db: bool = False, raise_exc: bool = False, 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. chunk_size: size of chunk to be proceed. save_to_db: flag, if True messages is being saved to db. raise_exc: flag, if True sqs client exception are raised. 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, chunk_size): messages_chunk = self.send_messages(messages_chunk, raise_exc) if messages_chunk and save_to_db: self.save_messages_to_db(messages_chunk, push_date_getter, track_id_getter)