import json from typing import Union from marshmallow import Schema, ValidationError, fields, post_dump, validate, validates from apollo_main_db.push_notifications.models import NotificationSetting, PushMessage, VendorEnum from core.serializers import BasePaginationOutputSchema, ModelSchema, PaginationMixin, QueryParams, QueryWebArgs from main_db.base import session as db_session class UserDeviceInput(Schema): """Schema for user device input validation.""" token = fields.String(required=True, validate=validate.Length(min=1, max=255)) class PushMessageSchema(QueryWebArgs, PaginationMixin): pass class PushMessageExtraData(Schema): """Schema for serializing extra data from PushMessage data field""" artist_name = fields.String() track_name = fields.String() position = fields.Integer() target = fields.String() isrc = fields.String() country_code = fields.String() change = fields.Integer() class PushMessageOutputItemSchema(Schema): """Schema for serializing objects of PushMessage model""" topic = fields.Function(lambda o: o.topic.name.lower()) vendor = fields.Function(lambda o: o.vendor.name.lower()) extra_data = fields.Method("get_extra_data") track_id = fields.Method("get_track_id") date = fields.DateTime(attribute="created_at") time_delta = fields.Method("calculate_timedelta") class Meta: fields = ("id", "title", "message", "extra_data", "track_id", "date", "is_new", "topic", "vendor", "time_delta") model = PushMessage def get_extra_data(self, obj: PushMessage) -> dict: try: data = json.loads(obj.data) except (TypeError, json.decoder.JSONDecodeError): data = {} return PushMessageExtraData().dumps(data).data def get_track_id(self, obj: PushMessage) -> Union[str, int]: if obj.vendor == VendorEnum.APPLE: try: return int(obj.track_id) except (ValueError, TypeError): return "" return obj.track_id def calculate_timedelta(self, obj: PushMessage) -> int or None: context_date = self.context.get("current_datetime") try: timedelta_obj = context_date - obj.created_at return int(timedelta_obj.total_seconds()) except (ValueError, TypeError, AttributeError): return @staticmethod def process_extra_data(data: dict) -> dict: extra_data = data.pop("extra_data", "{}") processed_data = dict(data) for key, value in json.loads(extra_data).items(): if value is None: continue processed_data[key] = value return processed_data @post_dump(pass_many=True) def dump_extra_data(self, data: dict, many: bool) -> Union[list, dict]: if many: result = [] for d in data: item = self.process_extra_data(d) result.append(item) return result return self.process_extra_data(data) class PushMessageOutputSchema(BasePaginationOutputSchema): items = fields.List(fields.Nested(PushMessageOutputItemSchema)) class FilterPushMessage(QueryParams): is_new = fields.Boolean() class PushMessageInput(ModelSchema): id = fields.UUID() is_new = fields.Boolean() class Meta: model = PushMessage class PushMessageBulkInput(Schema): data = fields.Nested(PushMessageInput, many=True) class PushMessageSettingOutput(Schema): id = fields.Integer() code = fields.String() enabled = fields.Boolean() class PushMessageSettingInput(Schema): id = fields.Integer() class Meta: model = NotificationSetting @validates("id") def validate_id(self, value): if not db_session.query(self.Meta.model).filter_by(id=value).scalar(): raise ValidationError("ID isn't applicable")