models.py 7.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234
  1. import inspect
  2. import sys
  3. from collections.abc import Callable
  4. from dataclasses import dataclass, field
  5. from functools import lru_cache, partial
  6. from typing import Any, Literal
  7. from fastapi._compat import ModelField
  8. from fastapi.security.base import SecurityBase
  9. from fastapi.types import DependencyCacheKey
  10. if sys.version_info >= (3, 13): # pragma: no cover
  11. from inspect import iscoroutinefunction
  12. else: # pragma: no cover
  13. from asyncio import iscoroutinefunction
  14. def _unwrapped_call(call: Callable[..., Any] | None) -> Any:
  15. if call is None:
  16. return call # pragma: no cover
  17. unwrapped = inspect.unwrap(_impartial(call))
  18. return unwrapped
  19. def _impartial(func: Callable[..., Any]) -> Callable[..., Any]:
  20. while isinstance(func, partial):
  21. func = func.func
  22. return func
  23. @dataclass(slots=True)
  24. class Dependant:
  25. path_params: list[ModelField] = field(default_factory=list)
  26. query_params: list[ModelField] = field(default_factory=list)
  27. header_params: list[ModelField] = field(default_factory=list)
  28. cookie_params: list[ModelField] = field(default_factory=list)
  29. body_params: list[ModelField] = field(default_factory=list)
  30. dependencies: list["Dependant"] = field(default_factory=list)
  31. name: str | None = None
  32. call: Callable[..., Any] | None = None
  33. request_param_name: str | None = None
  34. websocket_param_name: str | None = None
  35. http_connection_param_name: str | None = None
  36. response_param_name: str | None = None
  37. background_tasks_param_name: str | None = None
  38. security_scopes_param_name: str | None = None
  39. own_oauth_scopes: list[str] | None = None
  40. parent_oauth_scopes: list[str] | None = None
  41. use_cache: bool = True
  42. path: str | None = None
  43. scope: Literal["function", "request"] | None = None
  44. _UsesScopesCache = dict[int, tuple[Dependant, bool]]
  45. _CALLABLE_CLASSIFICATION_CACHE_SIZE = 4096
  46. class _CallIdentity:
  47. __slots__ = ("call",)
  48. def __init__(self, call: Callable[..., Any]) -> None:
  49. self.call = call
  50. def __hash__(self) -> int:
  51. return id(self.call)
  52. def __eq__(self, other: object) -> bool:
  53. return isinstance(other, _CallIdentity) and self.call is other.call
  54. def _get_oauth_scopes(*, dependant: Dependant) -> list[str]:
  55. scopes = (
  56. dependant.parent_oauth_scopes.copy() if dependant.parent_oauth_scopes else []
  57. )
  58. # This doesn't use a set to preserve order, just in case
  59. for scope in dependant.own_oauth_scopes or []:
  60. if scope not in scopes:
  61. scopes.append(scope)
  62. return scopes
  63. def _get_cache_key(
  64. *,
  65. dependant: Dependant,
  66. uses_scopes_cache: _UsesScopesCache | None = None,
  67. ) -> DependencyCacheKey:
  68. scopes_for_cache = (
  69. tuple(sorted(set(_get_oauth_scopes(dependant=dependant))))
  70. if _uses_scopes(dependant=dependant, cache=uses_scopes_cache)
  71. else ()
  72. )
  73. return (
  74. dependant.call,
  75. scopes_for_cache,
  76. _get_computed_scope(dependant=dependant) or "",
  77. )
  78. def _uses_scopes(
  79. *, dependant: Dependant, cache: _UsesScopesCache | None = None
  80. ) -> bool:
  81. if cache is None:
  82. cache = {}
  83. cache_key = id(dependant)
  84. cached = cache.get(cache_key)
  85. if cached is not None and cached[0] is dependant:
  86. return cached[1]
  87. if dependant.own_oauth_scopes:
  88. result = True
  89. elif dependant.security_scopes_param_name is not None:
  90. result = True
  91. elif _is_security_scheme(dependant=dependant):
  92. result = True
  93. else:
  94. result = any(
  95. _uses_scopes(dependant=sub_dep, cache=cache)
  96. for sub_dep in dependant.dependencies
  97. )
  98. cache[cache_key] = (dependant, result)
  99. return result
  100. def _is_security_scheme(*, dependant: Dependant) -> bool:
  101. if dependant.call is None:
  102. return False # pragma: no cover
  103. unwrapped = _unwrapped_call(dependant.call)
  104. return isinstance(unwrapped, SecurityBase)
  105. def _get_security_scheme(*, dependant: Dependant) -> SecurityBase:
  106. # Mainly to get the type of SecurityBase, but it's the same dependant.call
  107. unwrapped = _unwrapped_call(dependant.call)
  108. assert isinstance(unwrapped, SecurityBase)
  109. return unwrapped
  110. @lru_cache(maxsize=_CALLABLE_CLASSIFICATION_CACHE_SIZE)
  111. def _is_gen_callable_cached(call_identity: _CallIdentity) -> bool:
  112. call = call_identity.call
  113. if inspect.isgeneratorfunction(_impartial(call)) or inspect.isgeneratorfunction(
  114. _unwrapped_call(call)
  115. ):
  116. return True
  117. if inspect.isclass(_unwrapped_call(call)):
  118. return False
  119. dunder_call = getattr(_impartial(call), "__call__", None) # noqa: B004
  120. if dunder_call is None:
  121. return False # pragma: no cover
  122. if inspect.isgeneratorfunction(
  123. _impartial(dunder_call)
  124. ) or inspect.isgeneratorfunction(_unwrapped_call(dunder_call)):
  125. return True
  126. dunder_unwrapped_call = getattr(_unwrapped_call(call), "__call__", None) # noqa: B004
  127. if dunder_unwrapped_call is None:
  128. return False # pragma: no cover
  129. return inspect.isgeneratorfunction(
  130. _impartial(dunder_unwrapped_call)
  131. ) or inspect.isgeneratorfunction(_unwrapped_call(dunder_unwrapped_call))
  132. def _is_gen_callable(call: Callable[..., Any] | None) -> bool:
  133. if call is None:
  134. return False # pragma: no cover
  135. return _is_gen_callable_cached(_CallIdentity(call))
  136. @lru_cache(maxsize=_CALLABLE_CLASSIFICATION_CACHE_SIZE)
  137. def _is_async_gen_callable_cached(call_identity: _CallIdentity) -> bool:
  138. call = call_identity.call
  139. if inspect.isasyncgenfunction(_impartial(call)) or inspect.isasyncgenfunction(
  140. _unwrapped_call(call)
  141. ):
  142. return True
  143. if inspect.isclass(_unwrapped_call(call)):
  144. return False
  145. dunder_call = getattr(_impartial(call), "__call__", None) # noqa: B004
  146. if dunder_call is None:
  147. return False # pragma: no cover
  148. if inspect.isasyncgenfunction(
  149. _impartial(dunder_call)
  150. ) or inspect.isasyncgenfunction(_unwrapped_call(dunder_call)):
  151. return True
  152. dunder_unwrapped_call = getattr(_unwrapped_call(call), "__call__", None) # noqa: B004
  153. if dunder_unwrapped_call is None:
  154. return False # pragma: no cover
  155. return inspect.isasyncgenfunction(
  156. _impartial(dunder_unwrapped_call)
  157. ) or inspect.isasyncgenfunction(_unwrapped_call(dunder_unwrapped_call))
  158. def _is_async_gen_callable(call: Callable[..., Any] | None) -> bool:
  159. if call is None:
  160. return False # pragma: no cover
  161. return _is_async_gen_callable_cached(_CallIdentity(call))
  162. @lru_cache(maxsize=_CALLABLE_CLASSIFICATION_CACHE_SIZE)
  163. def _is_coroutine_callable_cached(call_identity: _CallIdentity) -> bool:
  164. call = call_identity.call
  165. if inspect.isroutine(_impartial(call)) and iscoroutinefunction(_impartial(call)):
  166. return True
  167. if inspect.isroutine(_unwrapped_call(call)) and iscoroutinefunction(
  168. _unwrapped_call(call)
  169. ):
  170. return True
  171. if inspect.isclass(_unwrapped_call(call)):
  172. return False
  173. dunder_call = getattr(_impartial(call), "__call__", None) # noqa: B004
  174. if dunder_call is None:
  175. return False # pragma: no cover
  176. if iscoroutinefunction(_impartial(dunder_call)) or iscoroutinefunction(
  177. _unwrapped_call(dunder_call)
  178. ):
  179. return True
  180. dunder_unwrapped_call = getattr(_unwrapped_call(call), "__call__", None) # noqa: B004
  181. if dunder_unwrapped_call is None:
  182. return False # pragma: no cover
  183. return iscoroutinefunction(
  184. _impartial(dunder_unwrapped_call)
  185. ) or iscoroutinefunction(_unwrapped_call(dunder_unwrapped_call))
  186. def _is_coroutine_callable(call: Callable[..., Any] | None) -> bool:
  187. if call is None:
  188. return False # pragma: no cover
  189. return _is_coroutine_callable_cached(_CallIdentity(call))
  190. def _get_computed_scope(*, dependant: Dependant) -> str | None:
  191. if dependant.scope:
  192. return dependant.scope
  193. if _is_gen_callable(dependant.call) or _is_async_gen_callable(dependant.call):
  194. return "request"
  195. return None