from typing import Union from marshmallow import Schema, fields, post_load, ValidationError from marshmallow.validate import OneOf from charts.constants.common import COUNTRY_GLOBAL, CHART_TRACKS_DEFAULT_LIMIT, ChartDigestSoundRecordingItems, \ CHART_DIGEST_SOUND_RECORDING_ITEMS_VALUES from charts.schemas.fields.pagination import get_pagination_mixin ChartTracksPaginationMixin = get_pagination_mixin(page_limit=CHART_TRACKS_DEFAULT_LIMIT) class ChartDigestTracks: """Request & Response schema for chart/digest/tracks/""" class Request(Schema, ChartTracksPaginationMixin): """ definition_key: custom chart key built from {platform}_{type}_{frequency} definition_key example: spotify_viral_daily, apple_default_daily e.t.c. """ definition_key = fields.String(load_default=None) country_code = fields.String(load_default=COUNTRY_GLOBAL) chart_id = fields.String(load_default=None) chart_date = fields.Date(load_default=None) filter = fields.String( load_default=ChartDigestSoundRecordingItems.ALL.value, validate=OneOf(CHART_DIGEST_SOUND_RECORDING_ITEMS_VALUES) ) @post_load def _check_only_one_passed(self, data, *args, **kwargs): """Check passed params for definition_key and chart_id Only one of them should be passed """ _error_string = "one of 'definition_key' & 'chart_id' params should be passed." _definition_key = data["definition_key"] _chart_id = data["chart_id"] if _definition_key and _chart_id: raise ValidationError(f"Only {_error_string}") if not any([_definition_key, _chart_id]): raise ValidationError(f"At least {_error_string}") return data class Response(Schema): """Response schema with additional calculation of trend. Trend is shown for tracks that were in a chart on previous chart update. For entries and re-entries the trend is None """ uuid = fields.String() chartId = fields.String() position = fields.Integer() positionChange = fields.Integer() streams = fields.Integer() totalStreams = fields.Integer() totalShazams = fields.Integer() views = fields.Integer() numTracks = fields.Integer() totalUnits = fields.Integer() trackId = fields.String() isrc = fields.String() publicSoundRecordingId = fields.String() upc = fields.String() publicProductId = fields.String() videoId = fields.String() channelId = fields.String() channelName = fields.String() chartDate = fields.String() peakTimestamp = fields.String() peakPosition = fields.Integer() daysOnChart = fields.Integer() artistNames = fields.String() trend = fields.Method("_get_trend") newlyEntered = fields.Method("_blue_dot") def _get_trend(self, data, *args, **kwargs) -> Union[int, None]: is_new = self._is_new(data=data) return None if is_new else data["positionChange"] def _blue_dot(self, data, *args, **kwargs) -> bool: return self._is_new(data=data) def _is_new(self, data, *args, **kwargs): return data["isrc"] in self.context["entries"]