_signature.py 6.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189
  1. from __future__ import annotations
  2. import dataclasses
  3. from inspect import Parameter, Signature
  4. from typing import TYPE_CHECKING, Any, Callable
  5. from pydantic_core import PydanticUndefined
  6. from ._typing_extra import signature_no_eval
  7. from ._utils import is_valid_identifier
  8. if TYPE_CHECKING:
  9. from ..config import ExtraValues
  10. from ..fields import FieldInfo
  11. # Copied over from stdlib dataclasses
  12. class _HAS_DEFAULT_FACTORY_CLASS:
  13. def __repr__(self):
  14. return '<factory>'
  15. _HAS_DEFAULT_FACTORY = _HAS_DEFAULT_FACTORY_CLASS()
  16. def _field_name_for_signature(field_name: str, field_info: FieldInfo) -> str:
  17. """Extract the correct name to use for the field when generating a signature.
  18. Assuming the field has a valid alias, this will return the alias. Otherwise, it will return the field name.
  19. First priority is given to the alias, then the validation_alias, then the field name.
  20. Args:
  21. field_name: The name of the field
  22. field_info: The corresponding FieldInfo object.
  23. Returns:
  24. The correct name to use when generating a signature.
  25. """
  26. if isinstance(field_info.alias, str) and is_valid_identifier(field_info.alias):
  27. return field_info.alias
  28. if isinstance(field_info.validation_alias, str) and is_valid_identifier(field_info.validation_alias):
  29. return field_info.validation_alias
  30. return field_name
  31. def _process_param_defaults(param: Parameter) -> Parameter:
  32. """Modify the signature for a parameter in a dataclass where the default value is a FieldInfo instance.
  33. Args:
  34. param (Parameter): The parameter
  35. Returns:
  36. Parameter: The custom processed parameter
  37. """
  38. from ..fields import FieldInfo
  39. param_default = param.default
  40. if isinstance(param_default, FieldInfo):
  41. annotation = param.annotation
  42. # Replace the annotation if appropriate
  43. # inspect does "clever" things to show annotations as strings because we have
  44. # `from __future__ import annotations` in main, we don't want that
  45. if annotation == 'Any':
  46. annotation = Any
  47. # Replace the field default
  48. default = param_default.default
  49. if default is PydanticUndefined:
  50. if param_default.default_factory is None:
  51. default = Signature.empty
  52. else:
  53. # this is used by dataclasses to indicate a factory exists:
  54. default = dataclasses._HAS_DEFAULT_FACTORY # type: ignore
  55. return param.replace(
  56. annotation=annotation, name=_field_name_for_signature(param.name, param_default), default=default
  57. )
  58. return param
  59. def _generate_signature_parameters( # noqa: C901 (ignore complexity, could use a refactor)
  60. init: Callable[..., None],
  61. fields: dict[str, FieldInfo],
  62. validate_by_name: bool,
  63. extra: ExtraValues | None,
  64. ) -> dict[str, Parameter]:
  65. """Generate a mapping of parameter names to Parameter objects for a pydantic BaseModel or dataclass."""
  66. from itertools import islice
  67. present_params = signature_no_eval(init).parameters.values()
  68. merged_params: dict[str, Parameter] = {}
  69. var_kw = None
  70. use_var_kw = False
  71. for param in islice(present_params, 1, None): # skip self arg
  72. # inspect does "clever" things to show annotations as strings because we have
  73. # `from __future__ import annotations` in main, we don't want that
  74. if fields.get(param.name):
  75. # exclude params with init=False
  76. if getattr(fields[param.name], 'init', True) is False:
  77. continue
  78. param = param.replace(name=_field_name_for_signature(param.name, fields[param.name]))
  79. if param.annotation == 'Any':
  80. param = param.replace(annotation=Any)
  81. if param.kind is param.VAR_KEYWORD:
  82. var_kw = param
  83. continue
  84. merged_params[param.name] = param
  85. if var_kw: # if custom init has no var_kw, fields which are not declared in it cannot be passed through
  86. allow_names = validate_by_name
  87. for field_name, field in fields.items():
  88. # when alias is a str it should be used for signature generation
  89. param_name = _field_name_for_signature(field_name, field)
  90. if field_name in merged_params or param_name in merged_params:
  91. continue
  92. if not is_valid_identifier(param_name):
  93. if allow_names:
  94. param_name = field_name
  95. else:
  96. use_var_kw = True
  97. continue
  98. if field.is_required():
  99. default = Parameter.empty
  100. elif field.default_factory is not None:
  101. # Mimics stdlib dataclasses:
  102. default = _HAS_DEFAULT_FACTORY
  103. else:
  104. default = field.default
  105. merged_params[param_name] = Parameter(
  106. param_name,
  107. Parameter.KEYWORD_ONLY,
  108. annotation=field.rebuild_annotation(),
  109. default=default,
  110. )
  111. if extra == 'allow':
  112. use_var_kw = True
  113. if var_kw and use_var_kw:
  114. # Make sure the parameter for extra kwargs
  115. # does not have the same name as a field
  116. default_model_signature = [
  117. ('self', Parameter.POSITIONAL_ONLY),
  118. ('data', Parameter.VAR_KEYWORD),
  119. ]
  120. if [(p.name, p.kind) for p in present_params] == default_model_signature:
  121. # if this is the standard model signature, use extra_data as the extra args name
  122. var_kw_name = 'extra_data'
  123. else:
  124. # else start from var_kw
  125. var_kw_name = var_kw.name
  126. # generate a name that's definitely unique
  127. while var_kw_name in fields:
  128. var_kw_name += '_'
  129. merged_params[var_kw_name] = var_kw.replace(name=var_kw_name)
  130. return merged_params
  131. def generate_pydantic_signature(
  132. init: Callable[..., None],
  133. fields: dict[str, FieldInfo],
  134. validate_by_name: bool,
  135. extra: ExtraValues | None,
  136. is_dataclass: bool = False,
  137. ) -> Signature:
  138. """Generate signature for a pydantic BaseModel or dataclass.
  139. Args:
  140. init: The class init.
  141. fields: The model fields.
  142. validate_by_name: The `validate_by_name` value of the config.
  143. extra: The `extra` value of the config.
  144. is_dataclass: Whether the model is a dataclass.
  145. Returns:
  146. The dataclass/BaseModel subclass signature.
  147. """
  148. merged_params = _generate_signature_parameters(init, fields, validate_by_name, extra)
  149. if is_dataclass:
  150. merged_params = {k: _process_param_defaults(v) for k, v in merged_params.items()}
  151. return Signature(parameters=list(merged_params.values()), return_annotation=None)