decorator.py 10 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264
  1. from functools import wraps
  2. from typing import TYPE_CHECKING, Any, Callable, Dict, List, Mapping, Optional, Tuple, Type, TypeVar, Union, overload
  3. from pydantic.v1 import validator
  4. from pydantic.v1.config import Extra
  5. from pydantic.v1.errors import ConfigError
  6. from pydantic.v1.main import BaseModel, create_model
  7. from pydantic.v1.typing import get_all_type_hints
  8. from pydantic.v1.utils import to_camel
  9. __all__ = ('validate_arguments',)
  10. if TYPE_CHECKING:
  11. from pydantic.v1.typing import AnyCallable
  12. AnyCallableT = TypeVar('AnyCallableT', bound=AnyCallable)
  13. ConfigType = Union[None, Type[Any], Dict[str, Any]]
  14. @overload
  15. def validate_arguments(func: None = None, *, config: 'ConfigType' = None) -> Callable[['AnyCallableT'], 'AnyCallableT']:
  16. ...
  17. @overload
  18. def validate_arguments(func: 'AnyCallableT') -> 'AnyCallableT':
  19. ...
  20. def validate_arguments(func: Optional['AnyCallableT'] = None, *, config: 'ConfigType' = None) -> Any:
  21. """
  22. Decorator to validate the arguments passed to a function.
  23. """
  24. def validate(_func: 'AnyCallable') -> 'AnyCallable':
  25. vd = ValidatedFunction(_func, config)
  26. @wraps(_func)
  27. def wrapper_function(*args: Any, **kwargs: Any) -> Any:
  28. return vd.call(*args, **kwargs)
  29. wrapper_function.vd = vd # type: ignore
  30. wrapper_function.validate = vd.init_model_instance # type: ignore
  31. wrapper_function.raw_function = vd.raw_function # type: ignore
  32. wrapper_function.model = vd.model # type: ignore
  33. return wrapper_function
  34. if func:
  35. return validate(func)
  36. else:
  37. return validate
  38. ALT_V_ARGS = 'v__args'
  39. ALT_V_KWARGS = 'v__kwargs'
  40. V_POSITIONAL_ONLY_NAME = 'v__positional_only'
  41. V_DUPLICATE_KWARGS = 'v__duplicate_kwargs'
  42. class ValidatedFunction:
  43. def __init__(self, function: 'AnyCallableT', config: 'ConfigType'): # noqa C901
  44. from inspect import Parameter, signature
  45. parameters: Mapping[str, Parameter] = signature(function).parameters
  46. if parameters.keys() & {ALT_V_ARGS, ALT_V_KWARGS, V_POSITIONAL_ONLY_NAME, V_DUPLICATE_KWARGS}:
  47. raise ConfigError(
  48. f'"{ALT_V_ARGS}", "{ALT_V_KWARGS}", "{V_POSITIONAL_ONLY_NAME}" and "{V_DUPLICATE_KWARGS}" '
  49. f'are not permitted as argument names when using the "{validate_arguments.__name__}" decorator'
  50. )
  51. self.raw_function = function
  52. self.arg_mapping: Dict[int, str] = {}
  53. self.positional_only_args = set()
  54. self.v_args_name = 'args'
  55. self.v_kwargs_name = 'kwargs'
  56. type_hints = get_all_type_hints(function)
  57. takes_args = False
  58. takes_kwargs = False
  59. fields: Dict[str, Tuple[Any, Any]] = {}
  60. for i, (name, p) in enumerate(parameters.items()):
  61. if p.annotation is p.empty:
  62. annotation = Any
  63. else:
  64. annotation = type_hints[name]
  65. default = ... if p.default is p.empty else p.default
  66. if p.kind == Parameter.POSITIONAL_ONLY:
  67. self.arg_mapping[i] = name
  68. fields[name] = annotation, default
  69. fields[V_POSITIONAL_ONLY_NAME] = List[str], None
  70. self.positional_only_args.add(name)
  71. elif p.kind == Parameter.POSITIONAL_OR_KEYWORD:
  72. self.arg_mapping[i] = name
  73. fields[name] = annotation, default
  74. fields[V_DUPLICATE_KWARGS] = List[str], None
  75. elif p.kind == Parameter.KEYWORD_ONLY:
  76. fields[name] = annotation, default
  77. elif p.kind == Parameter.VAR_POSITIONAL:
  78. self.v_args_name = name
  79. fields[name] = Tuple[annotation, ...], None
  80. takes_args = True
  81. else:
  82. assert p.kind == Parameter.VAR_KEYWORD, p.kind
  83. self.v_kwargs_name = name
  84. fields[name] = Dict[str, annotation], None # type: ignore
  85. takes_kwargs = True
  86. # these checks avoid a clash between "args" and a field with that name
  87. if not takes_args and self.v_args_name in fields:
  88. self.v_args_name = ALT_V_ARGS
  89. # same with "kwargs"
  90. if not takes_kwargs and self.v_kwargs_name in fields:
  91. self.v_kwargs_name = ALT_V_KWARGS
  92. if not takes_args:
  93. # we add the field so validation below can raise the correct exception
  94. fields[self.v_args_name] = List[Any], None
  95. if not takes_kwargs:
  96. # same with kwargs
  97. fields[self.v_kwargs_name] = Dict[Any, Any], None
  98. self.create_model(fields, takes_args, takes_kwargs, config)
  99. def init_model_instance(self, *args: Any, **kwargs: Any) -> BaseModel:
  100. values = self.build_values(args, kwargs)
  101. return self.model(**values)
  102. def call(self, *args: Any, **kwargs: Any) -> Any:
  103. m = self.init_model_instance(*args, **kwargs)
  104. return self.execute(m)
  105. def build_values(self, args: Tuple[Any, ...], kwargs: Dict[str, Any]) -> Dict[str, Any]:
  106. values: Dict[str, Any] = {}
  107. if args:
  108. arg_iter = enumerate(args)
  109. while True:
  110. try:
  111. i, a = next(arg_iter)
  112. except StopIteration:
  113. break
  114. arg_name = self.arg_mapping.get(i)
  115. if arg_name is not None:
  116. values[arg_name] = a
  117. else:
  118. values[self.v_args_name] = [a] + [a for _, a in arg_iter]
  119. break
  120. var_kwargs: Dict[str, Any] = {}
  121. wrong_positional_args = []
  122. duplicate_kwargs = []
  123. fields_alias = [
  124. field.alias
  125. for name, field in self.model.__fields__.items()
  126. if name not in (self.v_args_name, self.v_kwargs_name)
  127. ]
  128. non_var_fields = set(self.model.__fields__) - {self.v_args_name, self.v_kwargs_name}
  129. for k, v in kwargs.items():
  130. if k in non_var_fields or k in fields_alias:
  131. if k in self.positional_only_args:
  132. wrong_positional_args.append(k)
  133. if k in values:
  134. duplicate_kwargs.append(k)
  135. values[k] = v
  136. else:
  137. var_kwargs[k] = v
  138. if var_kwargs:
  139. values[self.v_kwargs_name] = var_kwargs
  140. if wrong_positional_args:
  141. values[V_POSITIONAL_ONLY_NAME] = wrong_positional_args
  142. if duplicate_kwargs:
  143. values[V_DUPLICATE_KWARGS] = duplicate_kwargs
  144. return values
  145. def execute(self, m: BaseModel) -> Any:
  146. d = {k: v for k, v in m._iter() if k in m.__fields_set__ or m.__fields__[k].default_factory}
  147. var_kwargs = d.pop(self.v_kwargs_name, {})
  148. if self.v_args_name in d:
  149. args_: List[Any] = []
  150. in_kwargs = False
  151. kwargs = {}
  152. for name, value in d.items():
  153. if in_kwargs:
  154. kwargs[name] = value
  155. elif name == self.v_args_name:
  156. args_ += value
  157. in_kwargs = True
  158. else:
  159. args_.append(value)
  160. return self.raw_function(*args_, **kwargs, **var_kwargs)
  161. elif self.positional_only_args:
  162. args_ = []
  163. kwargs = {}
  164. for name, value in d.items():
  165. if name in self.positional_only_args:
  166. args_.append(value)
  167. else:
  168. kwargs[name] = value
  169. return self.raw_function(*args_, **kwargs, **var_kwargs)
  170. else:
  171. return self.raw_function(**d, **var_kwargs)
  172. def create_model(self, fields: Dict[str, Any], takes_args: bool, takes_kwargs: bool, config: 'ConfigType') -> None:
  173. pos_args = len(self.arg_mapping)
  174. class CustomConfig:
  175. pass
  176. if not TYPE_CHECKING: # pragma: no branch
  177. if isinstance(config, dict):
  178. CustomConfig = type('Config', (), config) # noqa: F811
  179. elif config is not None:
  180. CustomConfig = config # noqa: F811
  181. if hasattr(CustomConfig, 'fields') or hasattr(CustomConfig, 'alias_generator'):
  182. raise ConfigError(
  183. 'Setting the "fields" and "alias_generator" property on custom Config for '
  184. '@validate_arguments is not yet supported, please remove.'
  185. )
  186. class DecoratorBaseModel(BaseModel):
  187. @validator(self.v_args_name, check_fields=False, allow_reuse=True)
  188. def check_args(cls, v: Optional[List[Any]]) -> Optional[List[Any]]:
  189. if takes_args or v is None:
  190. return v
  191. raise TypeError(f'{pos_args} positional arguments expected but {pos_args + len(v)} given')
  192. @validator(self.v_kwargs_name, check_fields=False, allow_reuse=True)
  193. def check_kwargs(cls, v: Optional[Dict[str, Any]]) -> Optional[Dict[str, Any]]:
  194. if takes_kwargs or v is None:
  195. return v
  196. plural = '' if len(v) == 1 else 's'
  197. keys = ', '.join(map(repr, v.keys()))
  198. raise TypeError(f'unexpected keyword argument{plural}: {keys}')
  199. @validator(V_POSITIONAL_ONLY_NAME, check_fields=False, allow_reuse=True)
  200. def check_positional_only(cls, v: Optional[List[str]]) -> None:
  201. if v is None:
  202. return
  203. plural = '' if len(v) == 1 else 's'
  204. keys = ', '.join(map(repr, v))
  205. raise TypeError(f'positional-only argument{plural} passed as keyword argument{plural}: {keys}')
  206. @validator(V_DUPLICATE_KWARGS, check_fields=False, allow_reuse=True)
  207. def check_duplicate_kwargs(cls, v: Optional[List[str]]) -> None:
  208. if v is None:
  209. return
  210. plural = '' if len(v) == 1 else 's'
  211. keys = ', '.join(map(repr, v))
  212. raise TypeError(f'multiple values for argument{plural}: {keys}')
  213. class Config(CustomConfig):
  214. extra = getattr(CustomConfig, 'extra', Extra.forbid)
  215. self.model = create_model(to_camel(self.raw_function.__name__), __base__=DecoratorBaseModel, **fields)