decorator.py 11 KB

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