Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
undefined.py281 linesDownload Raw Back to dataclasses_json
1import abc2import dataclasses3import functools4import inspect5import sys6from dataclasses import Field, fields7from typing import Any, Callable, Dict, Optional, Tuple, Union, Type, get_type_hints8from enum import Enum9 10from marshmallow.exceptions import ValidationError  # type: ignore11 12from dataclasses_json.utils import CatchAllVar13 14KnownParameters = Dict[str, Any]15UnknownParameters = Dict[str, Any]16 17 18class _UndefinedParameterAction(abc.ABC):19    @staticmethod20    @abc.abstractmethod21    def handle_from_dict(cls, kvs: Dict[Any, Any]) -> Dict[str, Any]:22        """23        Return the parameters to initialize the class with.24        """25        pass26 27    @staticmethod28    def handle_to_dict(obj, kvs: Dict[Any, Any]) -> Dict[Any, Any]:29        """30        Return the parameters that will be written to the output dict31        """32        return kvs33 34    @staticmethod35    def handle_dump(obj) -> Dict[Any, Any]:36        """37        Return the parameters that will be added to the schema dump.38        """39        return {}40 41    @staticmethod42    def create_init(obj) -> Callable:43        return obj.__init__44 45    @staticmethod46    def _separate_defined_undefined_kvs(cls, kvs: Dict) -> \47            Tuple[KnownParameters, UnknownParameters]:48        """49        Returns a 2 dictionaries: defined and undefined parameters50        """51        class_fields = fields(cls)52        field_names = [field.name for field in class_fields]53        unknown_given_parameters = {k: v for k, v in kvs.items() if54                                    k not in field_names}55        known_given_parameters = {k: v for k, v in kvs.items() if56                                  k in field_names}57        return known_given_parameters, unknown_given_parameters58 59 60class _RaiseUndefinedParameters(_UndefinedParameterAction):61    """62    This action raises UndefinedParameterError if it encounters an undefined63    parameter during initialization.64    """65 66    @staticmethod67    def handle_from_dict(cls, kvs: Dict) -> Dict[str, Any]:68        known, unknown = \69            _UndefinedParameterAction._separate_defined_undefined_kvs(70                cls=cls, kvs=kvs)71        if len(unknown) > 0:72            raise UndefinedParameterError(73                f"Received undefined initialization arguments {unknown}")74        return known75 76 77CatchAll = Optional[CatchAllVar]78 79 80class _IgnoreUndefinedParameters(_UndefinedParameterAction):81    """82    This action does nothing when it encounters undefined parameters.83    The undefined parameters can not be retrieved after the class has been84    created.85    """86 87    @staticmethod88    def handle_from_dict(cls, kvs: Dict) -> Dict[str, Any]:89        known_given_parameters, _ = \90            _UndefinedParameterAction._separate_defined_undefined_kvs(91                cls=cls, kvs=kvs)92        return known_given_parameters93 94    @staticmethod95    def create_init(obj) -> Callable:96        original_init = obj.__init__97        init_signature = inspect.signature(original_init)98 99        @functools.wraps(obj.__init__)100        def _ignore_init(self, *args, **kwargs):101            known_kwargs, _ = \102                _CatchAllUndefinedParameters._separate_defined_undefined_kvs(103                    obj, kwargs)104            num_params_takeable = len(105                init_signature.parameters) - 1  # don't count self106            num_args_takeable = num_params_takeable - len(known_kwargs)107 108            args = args[:num_args_takeable]109            bound_parameters = init_signature.bind_partial(self, *args,110                                                           **known_kwargs)111            bound_parameters.apply_defaults()112 113            arguments = bound_parameters.arguments114            arguments.pop("self", None)115            final_parameters = \116                _IgnoreUndefinedParameters.handle_from_dict(obj, arguments)117            original_init(self, **final_parameters)118 119        return _ignore_init120 121 122class _CatchAllUndefinedParameters(_UndefinedParameterAction):123    """124    This class allows to add a field of type utils.CatchAll which acts as a125    dictionary into which all126    undefined parameters will be written.127    These parameters are not affected by LetterCase.128    If no undefined parameters are given, this dictionary will be empty.129    """130 131    class _SentinelNoDefault:132        pass133 134    @staticmethod135    def handle_from_dict(cls, kvs: Dict) -> Dict[str, Any]:136        known, unknown = _UndefinedParameterAction \137            ._separate_defined_undefined_kvs(cls=cls, kvs=kvs)138        catch_all_field = _CatchAllUndefinedParameters._get_catch_all_field(139            cls=cls)140 141        if catch_all_field.name in known:142 143            already_parsed = isinstance(known[catch_all_field.name], dict)144            default_value = _CatchAllUndefinedParameters._get_default(145                catch_all_field=catch_all_field)146            received_default = default_value == known[catch_all_field.name]147 148            value_to_write: Any149            if received_default and len(unknown) == 0:150                value_to_write = default_value151            elif received_default and len(unknown) > 0:152                value_to_write = unknown153            elif already_parsed:154                # Did not receive default155                value_to_write = known[catch_all_field.name]156                if len(unknown) > 0:157                    value_to_write.update(unknown)158            else:159                error_message = f"Received input field with " \160                                f"same name as catch-all field: " \161                                f"'{catch_all_field.name}': " \162                                f"'{known[catch_all_field.name]}'"163                raise UndefinedParameterError(error_message)164        else:165            value_to_write = unknown166 167        known[catch_all_field.name] = value_to_write168        return known169 170    @staticmethod171    def _get_default(catch_all_field: Field) -> Any:172        # access to the default factory currently causes173        # a false-positive mypy error (16. Dec 2019):174        # https://github.com/python/mypy/issues/6910175 176        # noinspection PyProtectedMember177        has_default = not isinstance(catch_all_field.default,178                                     dataclasses._MISSING_TYPE)179        # noinspection PyProtectedMember180        has_default_factory = not isinstance(catch_all_field.default_factory,181                                             # type: ignore182                                             dataclasses._MISSING_TYPE)183        # TODO: black this for proper formatting184        default_value: Union[185            Type[_CatchAllUndefinedParameters._SentinelNoDefault], Any] = _CatchAllUndefinedParameters\186            ._SentinelNoDefault187 188        if has_default:189            default_value = catch_all_field.default190        elif has_default_factory:191            # This might be unwanted if the default factory constructs192            # something expensive,193            # because we have to construct it again just for this test194            default_value = catch_all_field.default_factory()  # type: ignore195 196        return default_value197 198    @staticmethod199    def handle_to_dict(obj, kvs: Dict[Any, Any]) -> Dict[Any, Any]:200        catch_all_field = \201            _CatchAllUndefinedParameters._get_catch_all_field(obj.__class__)202        undefined_parameters = kvs.pop(catch_all_field.name)203        if isinstance(undefined_parameters, dict):204            kvs.update(205                undefined_parameters)  # If desired handle letter case here206        return kvs207 208    @staticmethod209    def handle_dump(obj) -> Dict[Any, Any]:210        catch_all_field = _CatchAllUndefinedParameters._get_catch_all_field(211            cls=obj)212        return getattr(obj, catch_all_field.name)213 214    @staticmethod215    def create_init(obj) -> Callable:216        original_init = obj.__init__217        init_signature = inspect.signature(original_init)218 219        @functools.wraps(obj.__init__)220        def _catch_all_init(self, *args, **kwargs):221            known_kwargs, unknown_kwargs = \222                _CatchAllUndefinedParameters._separate_defined_undefined_kvs(223                    obj, kwargs)224            num_params_takeable = len(225                init_signature.parameters) - 1  # don't count self226            if _CatchAllUndefinedParameters._get_catch_all_field(227                    obj).name not in known_kwargs:228                num_params_takeable -= 1229            num_args_takeable = num_params_takeable - len(known_kwargs)230 231            args, unknown_args = args[:num_args_takeable], args[232                                                           num_args_takeable:]233            bound_parameters = init_signature.bind_partial(self, *args,234                                                           **known_kwargs)235 236            unknown_args = {f"_UNKNOWN{i}": v for i, v in237                            enumerate(unknown_args)}238            arguments = bound_parameters.arguments239            arguments.update(unknown_args)240            arguments.update(unknown_kwargs)241            arguments.pop("self", None)242            final_parameters = _CatchAllUndefinedParameters.handle_from_dict(243                obj, arguments)244            original_init(self, **final_parameters)245 246        return _catch_all_init247 248    @staticmethod249    def _get_catch_all_field(cls) -> Field:250        cls_globals = vars(sys.modules[cls.__module__])251        types = get_type_hints(cls, globalns=cls_globals)252        catch_all_fields = list(253            filter(lambda f: types[f.name] == Optional[CatchAllVar], fields(cls)))254        number_of_catch_all_fields = len(catch_all_fields)255        if number_of_catch_all_fields == 0:256            raise UndefinedParameterError(257                "No field of type dataclasses_json.CatchAll defined")258        elif number_of_catch_all_fields > 1:259            raise UndefinedParameterError(260                f"Multiple catch-all fields supplied: "261                f"{number_of_catch_all_fields}.")262        else:263            return catch_all_fields[0]264 265 266class Undefined(Enum):267    """268    Choose the behavior what happens when an undefined parameter is encountered269    during class initialization.270    """271    INCLUDE = _CatchAllUndefinedParameters272    RAISE = _RaiseUndefinedParameters273    EXCLUDE = _IgnoreUndefinedParameters274 275 276class UndefinedParameterError(ValidationError):277    """278    Raised when something has gone wrong handling undefined parameters.279    """280    pass281 
codekingpro/portable-devtools · Team Ai