utils.py 38 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928929930931932933934935936937938939940941942943944945946947948949950951952953954955956957958959960961962963964965966967968969970971972973974975976977978979980981982983984985986987988989990991992993994995996997998999100010011002100310041005100610071008100910101011101210131014101510161017101810191020102110221023102410251026102710281029103010311032103310341035103610371038103910401041104210431044104510461047104810491050105110521053
  1. import dataclasses
  2. import inspect
  3. import sys
  4. from collections.abc import (
  5. AsyncGenerator,
  6. AsyncIterable,
  7. AsyncIterator,
  8. Callable,
  9. Generator,
  10. Iterable,
  11. Iterator,
  12. Mapping,
  13. Sequence,
  14. )
  15. from contextlib import AsyncExitStack, contextmanager
  16. from copy import copy, deepcopy
  17. from dataclasses import dataclass
  18. from typing import (
  19. Annotated,
  20. Any,
  21. ForwardRef,
  22. Literal,
  23. Union,
  24. cast,
  25. get_args,
  26. get_origin,
  27. )
  28. from fastapi import params
  29. from fastapi._compat import (
  30. ModelField,
  31. RequiredParam,
  32. Undefined,
  33. copy_field_info,
  34. create_body_model,
  35. evaluate_forwardref,
  36. field_annotation_is_scalar,
  37. field_annotation_is_scalar_sequence,
  38. field_annotation_is_sequence,
  39. get_cached_model_fields,
  40. get_missing_field_error,
  41. is_bytes_or_nonable_bytes_annotation,
  42. is_bytes_sequence_annotation,
  43. is_scalar_field,
  44. is_uploadfile_or_nonable_uploadfile_annotation,
  45. is_uploadfile_sequence_annotation,
  46. lenient_issubclass,
  47. sequence_types,
  48. serialize_sequence_value,
  49. value_is_sequence,
  50. )
  51. from fastapi.background import BackgroundTasks
  52. from fastapi.concurrency import (
  53. asynccontextmanager,
  54. contextmanager_in_threadpool,
  55. )
  56. from fastapi.dependencies.models import (
  57. Dependant,
  58. _get_cache_key,
  59. _get_computed_scope,
  60. _get_oauth_scopes,
  61. _is_async_gen_callable,
  62. _is_coroutine_callable,
  63. _is_gen_callable,
  64. _UsesScopesCache,
  65. )
  66. from fastapi.exceptions import DependencyScopeError
  67. from fastapi.logger import logger
  68. from fastapi.security.oauth2 import SecurityScopes
  69. from fastapi.types import DependencyCacheKey
  70. from fastapi.utils import create_model_field, get_path_param_names
  71. from pydantic import BaseModel, Json
  72. from pydantic.fields import FieldInfo
  73. from starlette.background import BackgroundTasks as StarletteBackgroundTasks
  74. from starlette.concurrency import run_in_threadpool
  75. from starlette.datastructures import (
  76. FormData,
  77. Headers,
  78. ImmutableMultiDict,
  79. QueryParams,
  80. UploadFile,
  81. )
  82. from starlette.requests import HTTPConnection, Request
  83. from starlette.responses import Response
  84. from starlette.websockets import WebSocket
  85. from typing_inspection.typing_objects import is_typealiastype
  86. multipart_not_installed_error = (
  87. 'Form data requires "python-multipart" to be installed. \n'
  88. 'You can install "python-multipart" with: \n\n'
  89. "pip install python-multipart\n"
  90. )
  91. multipart_incorrect_install_error = (
  92. 'Form data requires "python-multipart" to be installed. '
  93. 'It seems you installed "multipart" instead. \n'
  94. 'You can remove "multipart" with: \n\n'
  95. "pip uninstall multipart\n\n"
  96. 'And then install "python-multipart" with: \n\n'
  97. "pip install python-multipart\n"
  98. )
  99. def ensure_multipart_is_installed() -> None:
  100. try:
  101. from python_multipart import __version__
  102. # Import an attribute that can be mocked/deleted in testing
  103. assert __version__ > "0.0.12"
  104. except (ImportError, AssertionError):
  105. try:
  106. # __version__ is available in both multiparts, and can be mocked
  107. from multipart import ( # type: ignore[no-redef,import-untyped]
  108. __version__,
  109. )
  110. assert __version__
  111. try:
  112. # parse_options_header is only available in the right multipart
  113. from multipart.multipart import ( # type: ignore[import-untyped]
  114. parse_options_header,
  115. )
  116. assert parse_options_header
  117. except ImportError:
  118. logger.error(multipart_incorrect_install_error)
  119. raise RuntimeError(multipart_incorrect_install_error) from None
  120. except ImportError:
  121. logger.error(multipart_not_installed_error)
  122. raise RuntimeError(multipart_not_installed_error) from None
  123. def get_parameterless_sub_dependant(*, depends: params.Depends, path: str) -> Dependant:
  124. assert callable(depends.dependency), (
  125. "A parameter-less dependency must have a callable dependency"
  126. )
  127. own_oauth_scopes: list[str] = []
  128. if isinstance(depends, params.Security) and depends.scopes:
  129. own_oauth_scopes.extend(depends.scopes)
  130. return get_dependant(
  131. path=path,
  132. call=depends.dependency,
  133. scope=depends.scope,
  134. own_oauth_scopes=own_oauth_scopes,
  135. )
  136. def _get_flat_body_params(dependant: Dependant) -> list[ModelField]:
  137. body_params: list[ModelField] = []
  138. dependants = [dependant]
  139. while dependants:
  140. current_dependant = dependants.pop()
  141. body_params.extend(current_dependant.body_params)
  142. dependants.extend(reversed(current_dependant.dependencies))
  143. return body_params
  144. def _get_flat_fields_from_params(fields: list[ModelField]) -> list[ModelField]:
  145. if not fields:
  146. return fields
  147. first_field = fields[0]
  148. if len(fields) == 1 and lenient_issubclass(
  149. first_field.field_info.annotation, BaseModel
  150. ):
  151. fields_to_extract = get_cached_model_fields(first_field.field_info.annotation)
  152. return fields_to_extract
  153. return fields
  154. def get_flat_params(dependant: Dependant) -> list[ModelField]:
  155. path_params: list[ModelField] = []
  156. query_params: list[ModelField] = []
  157. header_params: list[ModelField] = []
  158. cookie_params: list[ModelField] = []
  159. visited: list[DependencyCacheKey] = []
  160. uses_scopes_cache: _UsesScopesCache = {}
  161. dependants = [dependant]
  162. while dependants:
  163. current_dependant = dependants.pop()
  164. cache_key = _get_cache_key(
  165. dependant=current_dependant,
  166. uses_scopes_cache=uses_scopes_cache,
  167. )
  168. if cache_key in visited:
  169. continue
  170. visited.append(cache_key)
  171. path_params.extend(current_dependant.path_params)
  172. query_params.extend(current_dependant.query_params)
  173. header_params.extend(current_dependant.header_params)
  174. cookie_params.extend(current_dependant.cookie_params)
  175. dependants.extend(reversed(current_dependant.dependencies))
  176. path_params = _get_flat_fields_from_params(path_params)
  177. query_params = _get_flat_fields_from_params(query_params)
  178. header_params = _get_flat_fields_from_params(header_params)
  179. cookie_params = _get_flat_fields_from_params(cookie_params)
  180. return path_params + query_params + header_params + cookie_params
  181. def _get_signature(call: Callable[..., Any]) -> inspect.Signature:
  182. try:
  183. signature = inspect.signature(call, eval_str=True)
  184. except NameError:
  185. # Handle type annotations with if TYPE_CHECKING, not used by FastAPI
  186. # e.g. dependency return types
  187. if sys.version_info >= (3, 14):
  188. from annotationlib import Format
  189. signature = inspect.signature(call, annotation_format=Format.FORWARDREF)
  190. else:
  191. signature = inspect.signature(call)
  192. return signature
  193. def get_typed_signature(call: Callable[..., Any]) -> inspect.Signature:
  194. signature = _get_signature(call)
  195. unwrapped = inspect.unwrap(call)
  196. globalns = getattr(unwrapped, "__globals__", {})
  197. typed_params = [
  198. inspect.Parameter(
  199. name=param.name,
  200. kind=param.kind,
  201. default=param.default,
  202. annotation=get_typed_annotation(param.annotation, globalns),
  203. )
  204. for param in signature.parameters.values()
  205. ]
  206. typed_signature = inspect.Signature(typed_params)
  207. return typed_signature
  208. def get_typed_annotation(annotation: Any, globalns: dict[str, Any]) -> Any:
  209. if isinstance(annotation, str):
  210. annotation = ForwardRef(annotation)
  211. annotation = evaluate_forwardref(annotation, globalns, globalns)
  212. if annotation is type(None):
  213. return None
  214. return annotation
  215. def get_typed_return_annotation(call: Callable[..., Any]) -> Any:
  216. signature = _get_signature(call)
  217. unwrapped = inspect.unwrap(call)
  218. annotation = signature.return_annotation
  219. if annotation is inspect.Signature.empty:
  220. return None
  221. globalns = getattr(unwrapped, "__globals__", {})
  222. return get_typed_annotation(annotation, globalns)
  223. _STREAM_ORIGINS = {
  224. AsyncIterable,
  225. AsyncIterator,
  226. AsyncGenerator,
  227. Iterable,
  228. Iterator,
  229. Generator,
  230. }
  231. def get_stream_item_type(annotation: Any) -> Any | None:
  232. origin = get_origin(annotation)
  233. if origin is not None and origin in _STREAM_ORIGINS:
  234. type_args = get_args(annotation)
  235. if type_args:
  236. return type_args[0]
  237. return Any
  238. return None
  239. def get_dependant(
  240. *,
  241. path: str,
  242. call: Callable[..., Any],
  243. name: str | None = None,
  244. own_oauth_scopes: list[str] | None = None,
  245. parent_oauth_scopes: list[str] | None = None,
  246. use_cache: bool = True,
  247. scope: Literal["function", "request"] | None = None,
  248. ) -> Dependant:
  249. dependant = Dependant(
  250. call=call,
  251. name=name,
  252. path=path,
  253. use_cache=use_cache,
  254. scope=scope,
  255. own_oauth_scopes=own_oauth_scopes,
  256. parent_oauth_scopes=parent_oauth_scopes,
  257. )
  258. current_scopes = (parent_oauth_scopes or []) + (own_oauth_scopes or [])
  259. path_param_names = get_path_param_names(path)
  260. endpoint_signature = get_typed_signature(call)
  261. signature_params = endpoint_signature.parameters
  262. for param_name, param in signature_params.items():
  263. is_path_param = param_name in path_param_names
  264. param_details = analyze_param(
  265. param_name=param_name,
  266. annotation=param.annotation,
  267. value=param.default,
  268. is_path_param=is_path_param,
  269. )
  270. if param_details.depends is not None:
  271. assert param_details.depends.dependency
  272. if (
  273. (
  274. _is_gen_callable(dependant.call)
  275. or _is_async_gen_callable(dependant.call)
  276. )
  277. and _get_computed_scope(dependant=dependant) == "request"
  278. and param_details.depends.scope == "function"
  279. ):
  280. assert dependant.call
  281. call_name = getattr(dependant.call, "__name__", "<unnamed_callable>")
  282. raise DependencyScopeError(
  283. f'The dependency "{call_name}" has a scope of '
  284. '"request", it cannot depend on dependencies with scope "function".'
  285. )
  286. sub_own_oauth_scopes: list[str] = []
  287. if isinstance(param_details.depends, params.Security):
  288. if param_details.depends.scopes:
  289. sub_own_oauth_scopes = list(param_details.depends.scopes)
  290. sub_dependant = get_dependant(
  291. path=path,
  292. call=param_details.depends.dependency,
  293. name=param_name,
  294. own_oauth_scopes=sub_own_oauth_scopes,
  295. parent_oauth_scopes=current_scopes,
  296. use_cache=param_details.depends.use_cache,
  297. scope=param_details.depends.scope,
  298. )
  299. dependant.dependencies.append(sub_dependant)
  300. continue
  301. if add_non_field_param_to_dependency(
  302. param_name=param_name,
  303. type_annotation=param_details.type_annotation,
  304. dependant=dependant,
  305. ):
  306. assert param_details.field is None, (
  307. f"Cannot specify multiple FastAPI annotations for {param_name!r}"
  308. )
  309. continue
  310. assert param_details.field is not None
  311. if isinstance(param_details.field.field_info, params.Body):
  312. dependant.body_params.append(param_details.field)
  313. else:
  314. add_param_to_fields(field=param_details.field, dependant=dependant)
  315. return dependant
  316. def add_non_field_param_to_dependency(
  317. *, param_name: str, type_annotation: Any, dependant: Dependant
  318. ) -> bool | None:
  319. if lenient_issubclass(type_annotation, Request):
  320. dependant.request_param_name = param_name
  321. return True
  322. elif lenient_issubclass(type_annotation, WebSocket):
  323. dependant.websocket_param_name = param_name
  324. return True
  325. elif lenient_issubclass(type_annotation, HTTPConnection):
  326. dependant.http_connection_param_name = param_name
  327. return True
  328. elif lenient_issubclass(type_annotation, Response):
  329. dependant.response_param_name = param_name
  330. return True
  331. elif lenient_issubclass(type_annotation, StarletteBackgroundTasks):
  332. dependant.background_tasks_param_name = param_name
  333. return True
  334. elif lenient_issubclass(type_annotation, SecurityScopes):
  335. dependant.security_scopes_param_name = param_name
  336. return True
  337. return None
  338. @dataclass
  339. class ParamDetails:
  340. type_annotation: Any
  341. depends: params.Depends | None
  342. field: ModelField | None
  343. def analyze_param(
  344. *,
  345. param_name: str,
  346. annotation: Any,
  347. value: Any,
  348. is_path_param: bool,
  349. ) -> ParamDetails:
  350. field_info = None
  351. depends = None
  352. type_annotation: Any = Any
  353. use_annotation: Any = Any
  354. if is_typealiastype(annotation):
  355. # unpack in case PEP 695 type syntax is used
  356. annotation = annotation.__value__
  357. if annotation is not inspect.Signature.empty:
  358. use_annotation = annotation
  359. type_annotation = annotation
  360. # Extract Annotated info
  361. if get_origin(use_annotation) is Annotated:
  362. annotated_args = get_args(annotation)
  363. type_annotation = annotated_args[0]
  364. fastapi_annotations = [
  365. arg
  366. for arg in annotated_args[1:]
  367. if isinstance(arg, (FieldInfo, params.Depends))
  368. ]
  369. fastapi_specific_annotations = [
  370. arg
  371. for arg in fastapi_annotations
  372. if isinstance(
  373. arg,
  374. (
  375. params.Param,
  376. params.Body,
  377. params.Depends,
  378. ),
  379. )
  380. ]
  381. if fastapi_specific_annotations:
  382. fastapi_annotation: FieldInfo | params.Depends | None = (
  383. fastapi_specific_annotations[-1]
  384. )
  385. else:
  386. fastapi_annotation = None
  387. # Set default for Annotated FieldInfo
  388. if isinstance(fastapi_annotation, FieldInfo):
  389. # Copy `field_info` because we mutate `field_info.default` below.
  390. field_info = copy_field_info(
  391. field_info=fastapi_annotation,
  392. annotation=use_annotation,
  393. )
  394. assert (
  395. field_info.default == Undefined or field_info.default == RequiredParam
  396. ), (
  397. f"`{field_info.__class__.__name__}` default value cannot be set in"
  398. f" `Annotated` for {param_name!r}. Set the default value with `=` instead."
  399. )
  400. if value is not inspect.Signature.empty:
  401. assert not is_path_param, "Path parameters cannot have default values"
  402. field_info.default = value
  403. else:
  404. field_info.default = RequiredParam
  405. # Get Annotated Depends
  406. elif isinstance(fastapi_annotation, params.Depends):
  407. depends = fastapi_annotation
  408. # Get Depends from default value
  409. if isinstance(value, params.Depends):
  410. assert depends is None, (
  411. "Cannot specify `Depends` in `Annotated` and default value"
  412. f" together for {param_name!r}"
  413. )
  414. assert field_info is None, (
  415. "Cannot specify a FastAPI annotation in `Annotated` and `Depends` as a"
  416. f" default value together for {param_name!r}"
  417. )
  418. depends = value
  419. # Get FieldInfo from default value
  420. elif isinstance(value, FieldInfo):
  421. assert field_info is None, (
  422. "Cannot specify FastAPI annotations in `Annotated` and default value"
  423. f" together for {param_name!r}"
  424. )
  425. field_info = value
  426. if isinstance(field_info, FieldInfo):
  427. field_info.annotation = type_annotation
  428. # Get Depends from type annotation
  429. if depends is not None and depends.dependency is None:
  430. # Copy `depends` before mutating it
  431. depends = copy(depends)
  432. depends = dataclasses.replace(depends, dependency=type_annotation)
  433. # Handle non-param type annotations like Request
  434. # Only apply special handling when there's no explicit Depends - if there's a Depends,
  435. # the dependency will be called and its return value used instead of the special injection
  436. if depends is None and lenient_issubclass(
  437. type_annotation,
  438. (
  439. Request,
  440. WebSocket,
  441. HTTPConnection,
  442. Response,
  443. StarletteBackgroundTasks,
  444. SecurityScopes,
  445. ),
  446. ):
  447. assert field_info is None, (
  448. f"Cannot specify FastAPI annotation for type {type_annotation!r}"
  449. )
  450. # Handle default assignations, neither field_info nor depends was not found in Annotated nor default value
  451. elif field_info is None and depends is None:
  452. default_value = value if value is not inspect.Signature.empty else RequiredParam
  453. if is_path_param:
  454. # We might check here that `default_value is RequiredParam`, but the fact is that the same
  455. # parameter might sometimes be a path parameter and sometimes not. See
  456. # `tests/test_infer_param_optionality.py` for an example.
  457. field_info = params.Path(annotation=use_annotation)
  458. elif is_uploadfile_or_nonable_uploadfile_annotation(
  459. type_annotation
  460. ) or is_uploadfile_sequence_annotation(type_annotation):
  461. field_info = params.File(annotation=use_annotation, default=default_value)
  462. elif not field_annotation_is_scalar(annotation=type_annotation):
  463. field_info = params.Body(annotation=use_annotation, default=default_value)
  464. else:
  465. field_info = params.Query(annotation=use_annotation, default=default_value)
  466. field = None
  467. # It's a field_info, not a dependency
  468. if field_info is not None:
  469. # Handle field_info.in_
  470. if is_path_param:
  471. assert isinstance(field_info, params.Path), (
  472. f"Cannot use `{field_info.__class__.__name__}` for path param"
  473. f" {param_name!r}"
  474. )
  475. elif (
  476. isinstance(field_info, params.Param)
  477. and getattr(field_info, "in_", None) is None
  478. ):
  479. field_info.in_ = params.ParamTypes.query
  480. use_annotation_from_field_info = use_annotation
  481. if isinstance(field_info, params.Form):
  482. ensure_multipart_is_installed()
  483. if not field_info.alias and getattr(field_info, "convert_underscores", None):
  484. alias = param_name.replace("_", "-")
  485. else:
  486. alias = field_info.alias or param_name
  487. field_info.alias = alias
  488. field = create_model_field(
  489. name=param_name,
  490. type_=use_annotation_from_field_info,
  491. default=field_info.default,
  492. alias=alias,
  493. field_info=field_info,
  494. )
  495. if is_path_param:
  496. assert is_scalar_field(field=field), (
  497. "Path params must be of one of the supported types"
  498. )
  499. elif isinstance(field_info, params.Query):
  500. assert (
  501. is_scalar_field(field)
  502. or field_annotation_is_scalar_sequence(field.field_info.annotation)
  503. or lenient_issubclass(field.field_info.annotation, BaseModel)
  504. ), f"Query parameter {param_name!r} must be one of the supported types"
  505. return ParamDetails(type_annotation=type_annotation, depends=depends, field=field)
  506. def add_param_to_fields(*, field: ModelField, dependant: Dependant) -> None:
  507. field_info = field.field_info
  508. field_info_in = getattr(field_info, "in_", None)
  509. if field_info_in == params.ParamTypes.path:
  510. dependant.path_params.append(field)
  511. elif field_info_in == params.ParamTypes.query:
  512. dependant.query_params.append(field)
  513. elif field_info_in == params.ParamTypes.header:
  514. dependant.header_params.append(field)
  515. else:
  516. assert field_info_in == params.ParamTypes.cookie, (
  517. f"non-body parameters must be in path, query, header or cookie: {field.name}"
  518. )
  519. dependant.cookie_params.append(field)
  520. async def _solve_generator(
  521. *, dependant: Dependant, stack: AsyncExitStack, sub_values: dict[str, Any]
  522. ) -> Any:
  523. assert dependant.call
  524. if _is_async_gen_callable(dependant.call):
  525. cm = asynccontextmanager(dependant.call)(**sub_values)
  526. elif _is_gen_callable(dependant.call):
  527. cm = contextmanager_in_threadpool(contextmanager(dependant.call)(**sub_values))
  528. return await stack.enter_async_context(cm)
  529. @dataclass
  530. class SolvedDependency:
  531. values: dict[str, Any]
  532. errors: list[Any]
  533. background_tasks: StarletteBackgroundTasks | None
  534. response: Response
  535. dependency_cache: dict[DependencyCacheKey, Any]
  536. async def solve_dependencies(
  537. *,
  538. request: Request | WebSocket,
  539. dependant: Dependant,
  540. body: dict[str, Any] | FormData | bytes | None = None,
  541. background_tasks: StarletteBackgroundTasks | None = None,
  542. response: Response | None = None,
  543. dependency_overrides_provider: Any | None = None,
  544. dependency_cache: dict[DependencyCacheKey, Any] | None = None,
  545. # TODO: remove this parameter later, no longer used, not removing it yet as some
  546. # people might be monkey patching this function (although that's not supported)
  547. async_exit_stack: AsyncExitStack,
  548. embed_body_fields: bool,
  549. _uses_scopes_cache: _UsesScopesCache | None = None,
  550. ) -> SolvedDependency:
  551. request_astack = request.scope.get("fastapi_inner_astack")
  552. assert isinstance(request_astack, AsyncExitStack), (
  553. "fastapi_inner_astack not found in request scope"
  554. )
  555. function_astack = request.scope.get("fastapi_function_astack")
  556. assert isinstance(function_astack, AsyncExitStack), (
  557. "fastapi_function_astack not found in request scope"
  558. )
  559. values: dict[str, Any] = {}
  560. errors: list[Any] = []
  561. if response is None:
  562. response = Response()
  563. del response.headers["content-length"]
  564. response.status_code = None # type: ignore
  565. if dependency_cache is None:
  566. dependency_cache = {}
  567. if _uses_scopes_cache is None:
  568. _uses_scopes_cache = {}
  569. for sub_dependant in dependant.dependencies:
  570. sub_dependant.call = cast(Callable[..., Any], sub_dependant.call)
  571. call = sub_dependant.call
  572. use_sub_dependant = sub_dependant
  573. if (
  574. dependency_overrides_provider
  575. and dependency_overrides_provider.dependency_overrides
  576. ):
  577. original_call = sub_dependant.call
  578. call = getattr(
  579. dependency_overrides_provider, "dependency_overrides", {}
  580. ).get(original_call, original_call)
  581. use_path: str = sub_dependant.path # type: ignore
  582. use_sub_dependant = get_dependant(
  583. path=use_path,
  584. call=call,
  585. name=sub_dependant.name,
  586. parent_oauth_scopes=_get_oauth_scopes(dependant=sub_dependant),
  587. scope=sub_dependant.scope,
  588. )
  589. solved_result = await solve_dependencies(
  590. request=request,
  591. dependant=use_sub_dependant,
  592. body=body,
  593. background_tasks=background_tasks,
  594. response=response,
  595. dependency_overrides_provider=dependency_overrides_provider,
  596. dependency_cache=dependency_cache,
  597. async_exit_stack=async_exit_stack,
  598. embed_body_fields=embed_body_fields,
  599. _uses_scopes_cache=_uses_scopes_cache,
  600. )
  601. background_tasks = solved_result.background_tasks
  602. if solved_result.errors:
  603. errors.extend(solved_result.errors)
  604. continue
  605. sub_dependant_cache_key = _get_cache_key(
  606. dependant=sub_dependant,
  607. uses_scopes_cache=_uses_scopes_cache,
  608. )
  609. if sub_dependant.use_cache and sub_dependant_cache_key in dependency_cache:
  610. solved = dependency_cache[sub_dependant_cache_key]
  611. elif _is_gen_callable(use_sub_dependant.call) or _is_async_gen_callable(
  612. use_sub_dependant.call
  613. ):
  614. use_astack = request_astack
  615. if sub_dependant.scope == "function":
  616. use_astack = function_astack
  617. solved = await _solve_generator(
  618. dependant=use_sub_dependant,
  619. stack=use_astack,
  620. sub_values=solved_result.values,
  621. )
  622. elif _is_coroutine_callable(use_sub_dependant.call):
  623. solved = await call(**solved_result.values)
  624. else:
  625. solved = await run_in_threadpool(call, **solved_result.values)
  626. if sub_dependant.name is not None:
  627. values[sub_dependant.name] = solved
  628. if sub_dependant_cache_key not in dependency_cache:
  629. dependency_cache[sub_dependant_cache_key] = solved
  630. path_values, path_errors = request_params_to_args(
  631. dependant.path_params, request.path_params
  632. )
  633. query_values, query_errors = request_params_to_args(
  634. dependant.query_params, request.query_params
  635. )
  636. header_values, header_errors = request_params_to_args(
  637. dependant.header_params, request.headers
  638. )
  639. cookie_values, cookie_errors = request_params_to_args(
  640. dependant.cookie_params, request.cookies
  641. )
  642. values.update(path_values)
  643. values.update(query_values)
  644. values.update(header_values)
  645. values.update(cookie_values)
  646. errors += path_errors + query_errors + header_errors + cookie_errors
  647. if dependant.body_params:
  648. (
  649. body_values,
  650. body_errors,
  651. ) = await request_body_to_args( # body_params checked above
  652. body_fields=dependant.body_params,
  653. received_body=body,
  654. embed_body_fields=embed_body_fields,
  655. )
  656. values.update(body_values)
  657. errors.extend(body_errors)
  658. if dependant.http_connection_param_name:
  659. values[dependant.http_connection_param_name] = request
  660. if dependant.request_param_name and isinstance(request, Request):
  661. values[dependant.request_param_name] = request
  662. elif dependant.websocket_param_name and isinstance(request, WebSocket):
  663. values[dependant.websocket_param_name] = request
  664. if dependant.background_tasks_param_name:
  665. if background_tasks is None:
  666. background_tasks = BackgroundTasks()
  667. values[dependant.background_tasks_param_name] = background_tasks
  668. if dependant.response_param_name:
  669. values[dependant.response_param_name] = response
  670. if dependant.security_scopes_param_name:
  671. values[dependant.security_scopes_param_name] = SecurityScopes(
  672. scopes=_get_oauth_scopes(dependant=dependant)
  673. )
  674. return SolvedDependency(
  675. values=values,
  676. errors=errors,
  677. background_tasks=background_tasks,
  678. response=response,
  679. dependency_cache=dependency_cache,
  680. )
  681. def _validate_value_with_model_field(
  682. *, field: ModelField, value: Any, values: dict[str, Any], loc: tuple[str, ...]
  683. ) -> tuple[Any, list[Any]]:
  684. if value is None:
  685. if field.field_info.is_required():
  686. return None, [get_missing_field_error(loc=loc)]
  687. else:
  688. return deepcopy(field.default), []
  689. return field.validate(value, values, loc=loc)
  690. def _is_json_field(field: ModelField) -> bool:
  691. return any(type(item) is Json for item in field.field_info.metadata)
  692. def _get_multidict_value(
  693. field: ModelField, values: Mapping[str, Any], alias: str | None = None
  694. ) -> Any:
  695. alias = alias or get_validation_alias(field)
  696. if (
  697. (not _is_json_field(field))
  698. and field_annotation_is_sequence(field.field_info.annotation)
  699. and isinstance(values, (ImmutableMultiDict, Headers))
  700. ):
  701. value = values.getlist(alias)
  702. else:
  703. value = values.get(alias, None)
  704. if (
  705. value is None
  706. or (
  707. isinstance(field.field_info, params.Form)
  708. and isinstance(value, str) # For type checks
  709. and value == ""
  710. )
  711. or (
  712. field_annotation_is_sequence(field.field_info.annotation)
  713. and len(value) == 0
  714. )
  715. ):
  716. if field.field_info.is_required():
  717. return
  718. else:
  719. return deepcopy(field.default)
  720. return value
  721. def request_params_to_args(
  722. fields: Sequence[ModelField],
  723. received_params: Mapping[str, Any] | QueryParams | Headers,
  724. ) -> tuple[dict[str, Any], list[Any]]:
  725. values: dict[str, Any] = {}
  726. errors: list[dict[str, Any]] = []
  727. if not fields:
  728. return values, errors
  729. first_field = fields[0]
  730. fields_to_extract = fields
  731. single_not_embedded_field = False
  732. default_convert_underscores = True
  733. if len(fields) == 1 and lenient_issubclass(
  734. first_field.field_info.annotation, BaseModel
  735. ):
  736. fields_to_extract = get_cached_model_fields(first_field.field_info.annotation)
  737. single_not_embedded_field = True
  738. # If headers are in a Pydantic model, the way to disable convert_underscores
  739. # would be with Header(convert_underscores=False) at the Pydantic model level
  740. default_convert_underscores = getattr(
  741. first_field.field_info, "convert_underscores", True
  742. )
  743. params_to_process: dict[str, Any] = {}
  744. processed_keys = set()
  745. for field in fields_to_extract:
  746. alias = None
  747. if isinstance(received_params, Headers):
  748. # Handle fields extracted from a Pydantic Model for a header, each field
  749. # doesn't have a FieldInfo of type Header with the default convert_underscores=True
  750. convert_underscores = getattr(
  751. field.field_info, "convert_underscores", default_convert_underscores
  752. )
  753. if convert_underscores:
  754. alias = get_validation_alias(field)
  755. if alias == field.name:
  756. alias = alias.replace("_", "-")
  757. value = _get_multidict_value(field, received_params, alias=alias)
  758. if value is not None:
  759. params_to_process[get_validation_alias(field)] = value
  760. processed_keys.add(alias or get_validation_alias(field))
  761. # For headers with convert_underscores=True, mark both the converted
  762. # header name and the original field alias as processed to avoid
  763. # accepting the original alias as an extra header.
  764. processed_keys.add(get_validation_alias(field))
  765. for key in received_params.keys():
  766. if key not in processed_keys:
  767. if isinstance(received_params, (ImmutableMultiDict, Headers)):
  768. value = received_params.getlist(key)
  769. if isinstance(value, list) and (len(value) == 1):
  770. params_to_process[key] = value[0]
  771. else:
  772. params_to_process[key] = value
  773. else:
  774. params_to_process[key] = received_params.get(key)
  775. if single_not_embedded_field:
  776. field_info = first_field.field_info
  777. assert isinstance(field_info, params.Param), (
  778. "Params must be subclasses of Param"
  779. )
  780. loc: tuple[str, ...] = (field_info.in_.value,)
  781. v_, errors_ = _validate_value_with_model_field(
  782. field=first_field, value=params_to_process, values=values, loc=loc
  783. )
  784. return {first_field.name: v_}, errors_
  785. for field in fields:
  786. value = _get_multidict_value(field, received_params)
  787. field_info = field.field_info
  788. assert isinstance(field_info, params.Param), (
  789. "Params must be subclasses of Param"
  790. )
  791. loc = (field_info.in_.value, get_validation_alias(field))
  792. v_, errors_ = _validate_value_with_model_field(
  793. field=field, value=value, values=values, loc=loc
  794. )
  795. if errors_:
  796. errors.extend(errors_)
  797. else:
  798. values[field.name] = v_
  799. return values, errors
  800. def is_union_of_base_models(field_type: Any) -> bool:
  801. """Check if field type is a Union where all members are BaseModel subclasses."""
  802. from fastapi.types import UnionType
  803. origin = get_origin(field_type)
  804. # Check if it's a Union type (covers both typing.Union and types.UnionType in Python 3.10+)
  805. if origin is not Union and origin is not UnionType:
  806. return False
  807. union_args = get_args(field_type)
  808. for arg in union_args:
  809. if not lenient_issubclass(arg, BaseModel):
  810. return False
  811. return True
  812. def _should_embed_body_fields(fields: list[ModelField]) -> bool:
  813. if not fields:
  814. return False
  815. # More than one dependency could have the same field, it would show up as multiple
  816. # fields but it's the same one, so count them by name
  817. body_param_names_set = {field.name for field in fields}
  818. # A top level field has to be a single field, not multiple
  819. if len(body_param_names_set) > 1:
  820. return True
  821. first_field = fields[0]
  822. # If it explicitly specifies it is embedded, it has to be embedded
  823. if getattr(first_field.field_info, "embed", None):
  824. return True
  825. # If it's a Form (or File) field, it has to be a BaseModel (or a union of BaseModels) to be top level
  826. # otherwise it has to be embedded, so that the key value pair can be extracted
  827. if (
  828. isinstance(first_field.field_info, params.Form)
  829. and not lenient_issubclass(first_field.field_info.annotation, BaseModel)
  830. and not is_union_of_base_models(first_field.field_info.annotation)
  831. ):
  832. return True
  833. return False
  834. async def _extract_form_body(
  835. body_fields: list[ModelField],
  836. received_body: FormData,
  837. ) -> dict[str, Any]:
  838. values = {}
  839. for field in body_fields:
  840. value = _get_multidict_value(field, received_body)
  841. field_info = field.field_info
  842. if (
  843. isinstance(field_info, params.File)
  844. and is_bytes_or_nonable_bytes_annotation(field.field_info.annotation)
  845. and isinstance(value, UploadFile)
  846. ):
  847. value = await value.read()
  848. elif (
  849. is_bytes_sequence_annotation(field.field_info.annotation)
  850. and isinstance(field_info, params.File)
  851. and value_is_sequence(value)
  852. ):
  853. # For types
  854. assert isinstance(value, sequence_types)
  855. results: list[bytes | str] = []
  856. for sub_value in value:
  857. results.append(await sub_value.read())
  858. value = serialize_sequence_value(field=field, value=results)
  859. if value is not None:
  860. values[get_validation_alias(field)] = value
  861. field_aliases = {get_validation_alias(field) for field in body_fields}
  862. for key in received_body.keys():
  863. if key not in field_aliases:
  864. param_values = received_body.getlist(key)
  865. if len(param_values) == 1:
  866. values[key] = param_values[0]
  867. else:
  868. values[key] = param_values
  869. return values
  870. async def request_body_to_args(
  871. body_fields: list[ModelField],
  872. received_body: dict[str, Any] | FormData | bytes | None,
  873. embed_body_fields: bool,
  874. ) -> tuple[dict[str, Any], list[dict[str, Any]]]:
  875. values: dict[str, Any] = {}
  876. errors: list[dict[str, Any]] = []
  877. assert body_fields, "request_body_to_args() should be called with fields"
  878. single_not_embedded_field = len(body_fields) == 1 and not embed_body_fields
  879. first_field = body_fields[0]
  880. body_to_process = received_body
  881. fields_to_extract: list[ModelField] = body_fields
  882. if (
  883. single_not_embedded_field
  884. and lenient_issubclass(first_field.field_info.annotation, BaseModel)
  885. and isinstance(received_body, FormData)
  886. ):
  887. fields_to_extract = get_cached_model_fields(first_field.field_info.annotation)
  888. if isinstance(received_body, FormData):
  889. body_to_process = await _extract_form_body(fields_to_extract, received_body)
  890. if single_not_embedded_field:
  891. loc: tuple[str, ...] = ("body",)
  892. v_, errors_ = _validate_value_with_model_field(
  893. field=first_field, value=body_to_process, values=values, loc=loc
  894. )
  895. return {first_field.name: v_}, errors_
  896. for field in body_fields:
  897. loc = ("body", get_validation_alias(field))
  898. value: Any | None = None
  899. if body_to_process is not None and not isinstance(body_to_process, bytes):
  900. try:
  901. value = body_to_process.get(get_validation_alias(field))
  902. # If the received body is a list, not a dict
  903. except AttributeError:
  904. errors.append(get_missing_field_error(loc))
  905. continue
  906. v_, errors_ = _validate_value_with_model_field(
  907. field=field, value=value, values=values, loc=loc
  908. )
  909. if errors_:
  910. errors.extend(errors_)
  911. else:
  912. values[field.name] = v_
  913. return values, errors
  914. def _get_body_field(
  915. *, body_params: list[ModelField], name: str, embed_body_fields: bool
  916. ) -> ModelField | None:
  917. """
  918. Get a ModelField representing the request body for a path operation, combining
  919. all body parameters into a single field if necessary.
  920. Used to check if it's form data (with `isinstance(body_field, params.Form)`)
  921. or JSON and to generate the JSON Schema for a request body.
  922. This is **not** used to validate/parse the request body, that's done with each
  923. individual body parameter.
  924. """
  925. if not body_params:
  926. return None
  927. first_param = body_params[0]
  928. if not embed_body_fields:
  929. return first_param
  930. model_name = "Body_" + name
  931. BodyModel = create_body_model(fields=body_params, model_name=model_name)
  932. required = any(True for f in body_params if f.field_info.is_required())
  933. BodyFieldInfo_kwargs: dict[str, Any] = {
  934. "annotation": BodyModel,
  935. "alias": "body",
  936. }
  937. if not required:
  938. BodyFieldInfo_kwargs["default"] = None
  939. if any(isinstance(f.field_info, params.File) for f in body_params):
  940. BodyFieldInfo: type[params.Body] = params.File
  941. elif any(isinstance(f.field_info, params.Form) for f in body_params):
  942. BodyFieldInfo = params.Form
  943. else:
  944. BodyFieldInfo = params.Body
  945. body_param_media_types = [
  946. f.field_info.media_type
  947. for f in body_params
  948. if isinstance(f.field_info, params.Body)
  949. ]
  950. if len(set(body_param_media_types)) == 1:
  951. BodyFieldInfo_kwargs["media_type"] = body_param_media_types[0]
  952. final_field = create_model_field(
  953. name="body",
  954. type_=BodyModel,
  955. alias="body",
  956. field_info=BodyFieldInfo(**BodyFieldInfo_kwargs),
  957. )
  958. return final_field
  959. def get_validation_alias(field: ModelField) -> str:
  960. va = getattr(field, "validation_alias", None)
  961. return va or field.alias