diff --git a/fastapi/_compat/main.py b/fastapi/_compat/main.py index 3f758f072..777f93d61 100644 --- a/fastapi/_compat/main.py +++ b/fastapi/_compat/main.py @@ -3,6 +3,7 @@ from typing import ( Any, Dict, List, + Optional, Sequence, Tuple, Type, @@ -19,6 +20,7 @@ from .model_field import ModelField if PYDANTIC_V2: from .v2 import BaseConfig as BaseConfig from .v2 import FieldInfo as FieldInfo + from .v2 import GenerateJsonSchema from .v2 import PydanticSchemaGenerationError as PydanticSchemaGenerationError from .v2 import RequiredParam as RequiredParam from .v2 import Undefined as Undefined @@ -33,6 +35,7 @@ if PYDANTIC_V2: else: from .v1 import BaseConfig as BaseConfig # type: ignore[assignment] from .v1 import FieldInfo as FieldInfo + from .v1 import GenerateJsonSchema from .v1 import ( # type: ignore[assignment] PydanticSchemaGenerationError as PydanticSchemaGenerationError, ) @@ -231,6 +234,7 @@ def get_definitions( fields: List[ModelField], model_name_map: ModelNameMap, separate_input_output_schemas: bool = True, + schema_generator: Optional[GenerateJsonSchema] = None, ) -> Tuple[ Dict[Tuple[ModelField, Literal["validation", "serialization"]], v1.JsonSchemaValue], Dict[str, Dict[str, Any]], @@ -251,6 +255,7 @@ def get_definitions( fields=v2_fields, model_name_map=model_name_map, separate_input_output_schemas=separate_input_output_schemas, + schema_generator=schema_generator, ) all_definitions = {**v1_definitions, **v2_definitions} all_field_maps = {**v1_field_maps, **v2_field_maps} diff --git a/fastapi/_compat/v2.py b/fastapi/_compat/v2.py index 29606b9f3..aea7037b3 100644 --- a/fastapi/_compat/v2.py +++ b/fastapi/_compat/v2.py @@ -7,6 +7,7 @@ from typing import ( Any, Dict, List, + Optional, Sequence, Set, Tuple, @@ -199,11 +200,12 @@ def get_definitions( fields: Sequence[ModelField], model_name_map: ModelNameMap, separate_input_output_schemas: bool = True, + schema_generator: Optional[GenerateJsonSchema] = None, ) -> Tuple[ Dict[Tuple[ModelField, Literal["validation", "serialization"]], JsonSchemaValue], Dict[str, Dict[str, Any]], ]: - schema_generator = GenerateJsonSchema(ref_template=REF_TEMPLATE) + schema_generator = schema_generator or GenerateJsonSchema(ref_template=REF_TEMPLATE) override_mode: Union[Literal["validation"], None] = ( None if separate_input_output_schemas else "validation" ) diff --git a/fastapi/openapi/utils.py b/fastapi/openapi/utils.py index dbc93d289..d006c8586 100644 --- a/fastapi/openapi/utils.py +++ b/fastapi/openapi/utils.py @@ -5,6 +5,7 @@ from typing import Any, Dict, List, Optional, Sequence, Set, Tuple, Type, Union, from fastapi import routing from fastapi._compat import ( + PYDANTIC_V2, JsonSchemaValue, ModelField, Undefined, @@ -38,6 +39,11 @@ from typing_extensions import Literal from .._compat import _is_model_field +if PYDANTIC_V2: + from .._compat.v2 import GenerateJsonSchema +else: + from .._compat.v1 import GenerateJsonSchema + validation_error_definition = { "title": "ValidationError", "type": "object", @@ -480,6 +486,7 @@ def get_openapi( license_info: Optional[Dict[str, Union[str, Any]]] = None, separate_input_output_schemas: bool = True, external_docs: Optional[Dict[str, Any]] = None, + schema_generator: Optional[GenerateJsonSchema] = None, ) -> Dict[str, Any]: info: Dict[str, Any] = {"title": title, "version": version} if summary: @@ -505,6 +512,7 @@ def get_openapi( fields=all_fields, model_name_map=model_name_map, separate_input_output_schemas=separate_input_output_schemas, + schema_generator=schema_generator, ) for route in routes or []: if isinstance(route, routing.APIRoute):