from apollo_utils.core.constants import MARKET_WORLDWIDE from apollo_utils.service.exceptions import JsonValidationError, QueryParamsValidationError, UnsupportedMediaType from apollo_utils.service.schemas.fields.list import BaseUniqueArgsListField, SplitUniqueListField from apollo_utils.service.utils.market import fix_country_code_worldwide from flask import request from marshmallow import Schema, ValidationError, fields, post_load, pre_load, validate from typing import List from src.cache.constants import CacheMode from src.constants.core import DEFAULT_PAGINATION_LIMIT, MARKET_GLOBAL, SPOTIFY_MARKET_GLOBAL class EmptySchema(Schema): """flask_apispec has a bug that schema can not be None, it throws error in this case, so we can use this instead.""" pass class QueryDefaultParams(Schema): """ Base class for serializers that validate query string parameters. """ def load_from_request(self) -> dict: """ Parse, clean and validate query parameters from request. If serializer has List fields, corresponding query parameters are correctly parsed as lists. Returns: Dictionary with validated and cleaned query parameters. Raises: QueryParamsValidationError: If parameters validation failed. """ try: cleaned_data = self.load(request.args) except ValidationError as errors: raise QueryParamsValidationError(extra=errors.data) return cleaned_data class QueryWebArgs(Schema): """Base class for schemes used with flask-apispec @use_kwargs decorator to load data from request query. Allows to raise Bad Request in case of having errors. """ class Meta: strict = True def handle_error(self, exc, data, **kwargs): """Overwrite original ValidationError with custom exception.""" raise QueryParamsValidationError(extra=exc.messages) class JsonBodySchema(Schema): """Base class for schemes used with flask-apispec @use_kwargs decorator to load data from request json. Allows to raise Bad Request in case of having errors. """ @pre_load def check_type(self, data, **kwargs): if not request.is_json: raise UnsupportedMediaType( 'Unsupported media type "{}". Expected "application/json".'.format(request.mimetype) ) return data def handle_error(self, exc, data, **kwargs): """Overwrite original ValidationError with custom exception.""" raise JsonValidationError(extra=exc.messages) class UniqueArgsListField(BaseUniqueArgsListField): """Class for getting data for list field from query string parameters, by default removes duplicates.""" def _get_attr(self, attr): return request.args.getlist(attr) class MarketListOptionalMixin: markets_list = SplitUniqueListField( fields.String(validate=validate.Length(max=10, min=2)), data_key="market", required=False, missing=None ) class GlobalMarketCodeToFullMixin: """Parse _gl to global.""" @post_load def parse_market(self, data, **kwargs): market = data.get("market") or data.get("market_list") if isinstance(market, list): if SPOTIFY_MARKET_GLOBAL in market: market.remove(SPOTIFY_MARKET_GLOBAL) market.append(MARKET_GLOBAL) elif market == SPOTIFY_MARKET_GLOBAL: data["market"] = MARKET_GLOBAL return data def get_country_code_worldwide_mixin( field_name_list: List[str] or tuple, worldwide_country_code: str = MARKET_WORLDWIDE ): class CountryCodeToGlobalMixin: @post_load def parse_market(self, data, **kwargs): for field_name in field_name_list: original_country_code = data.get(field_name) if original_country_code: data[field_name] = fix_country_code_worldwide(original_country_code, worldwide_country_code) return data return CountryCodeToGlobalMixin FullMarketCodeToGlobalMixin = get_country_code_worldwide_mixin(["market"], SPOTIFY_MARKET_GLOBAL) def get_pagination_mixin(page_offset: int = 0, page_limit: int or None = DEFAULT_PAGINATION_LIMIT) -> type: class PaginationParams: offset = fields.Integer(validate=validate.Range(min=0), missing=page_offset) limit = fields.Integer(validate=validate.Range(min=1), missing=page_limit) include_count = fields.Boolean(missing=True) return PaginationParams PaginationMixin = get_pagination_mixin() class BasePaginationOutputSchema(Schema): count = fields.Integer(validate=validate.Range(min=0)) next = fields.String() previous = fields.String() class PaginationLinksNoneResponse(Schema): count = fields.Integer(default=0) next = fields.String(default=None) previous = fields.String(default=None) class CacheModeMixin: """Cache Mode validation mixin.""" _cache_mode = fields.String(validate=validate.OneOf(CacheMode.values()), allow_none=True)