import re from abc import abstractmethod from typing import List, Any from marshmallow import fields, ValidationError from apollo_utils.core.utils.misc import is_collection class BaseUniqueArgsListField(fields.List): """Base Class for getting data for list field from query string parameters, by default removes duplicates.""" def __init__(self, *args, **kwargs): self._unify = kwargs.pop("unify", True) super().__init__(*args, **kwargs) def _deserialize(self, value, attr, data, *args, **kwargs): query_value = self._get_attr(attr) value = query_value if query_value else value value = super()._deserialize(value, attr, data) if value and self._unify: value = sorted(list(set(value)), key=value.index) return value @abstractmethod def _get_attr(self, attr): """ Flask example: request.args.getlist(attr) """ raise NotImplementedError class SplitUniqueListField(fields.List): """List field to handle list values, and strings with comma-separated parameters. Examples: ?id=1&id=2 or ?id=1,2. Field also unifies elements of the passed list if 'unify' parameter is True (default value). """ def __init__(self, *args, **kwargs): self._split = kwargs.pop("split", True) self._unify = kwargs.pop("unify", True) self._value_list = kwargs.pop("value_list", ()) self._value_list_invalid_error = kwargs.pop("value_list_invalid_error", False) self._filter_mask = kwargs.pop("filter_mask", None) self._force_list = kwargs.pop("force_list", False) super().__init__(*args, **kwargs) def _do_split(self, value): if value and isinstance(value, list) and len(value) == 1 and isinstance(value[0], str): return value[0].split(",") return value def _do_unify(self, value): if value: return sorted(list(set(value)), key=value.index) return value def _do_value_list(self, value, attr): if value and self._value_list and is_collection(value): value = [i for i in value if i in self._value_list] if not value and self._value_list_invalid_error: raise ValidationError(f"{attr} values are incorrect, please use at least one of {self._value_list}") return value def _deserialize(self, value, attr, data, **kwargs) -> List[Any]: if self._force_list and value and not isinstance(value, list): value = [value] if self._split: value = self._do_split(value) if self._value_list: value = self._do_value_list(value, attr) value = super()._deserialize(value, attr, data) if self._filter_mask: value = list(i for i in value if re.match(self._filter_mask, i)) if not value and self.required: raise ValidationError(f"Missing data for required field {attr}") if self._unify: value = self._do_unify(value) return value