codekingpro/portable-devtools
115k
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 