"""Common schemas used across the whole project.""" from copy import deepcopy from marshmallow import Schema, fields, post_dump, pre_dump from oto import response as oto_response from sound_recordings.constants.store import ALL_SOURCES from sound_recordings.schemas import fields as custom_fields class BaseSchema(Schema): """Marshmallow base schema.""" class Meta: """Metadata for base schema.""" strict = True @classmethod def normalize(cls, obj, permissions_filter=None, many=False): """Normalize input to match schema.""" instance = cls() if permissions_filter: instance.context.update(permissions_filter) if permissions_filter["subaccount_ids"]: instance.context.update( {"account_id": permissions_filter["subaccount_ids"][0]} ) elif permissions_filter["label_ids"]: instance.context.update( {"account_id": permissions_filter["label_ids"][0]} ) return instance.dump(obj, many=many) @classmethod def normalized_response(cls, obj, permissions_filter=None, many=False): """Return an normalized OTO Response object.""" normalized = cls.normalize(obj, permissions_filter, many) return oto_response.Response(normalized) class SourceErrorSchema(BaseSchema): """Marshmallow schema for source error.""" types = fields.List(fields.String) code = fields.String() message = fields.String() class SourceSchema(BaseSchema): """Marshmallow schema for source item.""" @pre_dump def fill_in_name(self, data_in, **kwargs): """Pre-dump to fill in the name of a source if missing.""" result = deepcopy(data_in) result["name"] = next( (item for item in ALL_SOURCES if item["id"] == data_in["id"]), data_in )["name"] return result id_ = fields.Integer(attribute="id", data_key="id") name = fields.String() error = fields.Nested(SourceErrorSchema) class StreamDataSchema(BaseSchema): """Marshmallow schema for stream details.""" skip_rate = fields.Number(data_key="skip_rate") growth_percentage = fields.Number(data_key="growth_percentage") all_time = fields.Integer(data_key="all_time") class StreamBreakdownItemSchema(BaseSchema): """Schema for streams by date.""" streams = fields.Integer() skip_rate = fields.Float() saves = fields.Integer() date = custom_fields.UnifiedDate() class DownloadBreakdownItemSchema(BaseSchema): """Schema for downloads by date.""" downloads = fields.Integer() date = custom_fields.UnifiedDate() class StreamBreakdownSchema(StreamDataSchema): """Marshmallow schema for stream breakdown.""" items = fields.Nested(StreamBreakdownItemSchema, many=True) class DownloadBreakdownSchema(BaseSchema): """Marshmallow schema for stream breakdown.""" items = fields.Nested(DownloadBreakdownItemSchema, many=True) class PlaylistBreakdownSchema(BaseSchema): """Schema for playlists with streams.""" playlist_name = fields.String() followers = fields.Integer() playlist_url = fields.String() playlist_image = fields.String() store_id = fields.Integer() streams = fields.Integer() class TopCountriesStreamsBreakdownSchema(BaseSchema): """Schema for top countries with streams.""" country_code = fields.String() streams = fields.Integer() class StreamsCountriesSchema(BaseSchema): """Schema for streams by country.""" country_code = fields.String() items = fields.Nested(StreamBreakdownItemSchema, many=True) class DownloadsCountriesSchema(BaseSchema): """Schema for downloads by country.""" code = fields.String() items = fields.Nested(DownloadBreakdownItemSchema, many=True) class StreamsStoresSchema(BaseSchema): """Schema for streams by store.""" store_id = fields.Integer(attribute="id", data_key="id") name = fields.String() items = fields.Nested(StreamBreakdownItemSchema, many=True) class DownloadsStoresSchema(BaseSchema): """Schema for downloads by store.""" store_id = fields.Integer(attribute="id", data_key="id") name = fields.String() items = fields.Nested(DownloadBreakdownItemSchema, many=True) class TopCountriesDownloadsBreakdownSchema(BaseSchema): """Schema for top countries with downloads.""" country_code = fields.String() downloads = fields.Integer() class LabelEntitySchema(BaseSchema): """Base schema for entities with label and subaccount fields.""" label_id = fields.Integer(data_key="label_id") subaccount_name = fields.String(data_key="subaccount") subaccount_id = fields.Integer(data_key="subaccount_id") @post_dump def _remove_subaccount_name(self, data, **kwargs): """Remove subaccount name if account_id matches subaccount_id.""" account_id = self.context.get("account_id") if account_id and data.get("subaccount"): if str(account_id) == str(data.get("subaccount_id")): del data["subaccount"] return data class ArtistCoreMetadataSchema(BaseSchema): """Marshmallow schema for artist core metadata.""" artist_name = fields.String(data_key="artist_name") artist_type = fields.String(data_key="type", attribute="type") class ProductCoreMetadataSchema(LabelEntitySchema): """Marshmallow schema for product core metadata.""" aliases = { "release_name": "product_name", "display_upc": "upc", "sale_start_date": "sales_start_date", "name": "product_name", "product_version": "version", } product_id = fields.String(data_key="product_id") upc = fields.String(data_key="upc") product_name = fields.String(data_key="product_name") version = fields.String(data_key="version") delivered_version = fields.String(data_key="delivered_version") artist_name = fields.String(data_key="artist_name") format_ = fields.String(data_key="format", attribute="format") image_location = fields.Url(data_key="image_location") release_date = custom_fields.UnifiedDate(data_key="release_date") sales_start_date = custom_fields.UnifiedDate(data_key="sales_start_date") artists = fields.Nested(ArtistCoreMetadataSchema, many=True) display_upc = fields.String(required=False) release_name = fields.String(required=False) sale_start_date = custom_fields.UnifiedDate(required=False) product_version = fields.String(required=False) @post_dump def _handle_aliases(self, data, **kwargs): """Wrong input name: right output name.""" for key, value in self.aliases.items(): if data.get(key) and not data.get(value): data[value] = data[key] del data[key] if data.get(key) and data.get(value): del data[key] return data class TrackCoreMetadataSchema(LabelEntitySchema): """Marshmallow for track core metadata.""" aliases = {"name": "track_name", "track_version": "version"} isrc = fields.String() artist_name = fields.String(data_key="artist_name") image_location = fields.Url(data_key="image_location") name = fields.String(required=False) version = fields.String(required=False) artists = fields.Nested(ArtistCoreMetadataSchema, many=True) # Aliases track_version = fields.String(required=False) track_name = fields.String(required=False) @post_dump def _handle_aliases(self, data, **kwargs): """Wrong input name: right output name.""" for key, value in self.aliases.items(): if data.get(key) and not data.get(value): data[value] = data[key] del data[key] if data.get(key) and data.get(value): del data[key] return data class SoundRecordingCoreMetadataSchema(LabelEntitySchema): """Marshmallow for track core metadata.""" isrc = fields.String() class StoreItemSchema(BaseSchema): """Schema for streams by store.""" store_id = fields.Integer(attribute="id", data_key="id") name = fields.String() items = fields.Nested(StreamBreakdownItemSchema, many=True) class StreamsSOSSchema(BaseSchema): """Schema for streams by sos.""" source = fields.String() items = fields.Nested( StreamBreakdownItemSchema(exclude=("skip_rate", "saves")), many=True )