_validate_call.py 5.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141
  1. from __future__ import annotations as _annotations
  2. import functools
  3. import inspect
  4. from collections.abc import Awaitable
  5. from functools import partial
  6. from typing import Any, Callable
  7. import pydantic_core
  8. from ..config import ConfigDict
  9. from ..plugin._schema_validator import create_schema_validator
  10. from ._config import ConfigWrapper
  11. from ._generate_schema import GenerateSchema, ValidateCallSupportedTypes
  12. from ._namespace_utils import MappingNamespace, NsResolver, ns_for_function
  13. from ._typing_extra import signature_no_eval
  14. def extract_function_name(func: ValidateCallSupportedTypes) -> str:
  15. """Extract the name of a `ValidateCallSupportedTypes` object."""
  16. return f'partial({func.func.__name__})' if isinstance(func, functools.partial) else func.__name__
  17. def extract_function_qualname(func: ValidateCallSupportedTypes) -> str:
  18. """Extract the qualname of a `ValidateCallSupportedTypes` object."""
  19. return f'partial({func.func.__qualname__})' if isinstance(func, functools.partial) else func.__qualname__
  20. def update_wrapper_attributes(wrapped: ValidateCallSupportedTypes, wrapper: Callable[..., Any]):
  21. """Update the `wrapper` function with the attributes of the `wrapped` function. Return the updated function."""
  22. if inspect.iscoroutinefunction(wrapped):
  23. @functools.wraps(wrapped)
  24. async def wrapper_function(*args, **kwargs): # type: ignore
  25. return await wrapper(*args, **kwargs)
  26. else:
  27. @functools.wraps(wrapped)
  28. def wrapper_function(*args, **kwargs):
  29. return wrapper(*args, **kwargs)
  30. # We need to manually update this because `partial` object has no `__name__` and `__qualname__`.
  31. wrapper_function.__name__ = extract_function_name(wrapped)
  32. wrapper_function.__qualname__ = extract_function_qualname(wrapped)
  33. wrapper_function.raw_function = wrapped # type: ignore
  34. return wrapper_function
  35. class ValidateCallWrapper:
  36. """This is a wrapper around a function that validates the arguments passed to it, and optionally the return value."""
  37. __slots__ = (
  38. 'function',
  39. 'validate_return',
  40. 'schema_type',
  41. 'module',
  42. 'qualname',
  43. 'ns_resolver',
  44. 'config_wrapper',
  45. '__pydantic_complete__',
  46. '__pydantic_validator__',
  47. '__return_pydantic_validator__',
  48. )
  49. def __init__(
  50. self,
  51. function: ValidateCallSupportedTypes,
  52. config: ConfigDict | None,
  53. validate_return: bool,
  54. parent_namespace: MappingNamespace | None,
  55. ) -> None:
  56. self.function = function
  57. self.validate_return = validate_return
  58. if isinstance(function, partial):
  59. self.schema_type = function.func
  60. self.module = function.func.__module__
  61. else:
  62. self.schema_type = function
  63. self.module = function.__module__
  64. self.qualname = extract_function_qualname(function)
  65. self.ns_resolver = NsResolver(
  66. namespaces_tuple=ns_for_function(self.schema_type, parent_namespace=parent_namespace)
  67. )
  68. self.config_wrapper = ConfigWrapper(config)
  69. if not self.config_wrapper.defer_build:
  70. self._create_validators()
  71. else:
  72. self.__pydantic_complete__ = False
  73. def _create_validators(self) -> None:
  74. gen_schema = GenerateSchema(self.config_wrapper, self.ns_resolver)
  75. schema = gen_schema.clean_schema(gen_schema.generate_schema(self.function))
  76. core_config = self.config_wrapper.core_config(title=self.qualname)
  77. self.__pydantic_validator__ = create_schema_validator(
  78. schema,
  79. self.schema_type,
  80. self.module,
  81. self.qualname,
  82. 'validate_call',
  83. core_config,
  84. self.config_wrapper.plugin_settings,
  85. )
  86. if self.validate_return:
  87. signature = signature_no_eval(self.function)
  88. return_type = signature.return_annotation if signature.return_annotation is not signature.empty else Any
  89. gen_schema = GenerateSchema(self.config_wrapper, self.ns_resolver)
  90. schema = gen_schema.clean_schema(gen_schema.generate_schema(return_type))
  91. validator = create_schema_validator(
  92. schema,
  93. self.schema_type,
  94. self.module,
  95. self.qualname,
  96. 'validate_call',
  97. core_config,
  98. self.config_wrapper.plugin_settings,
  99. )
  100. if inspect.iscoroutinefunction(self.function):
  101. async def return_val_wrapper(aw: Awaitable[Any]) -> None:
  102. return validator.validate_python(await aw)
  103. self.__return_pydantic_validator__ = return_val_wrapper
  104. else:
  105. self.__return_pydantic_validator__ = validator.validate_python
  106. else:
  107. self.__return_pydantic_validator__ = None
  108. self.__pydantic_complete__ = True
  109. def __call__(self, *args: Any, **kwargs: Any) -> Any:
  110. if not self.__pydantic_complete__:
  111. self._create_validators()
  112. res = self.__pydantic_validator__.validate_python(pydantic_core.ArgsKwargs(args, kwargs))
  113. if self.__return_pydantic_validator__:
  114. return self.__return_pydantic_validator__(res)
  115. else:
  116. return res