from datetime import datetime from marshmallow import EXCLUDE, Schema, ValidationError, fields, post_load, pre_dump, validate, validates_schema from apollo_messages_views.config import EVENT_CODE, EVENT_DSP, EVENT_HOLDERS, EVENT_TYPES class MetaExcludeUnknown(Schema): """Schema to exclude unknown fields to have a predicted structure of event/message""" class Meta: unknown = EXCLUDE class MetaItem(MetaExcludeUnknown): """Schema for nested meta of an event/message""" dsp = fields.String(required=True, validate=validate.OneOf(EVENT_DSP)) country_code = fields.String(required=True, allow_none=True) type = fields.String(required=True, validate=validate.OneOf(EVENT_TYPES)) subject = fields.String(required=True) class RecipientItem(MetaExcludeUnknown): """Schema for nested recipient data ex: { "devices": [ { "expo_token": "expo_token_string", "id": "123" }, ] } """ class DeviceItem(MetaExcludeUnknown): expo_token = fields.String(required=True) id = fields.Integer() devices = fields.List(fields.Nested(DeviceItem, required=True), validate=validate.Length(min=1), required=True) def data_mixin(content_item): class DataItemMixin: class DataItem(MetaExcludeUnknown): content = fields.Nested(content_item, required=True) recipient = fields.Nested(RecipientItem, required=True) data = fields.Nested(DataItem, required=True) return DataItemMixin class Artist(Schema): id = fields.Str(required=True) name = fields.Str(required=True) order = fields.Int(load_only=True) class TrackItem(MetaExcludeUnknown): id = fields.Str(required=True) isrc = fields.Str(required=True) name = fields.Str(required=True) image_url = fields.Str(required=True) favorites_id = fields.Int() artists = fields.List(fields.Nested(Artist, required=True), validate=validate.Length(min=1), required=True) @pre_dump def sort_artists(self, in_data, **kwargs): artists = in_data.get("artists") if artists and artists[0].get("order") is not None: in_data["artists"].sort(key=lambda x: x["order"]) return in_data class TrackItemWithArtistNames(TrackItem): artists_names = fields.Method("get_artists_names") def get_artists_names(self, obj): return ", ".join(i["name"] for i in obj["artists"]) class PlaylistItem(MetaExcludeUnknown): id = fields.Str(required=True) name = fields.Str(required=True) updated_at = fields.Str(required=True) image_url = fields.Str(required=True) class BaseRawMessage(MetaExcludeUnknown): _check_raw = True id = fields.Integer(required=True, load_only=True) message_id = fields.Integer(required=False, allow_none=True) event_id = fields.Int(required=False, allow_none=True) account_id = fields.Int() app = fields.String(required=True, validate=validate.OneOf(EVENT_HOLDERS), load_only=True) code = fields.String(required=True, validate=validate.OneOf([EVENT_CODE]), load_only=True) ttl = fields.Integer(required=True) meta = fields.Nested(MetaItem, required=True) created_at = fields.DateTime(required=True, load_only=True) @validates_schema(pass_original=True) def check_raw(self, data, original_data, **kwargs): if not self._check_raw: return self._check_actual(data) self._check_holder(original_data) self._check_code(original_data) def _check_actual(self, in_data): current_dt = self.context.get("current_dt") or datetime.utcnow() if not (current_dt - in_data["created_at"]).total_seconds() <= in_data["ttl"]: raise ValidationError( f"{self.__class__.__name__} schema considered the event={in_data['id']} as outdated\n" f"created_at={in_data['created_at']}, ttl={in_data['ttl']}" ) def _check_holder(self, raw_data): holder = raw_data.get("app") if not holder or holder.lower() not in EVENT_HOLDERS: raise ValidationError( f"{self.__class__.__name__} schema considered the event={raw_data['id']} as invalid\n" f"app={holder} is not allowed, allowed values are: {EVENT_HOLDERS}" ) def _check_code(self, raw_data): code = raw_data.get("code") if not code or code.lower() != EVENT_CODE: raise ValidationError( f"{self.__class__.__name__} schema considered the event={raw_data['id']} as invalid\n" f"code={code} is not allowed, allowed value is {EVENT_CODE}" ) class BaseMessage(BaseRawMessage): _check_raw = False app = fields.String(validate=validate.OneOf(EVENT_HOLDERS), load_only=True) code = fields.String(validate=validate.OneOf([EVENT_CODE]), load_only=True) @post_load def add_message_id(self, in_data, **kwargs): # we need for basic_message.id -> feed_message.message_id only if not in_data.get("message_id"): in_data["message_id"] = in_data["id"] return in_data @pre_dump def set_ttl(self, in_data, **kwargs): in_data["ttl"] = in_data["ttl"] - (self.context["current_dt"] - in_data["created_at"]).total_seconds() return in_data