From 658c2603cea93c95f1c8813e650c7e4f5c38ae63 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Sebasti=C3=A1n=20Ram=C3=ADrez?= Date: Sun, 5 Oct 2025 23:36:20 +0200 Subject: [PATCH] =?UTF-8?q?=E2=99=BB=EF=B8=8F=20Add=20Pydantic=20v1=20sche?= =?UTF-8?q?ma=20functions=20to=20v2=20=5Fcompat=20to=20allow=20generating?= =?UTF-8?q?=20JSON=20Schema=20for=20all=20together?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- fastapi/_compat/main.py | 23 +++++++-- fastapi/_compat/v2.py | 100 ++++++++++++++++++++++++++++++++++++++-- 2 files changed, 115 insertions(+), 8 deletions(-) diff --git a/fastapi/_compat/main.py b/fastapi/_compat/main.py index 86661191a1..21043e72b2 100644 --- a/fastapi/_compat/main.py +++ b/fastapi/_compat/main.py @@ -204,11 +204,24 @@ def _model_rebuild(model: Type[BaseModel]) -> None: def get_compat_model_name_map(fields: List[ModelField]) -> ModelNameMap: - v1_models = [field for field in fields if isinstance(field, v1.ModelField)] - if v1_models: - models = v1.get_flat_models_from_fields(v1_models, known_models=set()) - return v1.get_model_name_map(models) # type: ignore[no-any-return] - return {} + v1_model_fields = [field for field in fields if isinstance(field, v1.ModelField)] + v1_flat_models = v1.get_flat_models_from_fields(v1_model_fields, known_models=set()) + all_flat_models = v1_flat_models + if PYDANTIC_V2: + from . import v2 + + v2_model_fields = [ + field for field in fields if isinstance(field, v2.ModelField) + ] + v2_flat_models = v2.get_flat_models_from_fields( + v2_model_fields, known_models=set() + ) + all_flat_models = all_flat_models.union(v2_flat_models) + + model_name_map = v2.get_model_name_map(all_flat_models) # type: ignore[no-any-return] + return model_name_map + model_name_map = v1.get_model_name_map(all_flat_models) # type: ignore[no-any-return] + return model_name_map # type: ignore[return-value] def get_definitions( diff --git a/fastapi/_compat/v2.py b/fastapi/_compat/v2.py index 943aa0eb65..b000679a58 100644 --- a/fastapi/_compat/v2.py +++ b/fastapi/_compat/v2.py @@ -1,10 +1,13 @@ +import re from copy import copy, deepcopy from dataclasses import dataclass +from enum import Enum from typing import ( Any, Dict, List, Sequence, + Set, Tuple, Type, Union, @@ -29,7 +32,7 @@ from pydantic.json_schema import JsonSchemaValue as JsonSchemaValue from pydantic_core import CoreSchema as CoreSchema from pydantic_core import PydanticUndefined, PydanticUndefinedType from pydantic_core import Url as Url -from typing_extensions import Annotated, Literal, get_origin +from typing_extensions import Annotated, Literal, get_args, get_origin try: from pydantic_core.core_schema import ( @@ -192,9 +195,19 @@ def get_definitions( override_mode: Union[Literal["validation"], None] = ( None if separate_input_output_schemas else "validation" ) + flat_models = get_flat_models_from_fields(fields, known_models=set()) + flat_model_fields = [ + ModelField(field_info=FieldInfo(annotation=model), name=model.__name__) + for model in flat_models + ] + input_types = {f.type_ for f in fields} + unique_flat_model_fields = { + f for f in flat_model_fields if f.type_ not in input_types + } + inputs = [ (field, override_mode or field.mode, field._type_adapter.core_schema) - for field in fields + for field in fields + list(unique_flat_model_fields) ] field_mapping, definitions = schema_generator.generate_definitions(inputs=inputs) for item_def in cast(Dict[str, Dict[str, Any]], definitions).values(): @@ -220,7 +233,7 @@ def _replace_refs( ref_name = schema["$ref"].split("/")[-1] if ref_name in old_name_to_new_name_map: new_name = old_name_to_new_name_map[ref_name] - new_schema["$ref"] = REF_TEMPLATE.format(new_name) + new_schema["$ref"] = REF_TEMPLATE.format(model=new_name) else: new_schema["$ref"] = schema["$ref"] continue @@ -260,6 +273,8 @@ def _remap_definitions_and_field_mappings( continue new_name = model_name_map[model] old_name = value["$ref"].split("/")[-1] + if old_name in {f"{new_name}-Input", f"{new_name}-Output"}: + continue old_name_to_new_name_map[old_name] = new_name new_field_mapping = {} @@ -344,3 +359,82 @@ def get_model_fields(model: Type[BaseModel]) -> List[ModelField]: ModelField(field_info=field_info, name=name) for name, field_info in model.model_fields.items() ] + +# Duplicate of several schema functions from Pydantic v1 to make them compatible with +# Pydantic v2 and allow mixing the models + +TypeModelOrEnum = Union[Type["BaseModel"], Type[Enum]] +TypeModelSet = Set[TypeModelOrEnum] + + +def normalize_name(name: str) -> str: + return re.sub(r"[^a-zA-Z0-9.\-_]", "_", name) + + +def get_model_name_map(unique_models: TypeModelSet) -> Dict[TypeModelOrEnum, str]: + name_model_map = {} + conflicting_names: Set[str] = set() + for model in unique_models: + model_name = normalize_name(model.__name__) + if model_name in conflicting_names: + model_name = get_long_model_name(model) + name_model_map[model_name] = model + elif model_name in name_model_map: + conflicting_names.add(model_name) + conflicting_model = name_model_map.pop(model_name) + name_model_map[get_long_model_name(conflicting_model)] = conflicting_model + name_model_map[get_long_model_name(model)] = model + else: + name_model_map[model_name] = model + return {v: k for k, v in name_model_map.items()} + + +def get_flat_models_from_model( + model: Type["BaseModel"], known_models: Union[TypeModelSet, None] = None +) -> TypeModelSet: + known_models = known_models or set() + fields = get_model_fields(model) + get_flat_models_from_fields(fields, known_models=known_models) + return known_models + + +def get_flat_models_from_annotation( + annotation: Any, known_models: Union[TypeModelSet, None] = None +) -> TypeModelSet: + known_models = known_models or set() + origin = get_origin(annotation) + if origin is not None: + for arg in get_args(annotation): + if lenient_issubclass(arg, (BaseModel, Enum)) and arg not in known_models: + known_models.add(arg) + if lenient_issubclass(arg, BaseModel): + get_flat_models_from_model(arg, known_models=known_models) + else: + get_flat_models_from_annotation(arg, known_models=known_models) + return known_models + + +def get_flat_models_from_field( + field: ModelField, known_models: TypeModelSet +) -> TypeModelSet: + field_type = field.type_ + if lenient_issubclass(field_type, BaseModel): + if field_type in known_models: + return known_models + known_models.add(field_type) + get_flat_models_from_model(field_type, known_models=known_models) + elif lenient_issubclass(field_type, Enum): + known_models.add(field_type) + return known_models + + +def get_flat_models_from_fields( + fields: Sequence[ModelField], known_models: TypeModelSet +) -> TypeModelSet: + for field in fields: + get_flat_models_from_field(field, known_models=known_models) + return known_models + + +def get_long_model_name(model: TypeModelOrEnum) -> str: + return f"{model.__module__}__{model.__qualname__}".replace(".", "__")