utils.py 29 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679
  1. import copy
  2. import http.client
  3. import inspect
  4. import warnings
  5. from collections.abc import Sequence
  6. from dataclasses import dataclass, field
  7. from typing import Any, Literal, cast
  8. from fastapi import routing
  9. from fastapi._compat import (
  10. ModelField,
  11. get_definitions,
  12. get_flat_models_from_fields,
  13. get_model_name_map,
  14. get_schema_from_model_field,
  15. lenient_issubclass,
  16. )
  17. from fastapi.datastructures import DefaultPlaceholder, _Unset
  18. from fastapi.dependencies.models import (
  19. Dependant,
  20. _get_cache_key,
  21. _get_oauth_scopes,
  22. _get_security_scheme,
  23. _is_security_scheme,
  24. _UsesScopesCache,
  25. )
  26. from fastapi.dependencies.utils import (
  27. _get_flat_fields_from_params,
  28. get_flat_params,
  29. get_validation_alias,
  30. )
  31. from fastapi.encoders import jsonable_encoder
  32. from fastapi.exceptions import FastAPIDeprecationWarning
  33. from fastapi.openapi.constants import METHODS_WITH_BODY, REF_PREFIX
  34. from fastapi.openapi.models import OpenAPI
  35. from fastapi.params import Body, ParamTypes
  36. from fastapi.responses import Response
  37. from fastapi.sse import _SSE_EVENT_SCHEMA
  38. from fastapi.types import DependencyCacheKey, ModelNameMap
  39. from fastapi.utils import (
  40. deep_dict_update,
  41. generate_operation_id_for_path,
  42. is_body_allowed_for_status_code,
  43. )
  44. from pydantic import BaseModel
  45. from starlette.responses import JSONResponse
  46. from starlette.routing import BaseRoute
  47. validation_error_definition = {
  48. "title": "ValidationError",
  49. "type": "object",
  50. "properties": {
  51. "loc": {
  52. "title": "Location",
  53. "type": "array",
  54. "items": {"anyOf": [{"type": "string"}, {"type": "integer"}]},
  55. },
  56. "msg": {"title": "Message", "type": "string"},
  57. "type": {"title": "Error Type", "type": "string"},
  58. "input": {"title": "Input"},
  59. "ctx": {"title": "Context", "type": "object"},
  60. },
  61. "required": ["loc", "msg", "type"],
  62. }
  63. validation_error_response_definition = {
  64. "title": "HTTPValidationError",
  65. "type": "object",
  66. "properties": {
  67. "detail": {
  68. "title": "Detail",
  69. "type": "array",
  70. "items": {"$ref": REF_PREFIX + "ValidationError"},
  71. }
  72. },
  73. }
  74. status_code_ranges: dict[str, str] = {
  75. "1XX": "Information",
  76. "2XX": "Success",
  77. "3XX": "Redirection",
  78. "4XX": "Client Error",
  79. "5XX": "Server Error",
  80. "DEFAULT": "Default Response",
  81. }
  82. @dataclass
  83. class _OpenAPIDependencyData:
  84. path_params: list[ModelField] = field(default_factory=list)
  85. query_params: list[ModelField] = field(default_factory=list)
  86. header_params: list[ModelField] = field(default_factory=list)
  87. cookie_params: list[ModelField] = field(default_factory=list)
  88. security_dependencies: list[tuple[Dependant, list[str]]] = field(
  89. default_factory=list
  90. )
  91. def _get_openapi_dependency_data(dependant: Dependant) -> _OpenAPIDependencyData:
  92. dependency_data = _OpenAPIDependencyData()
  93. visited: list[DependencyCacheKey] = []
  94. uses_scopes_cache: _UsesScopesCache = {}
  95. dependants: list[tuple[Dependant, list[str], bool]] = [(dependant, [], True)]
  96. while dependants:
  97. current_dependant, parent_oauth_scopes, is_root = dependants.pop()
  98. cache_key = _get_cache_key(
  99. dependant=current_dependant,
  100. uses_scopes_cache=uses_scopes_cache,
  101. )
  102. if cache_key in visited:
  103. continue
  104. visited.append(cache_key)
  105. dependency_data.path_params.extend(current_dependant.path_params)
  106. dependency_data.query_params.extend(current_dependant.query_params)
  107. dependency_data.header_params.extend(current_dependant.header_params)
  108. dependency_data.cookie_params.extend(current_dependant.cookie_params)
  109. oauth_scopes = parent_oauth_scopes.copy()
  110. for scope in _get_oauth_scopes(dependant=current_dependant):
  111. if scope not in oauth_scopes:
  112. oauth_scopes.append(scope)
  113. if not is_root and _is_security_scheme(dependant=current_dependant):
  114. dependency_data.security_dependencies.append(
  115. (current_dependant, oauth_scopes)
  116. )
  117. dependants.extend(
  118. (sub_dependant, oauth_scopes, False)
  119. for sub_dependant in reversed(current_dependant.dependencies)
  120. )
  121. return dependency_data
  122. def _get_openapi_security_definitions(
  123. security_dependencies: list[tuple[Dependant, list[str]]],
  124. ) -> tuple[dict[str, Any], list[dict[str, Any]]]:
  125. security_definitions = {}
  126. # Use a dict to merge scopes for same security scheme
  127. operation_security_dict: dict[str, list[str]] = {}
  128. for security_dependency, oauth_scopes in security_dependencies:
  129. security_scheme = _get_security_scheme(dependant=security_dependency)
  130. security_definition = jsonable_encoder(
  131. security_scheme.model,
  132. by_alias=True,
  133. exclude_none=True,
  134. )
  135. security_name = security_scheme.scheme_name
  136. security_definitions[security_name] = security_definition
  137. # Merge scopes for the same security scheme
  138. if security_name not in operation_security_dict:
  139. operation_security_dict[security_name] = []
  140. for scope in oauth_scopes:
  141. if scope not in operation_security_dict[security_name]:
  142. operation_security_dict[security_name].append(scope)
  143. operation_security = [
  144. {name: scopes} for name, scopes in operation_security_dict.items()
  145. ]
  146. return security_definitions, operation_security
  147. def _get_openapi_operation_parameters(
  148. *,
  149. dependency_data: _OpenAPIDependencyData,
  150. model_name_map: ModelNameMap,
  151. field_mapping: dict[
  152. tuple[ModelField, Literal["validation", "serialization"]], dict[str, Any]
  153. ],
  154. separate_input_output_schemas: bool = True,
  155. ) -> list[dict[str, Any]]:
  156. parameters = []
  157. path_params = _get_flat_fields_from_params(dependency_data.path_params)
  158. query_params = _get_flat_fields_from_params(dependency_data.query_params)
  159. header_params = _get_flat_fields_from_params(dependency_data.header_params)
  160. cookie_params = _get_flat_fields_from_params(dependency_data.cookie_params)
  161. parameter_groups = [
  162. (ParamTypes.path, path_params),
  163. (ParamTypes.query, query_params),
  164. (ParamTypes.header, header_params),
  165. (ParamTypes.cookie, cookie_params),
  166. ]
  167. default_convert_underscores = True
  168. if len(dependency_data.header_params) == 1:
  169. first_field = dependency_data.header_params[0]
  170. if lenient_issubclass(first_field.field_info.annotation, BaseModel):
  171. default_convert_underscores = getattr(
  172. first_field.field_info, "convert_underscores", True
  173. )
  174. for param_type, param_group in parameter_groups:
  175. for param in param_group:
  176. field_info = param.field_info
  177. # field_info = cast(Param, field_info)
  178. if not getattr(field_info, "include_in_schema", True):
  179. continue
  180. param_schema = get_schema_from_model_field(
  181. field=param,
  182. model_name_map=model_name_map,
  183. field_mapping=field_mapping,
  184. separate_input_output_schemas=separate_input_output_schemas,
  185. )
  186. name = get_validation_alias(param)
  187. convert_underscores = getattr(
  188. param.field_info,
  189. "convert_underscores",
  190. default_convert_underscores,
  191. )
  192. if (
  193. param_type == ParamTypes.header
  194. and name == param.name
  195. and convert_underscores
  196. ):
  197. name = param.name.replace("_", "-")
  198. parameter = {
  199. "name": name,
  200. "in": param_type.value,
  201. "required": param.field_info.is_required(),
  202. "schema": param_schema,
  203. }
  204. if field_info.description:
  205. parameter["description"] = field_info.description
  206. openapi_examples = getattr(field_info, "openapi_examples", None)
  207. example = getattr(field_info, "example", None)
  208. if openapi_examples:
  209. parameter["examples"] = jsonable_encoder(openapi_examples)
  210. elif example is not _Unset:
  211. parameter["example"] = jsonable_encoder(example)
  212. if getattr(field_info, "deprecated", None):
  213. parameter["deprecated"] = True
  214. parameters.append(parameter)
  215. return parameters
  216. def get_openapi_operation_request_body(
  217. *,
  218. body_field: ModelField | None,
  219. model_name_map: ModelNameMap,
  220. field_mapping: dict[
  221. tuple[ModelField, Literal["validation", "serialization"]], dict[str, Any]
  222. ],
  223. separate_input_output_schemas: bool = True,
  224. ) -> dict[str, Any] | None:
  225. if not body_field:
  226. return None
  227. assert isinstance(body_field, ModelField)
  228. body_schema = get_schema_from_model_field(
  229. field=body_field,
  230. model_name_map=model_name_map,
  231. field_mapping=field_mapping,
  232. separate_input_output_schemas=separate_input_output_schemas,
  233. )
  234. field_info = cast(Body, body_field.field_info)
  235. request_media_type = field_info.media_type
  236. required = body_field.field_info.is_required()
  237. request_body_oai: dict[str, Any] = {}
  238. if required:
  239. request_body_oai["required"] = required
  240. request_media_content: dict[str, Any] = {"schema": body_schema}
  241. if field_info.openapi_examples:
  242. request_media_content["examples"] = jsonable_encoder(
  243. field_info.openapi_examples
  244. )
  245. elif field_info.example is not _Unset:
  246. request_media_content["example"] = jsonable_encoder(field_info.example)
  247. request_body_oai["content"] = {request_media_type: request_media_content}
  248. return request_body_oai
  249. def generate_operation_id(
  250. *, route: routing._APIRouteLike, method: str
  251. ) -> str: # pragma: nocover
  252. warnings.warn(
  253. message="fastapi.openapi.utils.generate_operation_id() was deprecated, "
  254. "it is not used internally, and will be removed soon",
  255. category=FastAPIDeprecationWarning,
  256. stacklevel=2,
  257. )
  258. if route.operation_id:
  259. return route.operation_id
  260. path: str = route.path_format
  261. return generate_operation_id_for_path(name=route.name, path=path, method=method)
  262. def generate_operation_summary(*, route: routing._APIRouteLike, method: str) -> str:
  263. if route.summary:
  264. return route.summary
  265. return route.name.replace("_", " ").title()
  266. def get_openapi_operation_metadata(
  267. *, route: routing._APIRouteLike, method: str, operation_ids: set[str]
  268. ) -> dict[str, Any]:
  269. operation: dict[str, Any] = {}
  270. if route.tags:
  271. operation["tags"] = route.tags
  272. operation["summary"] = generate_operation_summary(route=route, method=method)
  273. if route.description:
  274. operation["description"] = route.description
  275. operation_id = route.operation_id or route.unique_id
  276. if operation_id in operation_ids:
  277. endpoint_name = getattr(route.endpoint, "__name__", "<unnamed_endpoint>")
  278. message = f"Duplicate Operation ID {operation_id} for function {endpoint_name}"
  279. file_name = getattr(route.endpoint, "__globals__", {}).get("__file__")
  280. if file_name:
  281. message += f" at {file_name}"
  282. warnings.warn(message, stacklevel=1)
  283. operation_ids.add(operation_id)
  284. operation["operationId"] = operation_id
  285. if route.deprecated:
  286. operation["deprecated"] = route.deprecated
  287. return operation
  288. def get_openapi_path(
  289. *,
  290. route: routing._APIRouteLike,
  291. operation_ids: set[str],
  292. model_name_map: ModelNameMap,
  293. field_mapping: dict[
  294. tuple[ModelField, Literal["validation", "serialization"]], dict[str, Any]
  295. ],
  296. separate_input_output_schemas: bool = True,
  297. ) -> tuple[dict[str, Any], dict[str, Any], dict[str, Any]]:
  298. path = {}
  299. security_schemes: dict[str, Any] = {}
  300. definitions: dict[str, Any] = {}
  301. assert route.methods is not None, "Methods must be a list"
  302. if isinstance(route.response_class, DefaultPlaceholder):
  303. current_response_class: type[Response] = route.response_class.value
  304. else:
  305. current_response_class = route.response_class
  306. assert current_response_class, "A response class is needed to generate OpenAPI"
  307. route_response_media_type: str | None = current_response_class.media_type
  308. if route.include_in_schema:
  309. dependency_data = _get_openapi_dependency_data(route.dependant)
  310. all_route_params = [
  311. field
  312. for fields in (
  313. dependency_data.path_params,
  314. dependency_data.query_params,
  315. dependency_data.header_params,
  316. dependency_data.cookie_params,
  317. )
  318. for field in _get_flat_fields_from_params(fields)
  319. ]
  320. for method in route.methods:
  321. operation = get_openapi_operation_metadata(
  322. route=route, method=method, operation_ids=operation_ids
  323. )
  324. parameters: list[dict[str, Any]] = []
  325. security_definitions, operation_security = (
  326. _get_openapi_security_definitions(
  327. security_dependencies=dependency_data.security_dependencies
  328. )
  329. )
  330. if operation_security:
  331. operation.setdefault("security", []).extend(operation_security)
  332. if security_definitions:
  333. security_schemes.update(security_definitions)
  334. operation_parameters = _get_openapi_operation_parameters(
  335. dependency_data=dependency_data,
  336. model_name_map=model_name_map,
  337. field_mapping=field_mapping,
  338. separate_input_output_schemas=separate_input_output_schemas,
  339. )
  340. parameters.extend(operation_parameters)
  341. if parameters:
  342. all_parameters = {
  343. (param["in"], param["name"]): param for param in parameters
  344. }
  345. required_parameters = {
  346. (param["in"], param["name"]): param
  347. for param in parameters
  348. if param.get("required")
  349. }
  350. # Make sure required definitions of the same parameter take precedence
  351. # over non-required definitions
  352. all_parameters.update(required_parameters)
  353. operation["parameters"] = list(all_parameters.values())
  354. if method in METHODS_WITH_BODY:
  355. request_body_oai = get_openapi_operation_request_body(
  356. body_field=route.body_field,
  357. model_name_map=model_name_map,
  358. field_mapping=field_mapping,
  359. separate_input_output_schemas=separate_input_output_schemas,
  360. )
  361. if request_body_oai:
  362. operation["requestBody"] = request_body_oai
  363. if route.callbacks:
  364. callbacks = {}
  365. for callback in route.callbacks:
  366. if isinstance(callback, routing.APIRoute):
  367. (
  368. cb_path,
  369. cb_security_schemes,
  370. cb_definitions,
  371. ) = get_openapi_path(
  372. route=cast(routing._APIRouteLike, callback),
  373. operation_ids=operation_ids,
  374. model_name_map=model_name_map,
  375. field_mapping=field_mapping,
  376. separate_input_output_schemas=separate_input_output_schemas,
  377. )
  378. callbacks[callback.name] = {callback.path: cb_path}
  379. operation["callbacks"] = callbacks
  380. if route.status_code is not None:
  381. status_code = str(route.status_code)
  382. else:
  383. # It would probably make more sense for all response classes to have an
  384. # explicit default status_code, and to extract it from them, instead of
  385. # doing this inspection tricks, that would probably be in the future
  386. # TODO: probably make status_code a default class attribute for all
  387. # responses in Starlette
  388. response_signature = inspect.signature(current_response_class.__init__)
  389. status_code_param = response_signature.parameters.get("status_code")
  390. if status_code_param is not None:
  391. if isinstance(status_code_param.default, int):
  392. status_code = str(status_code_param.default)
  393. operation.setdefault("responses", {}).setdefault(status_code, {})[
  394. "description"
  395. ] = route.response_description
  396. if is_body_allowed_for_status_code(route.status_code):
  397. # Check for JSONL streaming (generator endpoints)
  398. if route.is_json_stream:
  399. jsonl_content: dict[str, Any] = {}
  400. if route.stream_item_field:
  401. item_schema = get_schema_from_model_field(
  402. field=route.stream_item_field,
  403. model_name_map=model_name_map,
  404. field_mapping=field_mapping,
  405. separate_input_output_schemas=separate_input_output_schemas,
  406. )
  407. jsonl_content["itemSchema"] = item_schema
  408. else:
  409. jsonl_content["itemSchema"] = {}
  410. operation.setdefault("responses", {}).setdefault(
  411. status_code, {}
  412. ).setdefault("content", {})["application/jsonl"] = jsonl_content
  413. elif route.is_sse_stream:
  414. sse_content: dict[str, Any] = {}
  415. item_schema = copy.deepcopy(_SSE_EVENT_SCHEMA)
  416. if route.stream_item_field:
  417. content_schema = get_schema_from_model_field(
  418. field=route.stream_item_field,
  419. model_name_map=model_name_map,
  420. field_mapping=field_mapping,
  421. separate_input_output_schemas=separate_input_output_schemas,
  422. )
  423. item_schema["required"] = ["data"]
  424. item_schema["properties"]["data"] = {
  425. "type": "string",
  426. "contentMediaType": "application/json",
  427. "contentSchema": content_schema,
  428. }
  429. sse_content["itemSchema"] = item_schema
  430. operation.setdefault("responses", {}).setdefault(
  431. status_code, {}
  432. ).setdefault("content", {})["text/event-stream"] = sse_content
  433. elif route_response_media_type:
  434. response_schema = {"type": "string"}
  435. if lenient_issubclass(current_response_class, JSONResponse):
  436. if route.response_field:
  437. response_schema = get_schema_from_model_field(
  438. field=route.response_field,
  439. model_name_map=model_name_map,
  440. field_mapping=field_mapping,
  441. separate_input_output_schemas=separate_input_output_schemas,
  442. )
  443. else:
  444. response_schema = {}
  445. operation.setdefault("responses", {}).setdefault(
  446. status_code, {}
  447. ).setdefault("content", {}).setdefault(
  448. route_response_media_type, {}
  449. )["schema"] = response_schema
  450. if route.responses:
  451. operation_responses = operation.setdefault("responses", {})
  452. for (
  453. additional_status_code,
  454. additional_response,
  455. ) in route.responses.items():
  456. process_response = copy.deepcopy(additional_response)
  457. process_response.pop("model", None)
  458. status_code_key = str(additional_status_code).upper()
  459. if status_code_key == "DEFAULT":
  460. status_code_key = "default"
  461. openapi_response = operation_responses.setdefault(
  462. status_code_key, {}
  463. )
  464. assert isinstance(process_response, dict), (
  465. "An additional response must be a dict"
  466. )
  467. field = route.response_fields.get(additional_status_code)
  468. additional_field_schema: dict[str, Any] | None = None
  469. if field:
  470. additional_field_schema = get_schema_from_model_field(
  471. field=field,
  472. model_name_map=model_name_map,
  473. field_mapping=field_mapping,
  474. separate_input_output_schemas=separate_input_output_schemas,
  475. )
  476. media_type = route_response_media_type or "application/json"
  477. additional_schema = (
  478. process_response.setdefault("content", {})
  479. .setdefault(media_type, {})
  480. .setdefault("schema", {})
  481. )
  482. deep_dict_update(additional_schema, additional_field_schema)
  483. status_text: str | None = status_code_ranges.get(
  484. str(additional_status_code).upper()
  485. ) or http.client.responses.get(int(additional_status_code))
  486. description = (
  487. process_response.get("description")
  488. or openapi_response.get("description")
  489. or status_text
  490. or "Additional Response"
  491. )
  492. deep_dict_update(openapi_response, process_response)
  493. openapi_response["description"] = description
  494. http422 = "422"
  495. if (all_route_params or route.body_field) and not any(
  496. status in operation["responses"]
  497. for status in [http422, "4XX", "default"]
  498. ):
  499. operation["responses"][http422] = {
  500. "description": "Validation Error",
  501. "content": {
  502. "application/json": {
  503. "schema": {"$ref": REF_PREFIX + "HTTPValidationError"}
  504. }
  505. },
  506. }
  507. if "ValidationError" not in definitions:
  508. definitions.update(
  509. {
  510. "ValidationError": validation_error_definition,
  511. "HTTPValidationError": validation_error_response_definition,
  512. }
  513. )
  514. if route.openapi_extra:
  515. deep_dict_update(operation, route.openapi_extra)
  516. path[method.lower()] = operation
  517. return path, security_schemes, definitions
  518. def _get_api_route_for_openapi(
  519. route_context: routing.RouteContext,
  520. ) -> routing._APIRouteLike | None:
  521. if isinstance(route_context.original_route, routing.APIRoute):
  522. return cast(routing._APIRouteLike, route_context)
  523. return None
  524. def get_fields_from_routes(
  525. routes: Sequence[BaseRoute | routing.RouteContext],
  526. ) -> list[ModelField]:
  527. body_fields_from_routes: list[ModelField] = []
  528. responses_from_routes: list[ModelField] = []
  529. request_fields_from_routes: list[ModelField] = []
  530. callback_flat_models: list[ModelField] = []
  531. for route_context in routing.iter_route_contexts(routes):
  532. api_route = _get_api_route_for_openapi(route_context)
  533. if api_route is None:
  534. continue
  535. if api_route.include_in_schema:
  536. if api_route.body_field:
  537. assert isinstance(api_route.body_field, ModelField), (
  538. "A request body must be a Pydantic Field"
  539. )
  540. body_fields_from_routes.append(api_route.body_field)
  541. if api_route.response_field:
  542. responses_from_routes.append(api_route.response_field)
  543. if api_route.response_fields:
  544. responses_from_routes.extend(api_route.response_fields.values())
  545. if api_route.stream_item_field:
  546. responses_from_routes.append(api_route.stream_item_field)
  547. if api_route.callbacks:
  548. callback_flat_models.extend(get_fields_from_routes(api_route.callbacks))
  549. params = get_flat_params(api_route.dependant)
  550. request_fields_from_routes.extend(params)
  551. flat_models = callback_flat_models + list(
  552. body_fields_from_routes + responses_from_routes + request_fields_from_routes
  553. )
  554. return flat_models
  555. def get_openapi(
  556. *,
  557. title: str,
  558. version: str,
  559. openapi_version: str = "3.1.0",
  560. summary: str | None = None,
  561. description: str | None = None,
  562. routes: Sequence[BaseRoute | routing.RouteContext],
  563. webhooks: Sequence[BaseRoute | routing.RouteContext] | None = None,
  564. tags: list[dict[str, Any]] | None = None,
  565. servers: list[dict[str, str | Any]] | None = None,
  566. terms_of_service: str | None = None,
  567. contact: dict[str, str | Any] | None = None,
  568. license_info: dict[str, str | Any] | None = None,
  569. separate_input_output_schemas: bool = True,
  570. external_docs: dict[str, Any] | None = None,
  571. ) -> dict[str, Any]:
  572. info: dict[str, Any] = {"title": title, "version": version}
  573. if summary:
  574. info["summary"] = summary
  575. if description:
  576. info["description"] = description
  577. if terms_of_service:
  578. info["termsOfService"] = terms_of_service
  579. if contact:
  580. info["contact"] = contact
  581. if license_info:
  582. info["license"] = license_info
  583. output: dict[str, Any] = {"openapi": openapi_version, "info": info}
  584. if servers:
  585. output["servers"] = servers
  586. components: dict[str, dict[str, Any]] = {}
  587. paths: dict[str, dict[str, Any]] = {}
  588. webhook_paths: dict[str, dict[str, Any]] = {}
  589. operation_ids: set[str] = set()
  590. all_fields = get_fields_from_routes(list(routes) + list(webhooks or []))
  591. flat_models = get_flat_models_from_fields(all_fields, known_models=set())
  592. model_name_map = get_model_name_map(flat_models)
  593. field_mapping, definitions = get_definitions(
  594. fields=all_fields,
  595. model_name_map=model_name_map,
  596. separate_input_output_schemas=separate_input_output_schemas,
  597. )
  598. for route_context in routing.iter_route_contexts(routes):
  599. api_route = _get_api_route_for_openapi(route_context)
  600. if api_route is not None:
  601. result = get_openapi_path(
  602. route=api_route,
  603. operation_ids=operation_ids,
  604. model_name_map=model_name_map,
  605. field_mapping=field_mapping,
  606. separate_input_output_schemas=separate_input_output_schemas,
  607. )
  608. if result:
  609. path, security_schemes, path_definitions = result
  610. if path:
  611. paths.setdefault(api_route.path_format, {}).update(path)
  612. if security_schemes:
  613. components.setdefault("securitySchemes", {}).update(
  614. security_schemes
  615. )
  616. if path_definitions:
  617. definitions.update(path_definitions)
  618. for webhook_context in routing.iter_route_contexts(webhooks or []):
  619. api_webhook = _get_api_route_for_openapi(webhook_context)
  620. if api_webhook is not None:
  621. result = get_openapi_path(
  622. route=api_webhook,
  623. operation_ids=operation_ids,
  624. model_name_map=model_name_map,
  625. field_mapping=field_mapping,
  626. separate_input_output_schemas=separate_input_output_schemas,
  627. )
  628. if result:
  629. path, security_schemes, path_definitions = result
  630. if path:
  631. webhook_paths.setdefault(api_webhook.path_format, {}).update(path)
  632. if security_schemes:
  633. components.setdefault("securitySchemes", {}).update(
  634. security_schemes
  635. )
  636. if path_definitions:
  637. definitions.update(path_definitions)
  638. if definitions:
  639. components["schemas"] = {k: definitions[k] for k in sorted(definitions)}
  640. if components:
  641. output["components"] = components
  642. output["paths"] = paths
  643. if webhook_paths:
  644. output["webhooks"] = webhook_paths
  645. if tags:
  646. output["tags"] = tags
  647. if external_docs:
  648. output["externalDocs"] = external_docs
  649. return jsonable_encoder(OpenAPI(**output), by_alias=True, exclude_none=True) # type: ignore[no-any-return]