Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
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 
codekingpro/portable-devtools · Team Ai