codekingpro/portable-devtools
114k
1# flake8: noqa2 3import typing4import warnings5import sys6from copy import deepcopy7 8from dataclasses import MISSING, is_dataclass, fields as dc_fields9from datetime import datetime10from decimal import Decimal11from uuid import UUID12from enum import Enum13 14from typing_inspect import is_union_type # type: ignore15 16from marshmallow import fields, Schema, post_load # type: ignore17from marshmallow.exceptions import ValidationError # type: ignore18 19from dataclasses_json.core import (_is_supported_generic, _decode_dataclass,20 _ExtendedEncoder, _user_overrides_or_exts)21from dataclasses_json.utils import (_is_collection, _is_optional,22 _issubclass_safe, _timestamp_to_dt_aware,23 _is_new_type, _get_type_origin,24 _handle_undefined_parameters_safe,25 CatchAllVar)26 27 28class _TimestampField(fields.Field):29 def _serialize(self, value, attr, obj, **kwargs):30 if value is not None:31 return value.timestamp()32 else:33 if not self.required:34 return None35 else:36 raise ValidationError(self.default_error_messages["required"])37 38 def _deserialize(self, value, attr, data, **kwargs):39 if value is not None:40 return _timestamp_to_dt_aware(value)41 else:42 if not self.required:43 return None44 else:45 raise ValidationError(self.default_error_messages["required"])46 47 48class _IsoField(fields.Field):49 def _serialize(self, value, attr, obj, **kwargs):50 if value is not None:51 return value.isoformat()52 else:53 if not self.required:54 return None55 else:56 raise ValidationError(self.default_error_messages["required"])57 58 def _deserialize(self, value, attr, data, **kwargs):59 if value is not None:60 return datetime.fromisoformat(value)61 else:62 if not self.required:63 return None64 else:65 raise ValidationError(self.default_error_messages["required"])66 67 68class _UnionField(fields.Field):69 def __init__(self, desc, cls, field, *args, **kwargs):70 self.desc = desc71 self.cls = cls72 self.field = field73 super().__init__(*args, **kwargs)74 75 def _serialize(self, value, attr, obj, **kwargs):76 if self.allow_none and value is None:77 return None78 for type_, schema_ in self.desc.items():79 if _issubclass_safe(type(value), type_):80 if is_dataclass(value):81 res = schema_._serialize(value, attr, obj, **kwargs)82 res['__type'] = str(type_.__name__)83 return res84 break85 elif isinstance(value, _get_type_origin(type_)):86 return schema_._serialize(value, attr, obj, **kwargs)87 else:88 warnings.warn(89 f'The type "{type(value).__name__}" (value: "{value}") '90 f'is not in the list of possible types of typing.Union '91 f'(dataclass: {self.cls.__name__}, field: {self.field.name}). '92 f'Value cannot be serialized properly.')93 return super()._serialize(value, attr, obj, **kwargs)94 95 def _deserialize(self, value, attr, data, **kwargs):96 tmp_value = deepcopy(value)97 if isinstance(tmp_value, dict) and '__type' in tmp_value:98 dc_name = tmp_value['__type']99 for type_, schema_ in self.desc.items():100 if is_dataclass(type_) and type_.__name__ == dc_name:101 del tmp_value['__type']102 return schema_._deserialize(tmp_value, attr, data, **kwargs)103 elif isinstance(tmp_value, dict):104 warnings.warn(105 f'Attempting to deserialize "dict" (value: "{tmp_value}) '106 f'that does not have a "__type" type specifier field into'107 f'(dataclass: {self.cls.__name__}, field: {self.field.name}).'108 f'Deserialization may fail, or deserialization to wrong type may occur.'109 )110 return super()._deserialize(tmp_value, attr, data, **kwargs)111 else:112 for type_, schema_ in self.desc.items():113 if isinstance(tmp_value, _get_type_origin(type_)):114 return schema_._deserialize(tmp_value, attr, data, **kwargs)115 else:116 warnings.warn(117 f'The type "{type(tmp_value).__name__}" (value: "{tmp_value}") '118 f'is not in the list of possible types of typing.Union '119 f'(dataclass: {self.cls.__name__}, field: {self.field.name}). '120 f'Value cannot be deserialized properly.')121 return super()._deserialize(tmp_value, attr, data, **kwargs)122 123 124class _TupleVarLen(fields.List):125 """126 variable-length homogeneous tuples127 """128 def _deserialize(self, value, attr, data, **kwargs):129 optional_list = super()._deserialize(value, attr, data, **kwargs)130 return None if optional_list is None else tuple(optional_list)131 132 133TYPES = {134 typing.Mapping: fields.Mapping,135 typing.MutableMapping: fields.Mapping,136 typing.List: fields.List,137 typing.Dict: fields.Dict,138 typing.Tuple: fields.Tuple,139 typing.Callable: fields.Function,140 typing.Any: fields.Raw,141 dict: fields.Dict,142 list: fields.List,143 tuple: fields.Tuple,144 str: fields.Str,145 int: fields.Int,146 float: fields.Float,147 bool: fields.Bool,148 datetime: _TimestampField,149 UUID: fields.UUID,150 Decimal: fields.Decimal,151 CatchAllVar: fields.Dict,152}153 154A = typing.TypeVar('A')155JsonData = typing.Union[str, bytes, bytearray]156TEncoded = typing.Dict[str, typing.Any]157TOneOrMulti = typing.Union[typing.List[A], A]158TOneOrMultiEncoded = typing.Union[typing.List[TEncoded], TEncoded]159 160if sys.version_info >= (3, 7) or typing.TYPE_CHECKING:161 class SchemaF(Schema, typing.Generic[A]):162 """Lift Schema into a type constructor"""163 164 def __init__(self, *args, **kwargs):165 """166 Raises exception because this class should not be inherited.167 This class is helper only.168 """169 170 super().__init__(*args, **kwargs)171 raise NotImplementedError()172 173 @typing.overload174 def dump(self, obj: typing.List[A], many: typing.Optional[bool] = None) -> typing.List[TEncoded]: # type: ignore175 # mm has the wrong return type annotation (dict) so we can ignore the mypy error176 pass177 178 @typing.overload179 def dump(self, obj: A, many: typing.Optional[bool] = None) -> TEncoded:180 pass181 182 def dump(self, obj: TOneOrMulti, # type: ignore183 many: typing.Optional[bool] = None) -> TOneOrMultiEncoded:184 pass185 186 @typing.overload187 def dumps(self, obj: typing.List[A], many: typing.Optional[bool] = None, *args,188 **kwargs) -> str:189 pass190 191 @typing.overload192 def dumps(self, obj: A, many: typing.Optional[bool] = None, *args, **kwargs) -> str:193 pass194 195 def dumps(self, obj: TOneOrMulti, many: typing.Optional[bool] = None, *args, # type: ignore196 **kwargs) -> str:197 pass198 199 @typing.overload # type: ignore200 def load(self, data: typing.List[TEncoded],201 many: bool = True, partial: typing.Optional[bool] = None,202 unknown: typing.Optional[str] = None) -> \203 typing.List[A]:204 # ignore the mypy error of the decorator because mm does not define lists as an allowed input type205 pass206 207 @typing.overload208 def load(self, data: TEncoded,209 many: None = None, partial: typing.Optional[bool] = None,210 unknown: typing.Optional[str] = None) -> A:211 pass212 213 def load(self, data: TOneOrMultiEncoded,214 many: typing.Optional[bool] = None, partial: typing.Optional[bool] = None,215 unknown: typing.Optional[str] = None) -> TOneOrMulti:216 pass217 218 @typing.overload # type: ignore219 def loads(self, json_data: JsonData, # type: ignore220 many: typing.Optional[bool] = True, partial: typing.Optional[bool] = None, unknown: typing.Optional[str] = None,221 **kwargs) -> typing.List[A]:222 # ignore the mypy error of the decorator because mm does not define bytes as correct input data223 # mm has the wrong return type annotation (dict) so we can ignore the mypy error224 # for the return type overlap225 pass226 227 def loads(self, json_data: JsonData,228 many: typing.Optional[bool] = None, partial: typing.Optional[bool] = None, unknown: typing.Optional[str] = None,229 **kwargs) -> TOneOrMulti:230 pass231 232 233 SchemaType = SchemaF[A]234else:235 SchemaType = Schema236 237 238def build_type(type_, options, mixin, field, cls):239 def inner(type_, options):240 while True:241 if not _is_new_type(type_):242 break243 244 type_ = type_.__supertype__245 246 if is_dataclass(type_):247 if _issubclass_safe(type_, mixin):248 options['field_many'] = bool(249 _is_supported_generic(field.type) and _is_collection(250 field.type))251 return fields.Nested(type_.schema(), **options)252 else:253 warnings.warn(f"Nested dataclass field {field.name} of type "254 f"{field.type} detected in "255 f"{cls.__name__} that is not an instance of "256 f"dataclass_json. Did you mean to recursively "257 f"serialize this field? If so, make sure to "258 f"augment {type_} with either the "259 f"`dataclass_json` decorator or mixin.")260 return fields.Field(**options)261 262 origin = getattr(type_, '__origin__', type_)263 args = [inner(a, {}) for a in getattr(type_, '__args__', []) if264 a is not type(None)]265 266 if type_ == Ellipsis:267 return type_268 269 if _is_optional(type_):270 options["allow_none"] = True271 if origin is tuple:272 if len(args) == 2 and args[1] == Ellipsis:273 return _TupleVarLen(args[0], **options)274 else:275 return fields.Tuple(args, **options)276 if origin in TYPES:277 return TYPES[origin](*args, **options)278 279 if _issubclass_safe(origin, Enum):280 return fields.Enum(enum=origin, by_value=True, *args, **options)281 282 if is_union_type(type_):283 union_types = [a for a in getattr(type_, '__args__', []) if284 a is not type(None)]285 union_desc = dict(zip(union_types, args))286 return _UnionField(union_desc, cls, field, **options)287 288 warnings.warn(289 f"Unknown type {type_} at {cls.__name__}.{field.name}: {field.type} "290 f"It's advised to pass the correct marshmallow type to `mm_field`.")291 return fields.Field(**options)292 293 return inner(type_, options)294 295 296def schema(cls, mixin, infer_missing):297 schema = {}298 overrides = _user_overrides_or_exts(cls)299 # TODO check the undefined parameters and add the proper schema action300 # https://marshmallow.readthedocs.io/en/stable/quickstart.html301 for field in dc_fields(cls):302 metadata = overrides[field.name]303 if metadata.mm_field is not None:304 schema[field.name] = metadata.mm_field305 else:306 type_ = field.type307 options: typing.Dict[str, typing.Any] = {}308 missing_key = 'missing' if infer_missing else 'default'309 if field.default is not MISSING:310 options[missing_key] = field.default311 elif field.default_factory is not MISSING:312 options[missing_key] = field.default_factory()313 else:314 options['required'] = True315 316 if options.get(missing_key, ...) is None:317 options['allow_none'] = True318 319 if _is_optional(type_):320 options.setdefault(missing_key, None)321 options['allow_none'] = True322 if len(type_.__args__) == 2:323 # Union[str, int, None] is optional too, but it has more than 1 typed field.324 type_ = [tp for tp in type_.__args__ if tp is not type(None)][0]325 326 if metadata.letter_case is not None:327 options['data_key'] = metadata.letter_case(field.name)328 329 t = build_type(type_, options, mixin, field, cls)330 if field.metadata.get('dataclasses_json', {}).get('decoder'):331 # If the field defines a custom decoder, it should completely replace the Marshmallow field's conversion332 # logic.333 # From Marshmallow's documentation for the _deserialize method:334 # "Deserialize value. Concrete :class:`Field` classes should implement this method. "335 # This is the method that Field implementations override to perform the actual deserialization logic.336 # In this case we specifically override this method instead of `deserialize` to minimize potential337 # side effects, and only cancel the actual value deserialization.338 t._deserialize = lambda v, *_a, **_kw: v339 340 # if type(t) is not fields.Field: # If we use `isinstance` we would return nothing.341 if field.type != typing.Optional[CatchAllVar]:342 schema[field.name] = t343 344 return schema345 346 347def build_schema(cls: typing.Type[A],348 mixin,349 infer_missing,350 partial) -> typing.Type["SchemaType[A]"]:351 Meta = type('Meta',352 (),353 {'fields': tuple(field.name for field in dc_fields(cls) # type: ignore354 if355 field.name != 'dataclass_json_config' and field.type !=356 typing.Optional[CatchAllVar]),357 # TODO #180358 # 'render_module': global_config.json_module359 })360 361 @post_load362 def make_instance(self, kvs, **kwargs):363 return _decode_dataclass(cls, kvs, partial)364 365 def dumps(self, *args, **kwargs):366 if 'cls' not in kwargs:367 kwargs['cls'] = _ExtendedEncoder368 369 return Schema.dumps(self, *args, **kwargs)370 371 def dump(self, obj, *, many=None):372 many = self.many if many is None else bool(many)373 dumped = Schema.dump(self, obj, many=many)374 # TODO This is hacky, but the other option I can think of is to generate a different schema375 # depending on dump and load, which is even more hacky376 377 # The only problem is the catch-all field, we can't statically create a schema for it,378 # so we just update the dumped dict379 if many:380 for i, _obj in enumerate(obj):381 dumped[i].update(382 _handle_undefined_parameters_safe(cls=_obj, kvs={},383 usage="dump"))384 else:385 dumped.update(_handle_undefined_parameters_safe(cls=obj, kvs={},386 usage="dump"))387 return dumped388 389 schema_ = schema(cls, mixin, infer_missing)390 DataClassSchema: typing.Type["SchemaType[A]"] = type(391 f'{cls.__name__.capitalize()}Schema',392 (Schema,),393 {'Meta': Meta,394 f'make_{cls.__name__.lower()}': make_instance,395 'dumps': dumps,396 'dump': dump,397 **schema_})398 399 return DataClassSchema400 