Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
core.py476 linesDownload Raw Back to dataclasses_json
1import copy2import json3import sys4import warnings5from collections import defaultdict, namedtuple6from collections.abc import (Collection as ABCCollection, Mapping as ABCMapping, MutableMapping, MutableSequence,7                             MutableSet, Sequence, Set)8from dataclasses import (MISSING,9                         fields,10                         is_dataclass  # type: ignore11                         )12from datetime import datetime, timezone13from decimal import Decimal14from enum import Enum15from types import MappingProxyType16from typing import (Any, Collection, Mapping, Union, get_type_hints,17                    Tuple, TypeVar, Type)18from uuid import UUID19 20from typing_inspect import is_union_type  # type: ignore21 22from dataclasses_json import cfg23from dataclasses_json.utils import (_get_type_cons, _get_type_origin,24                                    _handle_undefined_parameters_safe,25                                    _is_collection, _is_mapping, _is_new_type,26                                    _is_optional, _isinstance_safe,27                                    _get_type_arg_param,28                                    _get_type_args, _is_counter,29                                    _NO_ARGS,30                                    _issubclass_safe, _is_tuple,31                                    _is_generic_dataclass)32 33Json = Union[dict, list, str, int, float, bool, None]34 35confs = ['encoder', 'decoder', 'mm_field', 'letter_case', 'exclude']36FieldOverride = namedtuple('FieldOverride', confs)  # type: ignore37collections_abc_type_to_implementation_type = MappingProxyType({38    ABCCollection: tuple,39    ABCMapping: dict,40    MutableMapping: dict,41    MutableSequence: list,42    MutableSet: set,43    Sequence: tuple,44    Set: frozenset,45})46 47 48class _ExtendedEncoder(json.JSONEncoder):49    def default(self, o) -> Json:50        result: Json51        if _isinstance_safe(o, Collection):52            if _isinstance_safe(o, Mapping):53                result = dict(o)54            else:55                result = list(o)56        elif _isinstance_safe(o, datetime):57            result = o.timestamp()58        elif _isinstance_safe(o, UUID):59            result = str(o)60        elif _isinstance_safe(o, Enum):61            result = o.value62        elif _isinstance_safe(o, Decimal):63            result = str(o)64        else:65            result = json.JSONEncoder.default(self, o)66        return result67 68 69def _user_overrides_or_exts(cls):70    global_metadata = defaultdict(dict)71    encoders = cfg.global_config.encoders72    decoders = cfg.global_config.decoders73    mm_fields = cfg.global_config.mm_fields74    for field in fields(cls):75        if field.type in encoders:76            global_metadata[field.name]['encoder'] = encoders[field.type]77        if field.type in decoders:78            global_metadata[field.name]['decoder'] = decoders[field.type]79        if field.type in mm_fields:80            global_metadata[field.name]['mm_field'] = mm_fields[field.type]81    try:82        cls_config = (cls.dataclass_json_config83                      if cls.dataclass_json_config is not None else {})84    except AttributeError:85        cls_config = {}86 87    overrides = {}88    for field in fields(cls):89        field_config = {}90        # first apply global overrides or extensions91        field_metadata = global_metadata[field.name]92        if 'encoder' in field_metadata:93            field_config['encoder'] = field_metadata['encoder']94        if 'decoder' in field_metadata:95            field_config['decoder'] = field_metadata['decoder']96        if 'mm_field' in field_metadata:97            field_config['mm_field'] = field_metadata['mm_field']98        # then apply class-level overrides or extensions99        field_config.update(cls_config)100        # last apply field-level overrides or extensions101        field_config.update(field.metadata.get('dataclasses_json', {}))102        overrides[field.name] = FieldOverride(*map(field_config.get, confs))103    return overrides104 105 106def _encode_json_type(value, default=_ExtendedEncoder().default):107    if isinstance(value, Json.__args__):  # type: ignore108        if isinstance(value, list):109            return [_encode_json_type(i) for i in value]110        elif isinstance(value, dict):111            return {k: _encode_json_type(v) for k, v in value.items()}112        else:113            return value114    return default(value)115 116 117def _encode_overrides(kvs, overrides, encode_json=False):118    override_kvs = {}119    for k, v in kvs.items():120        if k in overrides:121            exclude = overrides[k].exclude122            # If the exclude predicate returns true, the key should be123            #  excluded from encoding, so skip the rest of the loop124            if exclude and exclude(v):125                continue126            letter_case = overrides[k].letter_case127            original_key = k128            k = letter_case(k) if letter_case is not None else k129            if k in override_kvs:130                raise ValueError(131                    f"Multiple fields map to the same JSON "132                    f"key after letter case encoding: {k}"133                )134 135            encoder = overrides[original_key].encoder136            v = encoder(v) if encoder is not None else v137 138        if encode_json:139            v = _encode_json_type(v)140        override_kvs[k] = v141    return override_kvs142 143 144def _decode_letter_case_overrides(field_names, overrides):145    """Override letter case of field names for encode/decode"""146    names = {}147    for field_name in field_names:148        field_override = overrides.get(field_name)149        if field_override is not None:150            letter_case = field_override.letter_case151            if letter_case is not None:152                names[letter_case(field_name)] = field_name153    return names154 155 156def _decode_dataclass(cls, kvs, infer_missing):157    if _isinstance_safe(kvs, cls):158        return kvs159    overrides = _user_overrides_or_exts(cls)160    kvs = {} if kvs is None and infer_missing else kvs161    field_names = [field.name for field in fields(cls)]162    decode_names = _decode_letter_case_overrides(field_names, overrides)163    kvs = {decode_names.get(k, k): v for k, v in kvs.items()}164    missing_fields = {field for field in fields(cls) if field.name not in kvs}165 166    for field in missing_fields:167        if field.default is not MISSING:168            kvs[field.name] = field.default169        elif field.default_factory is not MISSING:170            kvs[field.name] = field.default_factory()171        elif infer_missing:172            kvs[field.name] = None173 174    # Perform undefined parameter action175    kvs = _handle_undefined_parameters_safe(cls, kvs, usage="from")176 177    init_kwargs = {}178    types = get_type_hints(cls)179    for field in fields(cls):180        # The field should be skipped from being added181        # to init_kwargs as it's not intended as a constructor argument.182        if not field.init:183            continue184 185        field_value = kvs[field.name]186        field_type = types[field.name]187        if field_value is None:188            if not _is_optional(field_type):189                warning = (190                    f"value of non-optional type {field.name} detected "191                    f"when decoding {cls.__name__}"192                )193                if infer_missing:194                    warnings.warn(195                        f"Missing {warning} and was defaulted to None by "196                        f"infer_missing=True. "197                        f"Set infer_missing=False (the default) to prevent "198                        f"this behavior.", RuntimeWarning199                    )200                else:201                    warnings.warn(202                        f"'NoneType' object {warning}.", RuntimeWarning203                    )204            init_kwargs[field.name] = field_value205            continue206 207        while True:208            if not _is_new_type(field_type):209                break210 211            field_type = field_type.__supertype__212 213        if (field.name in overrides214                and overrides[field.name].decoder is not None):215            # FIXME hack216            if field_type is type(field_value):217                init_kwargs[field.name] = field_value218            else:219                init_kwargs[field.name] = overrides[field.name].decoder(220                    field_value)221        elif is_dataclass(field_type):222            # FIXME this is a band-aid to deal with the value already being223            # serialized when handling nested marshmallow schema224            # proper fix is to investigate the marshmallow schema generation225            # code226            if is_dataclass(field_value):227                value = field_value228            else:229                value = _decode_dataclass(field_type, field_value,230                                          infer_missing)231            init_kwargs[field.name] = value232        elif _is_supported_generic(field_type) and field_type != str:233            init_kwargs[field.name] = _decode_generic(field_type,234                                                      field_value,235                                                      infer_missing)236        else:237            init_kwargs[field.name] = _support_extended_types(field_type,238                                                              field_value)239 240    return cls(**init_kwargs)241 242 243def _decode_type(type_, value, infer_missing):244    if _has_decoder_in_global_config(type_):245        return _get_decoder_in_global_config(type_)(value)246    if _is_supported_generic(type_):247        return _decode_generic(type_, value, infer_missing)248    if is_dataclass(type_) or is_dataclass(value):249        return _decode_dataclass(type_, value, infer_missing)250    return _support_extended_types(type_, value)251 252 253def _support_extended_types(field_type, field_value):254    if _issubclass_safe(field_type, datetime):255        # FIXME this is a hack to deal with mm already decoding256        # the issue is we want to leverage mm fields' missing argument257        # but need this for the object creation hook258        if isinstance(field_value, datetime):259            res = field_value260        else:261            tz = datetime.now(timezone.utc).astimezone().tzinfo262            res = datetime.fromtimestamp(field_value, tz=tz)263    elif _issubclass_safe(field_type, Decimal):264        res = (field_value265               if isinstance(field_value, Decimal)266               else Decimal(field_value))267    elif _issubclass_safe(field_type, UUID):268        res = (field_value269               if isinstance(field_value, UUID)270               else UUID(field_value))271    elif _issubclass_safe(field_type, (int, float, str, bool)):272        res = (field_value273               if isinstance(field_value, field_type)274               else field_type(field_value))275    else:276        res = field_value277    return res278 279 280def _is_supported_generic(type_):281    if type_ is _NO_ARGS:282        return False283    not_str = not _issubclass_safe(type_, str)284    is_enum = _issubclass_safe(type_, Enum)285    is_generic_dataclass = _is_generic_dataclass(type_)286    return (not_str and _is_collection(type_)) or _is_optional(287        type_) or is_union_type(type_) or is_enum or is_generic_dataclass288 289 290def _decode_generic(type_, value, infer_missing):291    if value is None:292        res = value293    elif _issubclass_safe(type_, Enum):294        # Convert to an Enum using the type as a constructor.295        # Assumes a direct match is found.296        res = type_(value)297    # FIXME this is a hack to fix a deeper underlying issue. A refactor is due.298    elif _is_collection(type_):299        if _is_mapping(type_) and not _is_counter(type_):300            k_type, v_type = _get_type_args(type_, (Any, Any))301            # a mapping type has `.keys()` and `.values()`302            # (see collections.abc)303            ks = _decode_dict_keys(k_type, value.keys(), infer_missing)304            vs = _decode_items(v_type, value.values(), infer_missing)305            xs = zip(ks, vs)306        elif _is_tuple(type_):307            types = _get_type_args(type_)308            if Ellipsis in types:309                xs = _decode_items(types[0], value, infer_missing)310            else:311                xs = _decode_items(_get_type_args(type_) or _NO_ARGS, value, infer_missing)312        elif _is_counter(type_):313            xs = dict(zip(_decode_items(_get_type_arg_param(type_, 0), value.keys(), infer_missing), value.values()))314        else:315            xs = _decode_items(_get_type_arg_param(type_, 0), value, infer_missing)316 317        collection_type = _resolve_collection_type_to_decode_to(type_)318        res = collection_type(xs)319    elif _is_generic_dataclass(type_):320        origin = _get_type_origin(type_)321        res = _decode_dataclass(origin, value, infer_missing)322    else:  # Optional or Union323        _args = _get_type_args(type_)324        if _args is _NO_ARGS:325            # Any, just accept326            res = value327        elif _is_optional(type_) and len(_args) == 2:  # Optional328            type_arg = _get_type_arg_param(type_, 0)329            res = _decode_type(type_arg, value, infer_missing)330        else:  # Union (already decoded or try to decode a dataclass)331            type_options = _get_type_args(type_)332            res = value  # assume already decoded333            if type(value) is dict and dict not in type_options:334                for type_option in type_options:335                    if is_dataclass(type_option):336                        try:337                            res = _decode_dataclass(type_option, value, infer_missing)338                            break339                        except (KeyError, ValueError, AttributeError):340                            continue341                if res == value:342                    warnings.warn(343                        f"Failed to decode {value} Union dataclasses."344                        f"Expected Union to include a matching dataclass and it didn't."345                    )346    return res347 348 349def _decode_dict_keys(key_type, xs, infer_missing):350    """351    Because JSON object keys must be strs, we need the extra step of decoding352    them back into the user's chosen python type353    """354    decode_function = key_type355    # handle NoneType keys... it's weird to type a Dict as NoneType keys356    # but it's valid...357    # Issue #341 and PR #346:358    #   This is a special case for Python 3.7 and Python 3.8.359    #   By some reason, "unbound" dicts are counted360    #   as having key type parameter to be TypeVar('KT')361    if key_type is None or key_type == Any or isinstance(key_type, TypeVar):362        decode_function = key_type = (lambda x: x)363    # handle a nested python dict that has tuples for keys. E.g. for364    # Dict[Tuple[int], int], key_type will be typing.Tuple[int], but365    # decode_function should be tuple, so map() doesn't break.366    #367    # Note: _get_type_origin() will return typing.Tuple for python368    # 3.6 and tuple for 3.7 and higher.369    elif _get_type_origin(key_type) in {tuple, Tuple}:370        decode_function = tuple371        key_type = key_type372 373    return map(decode_function, _decode_items(key_type, xs, infer_missing))374 375 376def _decode_items(type_args, xs, infer_missing):377    """378    This is a tricky situation where we need to check both the annotated379    type info (which is usually a type from `typing`) and check the380    value's type directly using `type()`.381 382    If the type_arg is a generic we can use the annotated type, but if the383    type_arg is a typevar we need to extract the reified type information384    hence the check of `is_dataclass(vs)`385    """386    def handle_pep0673(pre_0673_hint: str) -> Union[Type, str]:387        for module in sys.modules.values():388            if hasattr(module, type_args):389                maybe_resolved = getattr(module, type_args)390                warnings.warn(f"Assuming hint {pre_0673_hint} resolves to {maybe_resolved} "391                              "This is not necessarily the value that is in-scope.")392                return maybe_resolved393 394        warnings.warn(f"Could not resolve self-reference for type {pre_0673_hint}, "395                      f"decoded type might be incorrect or decode might fail altogether.")396        return pre_0673_hint397 398    # Before https://peps.python.org/pep-0673 (3.11+) self-type hints are simply strings399    if sys.version_info.minor < 11 and type_args is not type and type(type_args) is str:400        type_args = handle_pep0673(type_args)401 402    if _isinstance_safe(type_args, Collection) and not _issubclass_safe(type_args, Enum):403        if len(type_args) == len(xs):404            return list(_decode_type(type_arg, x, infer_missing) for type_arg, x in zip(type_args, xs))405        else:406            raise TypeError(f"Number of types specified in the collection type {str(type_args)} "407                            f"does not match number of elements in the collection. In case you are working with tuples"408                            f"take a look at this document "409                            f"docs.python.org/3/library/typing.html#annotating-tuples.")410    return list(_decode_type(type_args, x, infer_missing) for x in xs)411 412 413def _resolve_collection_type_to_decode_to(type_):414    # get the constructor if using corresponding generic type in `typing`415    # otherwise fallback on constructing using type_ itself416    try:417        collection_type = _get_type_cons(type_)418    except (TypeError, AttributeError):419        collection_type = type_420 421    # map abstract collection to concrete implementation422    return collections_abc_type_to_implementation_type.get(collection_type, collection_type)423 424 425def _asdict(obj, encode_json=False):426    """427    A re-implementation of `asdict` (based on the original in the `dataclasses`428    source) to support arbitrary Collection and Mapping types.429    """430    if is_dataclass(obj):431        result = []432        overrides = _user_overrides_or_exts(obj)433        for field in fields(obj):434            if overrides[field.name].encoder:435                value = getattr(obj, field.name)436            else:437                value = _asdict(438                    getattr(obj, field.name),439                    encode_json=encode_json440                )441            result.append((field.name, value))442 443        result = _handle_undefined_parameters_safe(cls=obj, kvs=dict(result),444                                                   usage="to")445        return _encode_overrides(dict(result), _user_overrides_or_exts(obj),446                                 encode_json=encode_json)447    elif isinstance(obj, Mapping):448        return dict((_asdict(k, encode_json=encode_json),449                     _asdict(v, encode_json=encode_json)) for k, v in450                    obj.items())451    # enum.IntFlag and enum.Flag are regarded as collections in Python 3.11, thus a check against Enum is needed452    elif isinstance(obj, Collection) and not isinstance(obj, (str, bytes, Enum)):453        return list(_asdict(v, encode_json=encode_json) for v in obj)454    # encoding of generics primarily relies on concrete types while decoding relies on type annotations. This makes455    # applying encoders/decoders from global configuration inconsistent.456    elif _has_encoder_in_global_config(type(obj)):457        return _get_encoder_in_global_config(type(obj))(obj)458    else:459        return copy.deepcopy(obj)460 461 462def _has_decoder_in_global_config(type_):463    return type_ in cfg.global_config.decoders464 465 466def _get_decoder_in_global_config(type_):467    return cfg.global_config.decoders[type_]468 469 470def _has_encoder_in_global_config(type_):471    return type_ in cfg.global_config.encoders472 473 474def _get_encoder_in_global_config(type_):475    return cfg.global_config.encoders[type_]476 
codekingpro/portable-devtools · Team Ai