from types import GeneratorType from typing import Any, Dict, Generator, List, Union from flask import request from marshmallow import MarshalResult, Schema, ValidationError, fields, post_load, pre_load, validate from sqlalchemy import Column from core.constants import DEFAULT_PAGINATION_LIMIT, MARKET_GLOBAL, VENDORS, Service from core.exceptions import JsonValidationError, QueryParamsValidationError, UnsupportedMediaType from main_db.base import session 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 QueryParams(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. """ data = {} for field_name, field in self.fields.items(): if field.load_from: field_name = field.load_from if field_name not in request.args: continue if isinstance(field, fields.List): data[field_name] = request.args.getlist(field_name) else: data[field_name] = request.args[field_name] cleaned_data, errors = self.load(data) if errors: raise QueryParamsValidationError(extra=errors) return cleaned_data 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. """ cleaned_data, errors = self.load(request.args) if errors: raise QueryParamsValidationError(extra=errors) 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. """ class Meta: strict = True @pre_load def check_type(self, in_data, **kwargs): if not request.is_json: raise UnsupportedMediaType( 'Unsupported media type "{}". Expected "application/json".'.format(request.mimetype) ) def handle_error(self, exc, data, **kwargs): """Overwrite original ValidationError with custom exception.""" raise JsonValidationError(extra=exc.messages) class ModelSchema(Schema): """Base class for SQLAlchemy model-based Schemas.""" _model = None def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) if not hasattr(self.Meta, "model"): raise ValueError("ModelSchema Meta should have model") self._model = self.Meta.model @post_load def make_instance(self, data: dict, **kwargs): """Deserialize data to an instance of the model. Update an existing row loaded by primary key(s) in the data else create a new row. """ instance = self.get_instance(data) create = self.context.get("create") if instance: for key, value in data.items(): setattr(instance, key, value) session.merge(instance) elif not instance and create: instance = self._model(**kwargs) session.add(instance) session.commit() return instance def get_instance(self, data: dict): """Retrieve an existing record by primary key(s).""" props = self.get_primary_keys() filters = {prop.key: data.get(prop.key) for prop in props} return session.query(self._model).filter_by(**filters).first() def get_primary_keys(self) -> List[Column]: """Get primary key properties for a SQLAlchemy model.""" mapper = self._model.__mapper__ return [mapper.get_property_by_column(column) for column in mapper.primary_key] class ListField(fields.List): """ Custom field for handling list of query params E.g. ?ids=foo,bar """ def _deserialize(self, value: str, attr: str, data: dict, **kwargs) -> List[Union[str, int]]: if value and isinstance(value, list) and len(value) == 1: value = value[0].split(",") return super()._deserialize(value, attr, data) class HealthCheckParams(QueryParams): """Health check endpoint params.""" include = ListField(fields.String(validate=validate.OneOf(Service.ALL_VALUES)), missing=[]) class ArgsListField(fields.List): """ Class for getting data for list field from query string parameters. """ def _deserialize(self, value, attr, data): query_value = request.args.getlist(attr) value = query_value if query_value else value return super()._deserialize(value, attr, data) class CompactMixin: """Mixin for adding compact response flag to query params schema.""" compact = fields.Bool(default=False, missing=False) class CompactListResponseMixin: """Compact response mixin. Allows to dump response in compact format: [n, field_name_1, ..., field_name_n, value_11, ..., value_1n, ..., value_nn] if compact is True. To create compact format correctly inherited schema must provide missing, default values for all non required fields. """ class Meta: ordered = True def dump(self, obj, many=None, update_fields=True, compact=False, **kwargs) -> MarshalResult: """Dump method with "compact" option to pass. If compact=True compact_dump runs instead dump. :returns namedtuple of data, errors. """ if not compact: return super().dump(obj, many=many, update_fields=update_fields, **kwargs) return self.compact_dump(obj) def compact_dump(self, obj_list: Union[Generator, List]) -> MarshalResult: """Method to dump response in compact format: [n, field_name_1, ..., field_name_n, value_11, ..., value_1n, ..., value_mn]. :param obj_list - list or generator to get dumping sequence from. :returns namedtuple of data, errors. :raise ValueError if got invalid obj_list type. """ is_generator = isinstance(obj_list, GeneratorType) is_list = isinstance(obj_list, list) if not is_generator and not is_list: raise ValueError(f"{self.__class__.__name__}.compact_dump got invalid object {obj_list} type.") fields_len = len(self.fields.keys()) errors = {} result = [] if is_generator: for i, obj in enumerate(obj_list): self._compact_dump_one(obj, i, result, errors, fields_len) elif is_list: # do it that way instead "for" for immediate deleting item from obj_list after dumping to reduce memory used ind = 0 while obj_list: self._compact_dump_one(obj_list[0], ind, result, errors, fields_len) ind += 1 del obj_list[0] if not result: return MarshalResult(result, errors) # add fields length and names fist result = [fields_len] + [f.dump_to or n for n, f in self.fields.items()] + result return MarshalResult(result, errors) def _compact_dump_one(self, obj: Any, ind: int, result: List[Any], errors: Dict[int, str], fields_len: int): """Method to dump one object to compact format. :param obj - object to dump. :param ind - index of object in dumping sequence. :param result - list to store result of dumping. :param errors - dict to store errors of dumping. :param fields_len - number of values to give in response for one item. :return None. :raises ValidationError - if number of values after dump is not equal to fields_len. """ data, error = self.dump(obj) dumped_values = data.values() if len(dumped_values) != fields_len: raise ValidationError( f"{self.__class__.__name__}.dump_one got invalid data for compact response: " f"values number after dump() is not equal schema's fields number {fields_len} " f"for {ind} item: {obj}." ) if error: errors[ind] = error result.extend(dumped_values) class IsrcMixin: isrc = fields.String(validate=validate.Length(min=1, max=45)) class IsrcRequiredMixin: isrc = fields.String(validate=validate.Length(min=1, max=45), required=True) class IsrcListRequiredMixin: isrc_list = ListField( fields.String(validate=validate.Length(min=1, max=45)), load_from="isrc", required=True, validate=validate.Length(min=1, max=1000), ) class MarketDefaultMixin: market = fields.String(validate=validate.Length(max=10, min=2), missing=MARKET_GLOBAL) class MarketRequiredMixin: market = fields.String(validate=validate.Length(max=10, min=2), required=True) class VendorRequiredMixin: vendor = fields.String(validate=validate.OneOf(VENDORS), required=True) def get_pagination_mixin(page_offset: int = 0, page_limit: int = 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 SearchParamMixin: search = fields.String(validate=validate.Length(min=2, max=50), required=False, missing=None)