from typing import List, Union from apollo_utils.service.schemas.fields.string import CustomStringField from marshmallow import EXCLUDE, Schema, fields, post_load, validate from multidict import MultiDictProxy from server.core.constants import APPLE, GLOBAL_MARKET, SPOTIFY, Service from server.core.exceptions import BadRequest class ListField(fields.List): """ Custom field for handling list of query params E.g. ?ids=foo,bar """ def __init__(self, *args, **kwargs): self._unify = kwargs.pop("unify", True) super().__init__(*args, **kwargs) def _deserialize(self, value: str, attr: str, data: MultiDictProxy, **kwargs) -> List[Union[str, int]]: value = data.getall(attr) if hasattr(data, "getall") else data.get(attr) if value and len(value) == 1: value = value[0].split(",") value = super()._deserialize(value, attr, data, **kwargs) if value and self._unify: value = sorted(list(set(value)), key=value.index) return value class QueryWebArgs(Schema): """Base class for schemes used with @querystring_schema decorator. Allows raise Bad Request in case of having errors (instead of 422). """ class Meta: strict = True unknown = EXCLUDE def handle_error(self, exc, data, **kwargs): """Overwrite original ValidationError with custom exception.""" raise BadRequest(extra=exc.messages) class TrackIsrcSearchDeserializer(QueryWebArgs): market = fields.String(required=False) isrc = fields.String(required=True) class TracksSearchV1Deserializer(QueryWebArgs): class TrackItem(QueryWebArgs): isrc = CustomStringField(upper=True, required=True) dsp = fields.String(required=True, validate=validate.OneOf([APPLE, SPOTIFY])) known_id = fields.Field(data_key="id") known_market = CustomStringField(lower=True, data_key="country_code") @post_load def prepare_id(self, data, **kwargs): known_id = data.get("known_id") if known_id: data["known_id"] = str(known_id) return data market = fields.String(required=False, missing=GLOBAL_MARKET) tracks = fields.List(fields.Nested(TrackItem()), missing=[]) class HealthCheckParams(QueryWebArgs): """Health check endpoint params.""" include = ListField(fields.String(validate=validate.OneOf(Service.ALL_VALUES)), missing=[]) class OffsetDefaultMixin: offset = fields.Integer(missing=0, validate=validate.Range(min=0))