shared.py 7.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222
  1. import types
  2. import typing
  3. import warnings
  4. from collections import deque
  5. from collections.abc import Mapping, Sequence
  6. from dataclasses import is_dataclass
  7. from typing import (
  8. Annotated,
  9. Any,
  10. TypeGuard,
  11. TypeVar,
  12. Union,
  13. get_args,
  14. get_origin,
  15. )
  16. from fastapi.types import UnionType
  17. from pydantic import BaseModel
  18. from pydantic.version import VERSION as PYDANTIC_VERSION
  19. from starlette.datastructures import UploadFile
  20. _T = TypeVar("_T")
  21. # Copy from Pydantic: pydantic/_internal/_typing_extra.py
  22. WithArgsTypes: tuple[Any, ...] = (
  23. typing._GenericAlias, # type: ignore[attr-defined] # ty: ignore[unresolved-attribute]
  24. types.GenericAlias,
  25. types.UnionType,
  26. ) # pyright: ignore[reportAttributeAccessIssue]
  27. PYDANTIC_VERSION_MINOR_TUPLE = tuple(int(x) for x in PYDANTIC_VERSION.split(".")[:2])
  28. sequence_annotation_to_type = {
  29. Sequence: list,
  30. list: list,
  31. tuple: tuple,
  32. set: set,
  33. frozenset: frozenset,
  34. deque: deque,
  35. }
  36. sequence_types: tuple[type[Any], ...] = tuple(sequence_annotation_to_type.keys())
  37. # Copy of Pydantic: pydantic/_internal/_utils.py with added TypeGuard
  38. def lenient_issubclass(
  39. cls: Any, class_or_tuple: type[_T] | tuple[type[_T], ...] | None
  40. ) -> TypeGuard[type[_T]]:
  41. try:
  42. return isinstance(cls, type) and issubclass(cls, class_or_tuple) # type: ignore[arg-type] # ty: ignore[invalid-argument-type]
  43. except TypeError: # pragma: no cover
  44. if isinstance(cls, WithArgsTypes):
  45. return False
  46. raise # pragma: no cover
  47. def _annotation_is_sequence(annotation: type[Any] | None) -> bool:
  48. if lenient_issubclass(annotation, (str, bytes)):
  49. return False
  50. return lenient_issubclass(annotation, sequence_types)
  51. def field_annotation_is_sequence(annotation: type[Any] | None) -> bool:
  52. origin = get_origin(annotation)
  53. if origin is Annotated:
  54. return field_annotation_is_sequence(get_args(annotation)[0])
  55. if origin is Union or origin is UnionType:
  56. for arg in get_args(annotation):
  57. if field_annotation_is_sequence(arg):
  58. return True
  59. return False
  60. return _annotation_is_sequence(annotation) or _annotation_is_sequence(
  61. get_origin(annotation)
  62. )
  63. def value_is_sequence(value: Any) -> bool:
  64. return isinstance(value, sequence_types) and not isinstance(value, (str, bytes))
  65. def _annotation_is_complex(annotation: type[Any] | None) -> bool:
  66. return (
  67. lenient_issubclass(annotation, (BaseModel, Mapping, UploadFile))
  68. or _annotation_is_sequence(annotation)
  69. or is_dataclass(annotation)
  70. )
  71. def field_annotation_is_complex(annotation: type[Any] | None) -> bool:
  72. origin = get_origin(annotation)
  73. if origin is Union or origin is UnionType:
  74. return any(field_annotation_is_complex(arg) for arg in get_args(annotation))
  75. if origin is Annotated:
  76. return field_annotation_is_complex(get_args(annotation)[0])
  77. return (
  78. _annotation_is_complex(annotation)
  79. or _annotation_is_complex(origin)
  80. or hasattr(origin, "__pydantic_core_schema__")
  81. or hasattr(origin, "__get_pydantic_core_schema__")
  82. )
  83. def field_annotation_is_scalar(annotation: Any) -> bool:
  84. # handle Ellipsis here to make tuple[int, ...] work nicely
  85. return annotation is Ellipsis or not field_annotation_is_complex(annotation)
  86. def field_annotation_is_scalar_sequence(annotation: type[Any] | None) -> bool:
  87. origin = get_origin(annotation)
  88. if origin is Annotated:
  89. return field_annotation_is_scalar_sequence(get_args(annotation)[0])
  90. if origin is Union or origin is UnionType:
  91. at_least_one_scalar_sequence = False
  92. for arg in get_args(annotation):
  93. if field_annotation_is_scalar_sequence(arg):
  94. at_least_one_scalar_sequence = True
  95. continue
  96. elif not field_annotation_is_scalar(arg):
  97. return False
  98. return at_least_one_scalar_sequence
  99. return field_annotation_is_sequence(annotation) and all(
  100. field_annotation_is_scalar(sub_annotation)
  101. for sub_annotation in get_args(annotation)
  102. )
  103. def is_bytes_or_nonable_bytes_annotation(annotation: Any) -> bool:
  104. if lenient_issubclass(annotation, bytes):
  105. return True
  106. origin = get_origin(annotation)
  107. if origin is Union or origin is UnionType:
  108. for arg in get_args(annotation):
  109. if lenient_issubclass(arg, bytes):
  110. return True
  111. return False
  112. def is_uploadfile_or_nonable_uploadfile_annotation(annotation: Any) -> bool:
  113. if lenient_issubclass(annotation, UploadFile):
  114. return True
  115. origin = get_origin(annotation)
  116. if origin is Union or origin is UnionType:
  117. for arg in get_args(annotation):
  118. if lenient_issubclass(arg, UploadFile):
  119. return True
  120. return False
  121. def is_bytes_sequence_annotation(annotation: Any) -> bool:
  122. origin = get_origin(annotation)
  123. if origin is Union or origin is UnionType:
  124. at_least_one = False
  125. for arg in get_args(annotation):
  126. if is_bytes_sequence_annotation(arg):
  127. at_least_one = True
  128. continue
  129. return at_least_one
  130. return field_annotation_is_sequence(annotation) and all(
  131. is_bytes_or_nonable_bytes_annotation(sub_annotation)
  132. for sub_annotation in get_args(annotation)
  133. )
  134. def is_uploadfile_sequence_annotation(annotation: Any) -> bool:
  135. origin = get_origin(annotation)
  136. if origin is Union or origin is UnionType:
  137. at_least_one = False
  138. for arg in get_args(annotation):
  139. if is_uploadfile_sequence_annotation(arg):
  140. at_least_one = True
  141. continue
  142. return at_least_one
  143. return field_annotation_is_sequence(annotation) and all(
  144. is_uploadfile_or_nonable_uploadfile_annotation(sub_annotation)
  145. for sub_annotation in get_args(annotation)
  146. )
  147. def is_pydantic_v1_model_instance(obj: Any) -> bool:
  148. # TODO: remove this function once the required version of Pydantic fully
  149. # removes pydantic.v1
  150. try:
  151. with warnings.catch_warnings():
  152. warnings.simplefilter("ignore", UserWarning)
  153. from pydantic import v1
  154. except ImportError: # pragma: no cover
  155. return False
  156. return isinstance(obj, v1.BaseModel)
  157. def is_pydantic_v1_model_class(cls: Any) -> bool:
  158. # TODO: remove this function once the required version of Pydantic fully
  159. # removes pydantic.v1
  160. try:
  161. with warnings.catch_warnings():
  162. warnings.simplefilter("ignore", UserWarning)
  163. from pydantic import v1
  164. except ImportError: # pragma: no cover
  165. return False
  166. return lenient_issubclass(cls, v1.BaseModel)
  167. def annotation_is_pydantic_v1(annotation: Any) -> bool:
  168. if is_pydantic_v1_model_class(annotation):
  169. return True
  170. origin = get_origin(annotation)
  171. if origin is Union or origin is UnionType:
  172. for arg in get_args(annotation):
  173. if is_pydantic_v1_model_class(arg):
  174. return True
  175. if field_annotation_is_sequence(annotation):
  176. for sub_annotation in get_args(annotation):
  177. if annotation_is_pydantic_v1(sub_annotation):
  178. return True
  179. return False