Browse Source

♻️ Refactor internals to use new compat utils for v1

pull/14168/head
Sebastián Ramírez 10 months ago
parent
commit
7cd3a71a0e
  1. 9
      fastapi/datastructures.py
  2. 51
      fastapi/dependencies/utils.py
  3. 4
      fastapi/encoders.py
  4. 2
      fastapi/openapi/utils.py
  5. 42
      fastapi/utils.py

9
fastapi/datastructures.py

@ -153,11 +153,10 @@ class UploadFile(StarletteUploadFile):
raise ValueError(f"Expected UploadFile, received: {type(__input_value)}") raise ValueError(f"Expected UploadFile, received: {type(__input_value)}")
return cast(UploadFile, __input_value) return cast(UploadFile, __input_value)
if not PYDANTIC_V2: # TODO: remove when deprecating Pydantic v1
@classmethod
@classmethod def __modify_schema__(cls, field_schema: Dict[str, Any]) -> None:
def __modify_schema__(cls, field_schema: Dict[str, Any]) -> None: field_schema.update({"type": "string", "format": "binary"})
field_schema.update({"type": "string", "format": "binary"})
@classmethod @classmethod
def __get_pydantic_json_schema__( def __get_pydantic_json_schema__(

51
fastapi/dependencies/utils.py

@ -48,6 +48,7 @@ from fastapi._compat import (
v1, v1,
value_is_sequence, value_is_sequence,
) )
from fastapi._compat.shared import annotation_is_pydantic_v1
from fastapi.background import BackgroundTasks from fastapi.background import BackgroundTasks
from fastapi.concurrency import ( from fastapi.concurrency import (
asynccontextmanager, asynccontextmanager,
@ -75,6 +76,8 @@ from starlette.responses import Response
from starlette.websockets import WebSocket from starlette.websockets import WebSocket
from typing_extensions import Annotated, get_args, get_origin from typing_extensions import Annotated, get_args, get_origin
from .._compat import _params_v1
if sys.version_info >= (3, 13): # pragma: no cover if sys.version_info >= (3, 13): # pragma: no cover
from inspect import iscoroutinefunction from inspect import iscoroutinefunction
else: # pragma: no cover else: # pragma: no cover
@ -316,7 +319,7 @@ def get_dependant(
) )
continue continue
assert param_details.field is not None assert param_details.field is not None
if isinstance(param_details.field.field_info, params.Body): if isinstance(param_details.field.field_info, (params.Body, _params_v1.Body)):
dependant.body_params.append(param_details.field) dependant.body_params.append(param_details.field)
else: else:
add_param_to_fields(field=param_details.field, dependant=dependant) add_param_to_fields(field=param_details.field, dependant=dependant)
@ -375,21 +378,30 @@ def analyze_param(
fastapi_annotations = [ fastapi_annotations = [
arg arg
for arg in annotated_args[1:] for arg in annotated_args[1:]
if isinstance(arg, (FieldInfo, params.Depends)) if isinstance(arg, (FieldInfo, v1.FieldInfo, params.Depends))
] ]
fastapi_specific_annotations = [ fastapi_specific_annotations = [
arg arg
for arg in fastapi_annotations for arg in fastapi_annotations
if isinstance(arg, (params.Param, params.Body, params.Depends)) if isinstance(
arg,
(
params.Param,
_params_v1.Param,
params.Body,
_params_v1.Body,
params.Depends,
),
)
] ]
if fastapi_specific_annotations: if fastapi_specific_annotations:
fastapi_annotation: Union[FieldInfo, params.Depends, None] = ( fastapi_annotation: Union[FieldInfo, v1.FieldInfo, params.Depends, None] = (
fastapi_specific_annotations[-1] fastapi_specific_annotations[-1]
) )
else: else:
fastapi_annotation = None fastapi_annotation = None
# Set default for Annotated FieldInfo # Set default for Annotated FieldInfo
if isinstance(fastapi_annotation, FieldInfo): if isinstance(fastapi_annotation, (FieldInfo, v1.FieldInfo)):
# Copy `field_info` because we mutate `field_info.default` below. # Copy `field_info` because we mutate `field_info.default` below.
field_info = copy_field_info( field_info = copy_field_info(
field_info=fastapi_annotation, annotation=use_annotation field_info=fastapi_annotation, annotation=use_annotation
@ -420,14 +432,15 @@ def analyze_param(
) )
depends = value depends = value
# Get FieldInfo from default value # Get FieldInfo from default value
elif isinstance(value, FieldInfo): elif isinstance(value, (FieldInfo, v1.FieldInfo)):
assert field_info is None, ( assert field_info is None, (
"Cannot specify FastAPI annotations in `Annotated` and default value" "Cannot specify FastAPI annotations in `Annotated` and default value"
f" together for {param_name!r}" f" together for {param_name!r}"
) )
field_info = value field_info = value
if PYDANTIC_V2: if PYDANTIC_V2:
field_info.annotation = type_annotation if isinstance(field_info, FieldInfo):
field_info.annotation = type_annotation
# Get Depends from type annotation # Get Depends from type annotation
if depends is not None and depends.dependency is None: if depends is not None and depends.dependency is None:
@ -464,7 +477,14 @@ def analyze_param(
) or is_uploadfile_sequence_annotation(type_annotation): ) or is_uploadfile_sequence_annotation(type_annotation):
field_info = params.File(annotation=use_annotation, default=default_value) field_info = params.File(annotation=use_annotation, default=default_value)
elif not field_annotation_is_scalar(annotation=type_annotation): elif not field_annotation_is_scalar(annotation=type_annotation):
field_info = params.Body(annotation=use_annotation, default=default_value) if annotation_is_pydantic_v1(use_annotation):
field_info = _params_v1.Body(
annotation=use_annotation, default=default_value
)
else:
field_info = params.Body(
annotation=use_annotation, default=default_value
)
else: else:
field_info = params.Query(annotation=use_annotation, default=default_value) field_info = params.Query(annotation=use_annotation, default=default_value)
@ -838,7 +858,7 @@ def is_union_of_base_models(field_type: Any) -> bool:
union_args = get_args(field_type) union_args = get_args(field_type)
for arg in union_args: for arg in union_args:
if not lenient_issubclass(arg, BaseModel): if not _is_model_class(arg):
return False return False
return True return True
@ -860,8 +880,8 @@ def _should_embed_body_fields(fields: List[ModelField]) -> bool:
# If it's a Form (or File) field, it has to be a BaseModel (or a union of BaseModels) to be top level # If it's a Form (or File) field, it has to be a BaseModel (or a union of BaseModels) to be top level
# otherwise it has to be embedded, so that the key value pair can be extracted # otherwise it has to be embedded, so that the key value pair can be extracted
if ( if (
isinstance(first_field.field_info, params.Form) isinstance(first_field.field_info, (params.Form, _params_v1.Form))
and not lenient_issubclass(first_field.type_, BaseModel) and not _is_model_class(first_field.type_)
and not is_union_of_base_models(first_field.type_) and not is_union_of_base_models(first_field.type_)
): ):
return True return True
@ -926,7 +946,7 @@ async def request_body_to_args(
if ( if (
single_not_embedded_field single_not_embedded_field
and lenient_issubclass(first_field.type_, BaseModel) and _is_model_class(first_field.type_)
and isinstance(received_body, FormData) and isinstance(received_body, FormData)
): ):
fields_to_extract = get_cached_model_fields(first_field.type_) fields_to_extract = get_cached_model_fields(first_field.type_)
@ -994,12 +1014,15 @@ def get_body_field(
elif any(isinstance(f.field_info, params.Form) for f in flat_dependant.body_params): elif any(isinstance(f.field_info, params.Form) for f in flat_dependant.body_params):
BodyFieldInfo = params.Form BodyFieldInfo = params.Form
else: else:
BodyFieldInfo = params.Body if annotation_is_pydantic_v1(BodyModel):
BodyFieldInfo = _params_v1.Body
else:
BodyFieldInfo = params.Body
body_param_media_types = [ body_param_media_types = [
f.field_info.media_type f.field_info.media_type
for f in flat_dependant.body_params for f in flat_dependant.body_params
if isinstance(f.field_info, params.Body) if isinstance(f.field_info, (params.Body, _params_v1.Body))
] ]
if len(set(body_param_media_types)) == 1: if len(set(body_param_media_types)) == 1:
BodyFieldInfo_kwargs["media_type"] = body_param_media_types[0] BodyFieldInfo_kwargs["media_type"] = body_param_media_types[0]

4
fastapi/encoders.py

@ -220,10 +220,10 @@ def jsonable_encoder(
include = set(include) include = set(include)
if exclude is not None and not isinstance(exclude, (set, dict)): if exclude is not None and not isinstance(exclude, (set, dict)):
exclude = set(exclude) exclude = set(exclude)
if isinstance(obj, BaseModel): if isinstance(obj, (BaseModel, v1.BaseModel)):
# TODO: remove when deprecating Pydantic v1 # TODO: remove when deprecating Pydantic v1
encoders: Dict[Any, Any] = {} encoders: Dict[Any, Any] = {}
if not PYDANTIC_V2: if isinstance(obj, v1.BaseModel):
encoders = getattr(obj.__config__, "json_encoders", {}) # type: ignore[attr-defined] encoders = getattr(obj.__config__, "json_encoders", {}) # type: ignore[attr-defined]
if custom_encoder: if custom_encoder:
encoders = {**encoders, **custom_encoder} encoders = {**encoders, **custom_encoder}

2
fastapi/openapi/utils.py

@ -176,7 +176,7 @@ def get_openapi_operation_request_body(
) -> Optional[Dict[str, Any]]: ) -> Optional[Dict[str, Any]]:
if not body_field: if not body_field:
return None return None
assert isinstance(body_field, ModelField) assert _is_model_field(body_field)
body_schema = get_schema_from_model_field( body_schema = get_schema_from_model_field(
field=body_field, field=body_field,
model_name_map=model_name_map, model_name_map=model_name_map,

42
fastapi/utils.py

@ -82,22 +82,24 @@ def create_model_field(
field_info: Optional[FieldInfo] = None, field_info: Optional[FieldInfo] = None,
alias: Optional[str] = None, alias: Optional[str] = None,
mode: Literal["validation", "serialization"] = "validation", mode: Literal["validation", "serialization"] = "validation",
version: Literal["1", "auto"] = "auto",
) -> ModelField: ) -> ModelField:
class_validators = class_validators or {} class_validators = class_validators or {}
kwargs = {"name": name, "field_info": field_info}
if lenient_issubclass(type_, v1.BaseModel): if lenient_issubclass(type_, v1.BaseModel) or version == "1":
model_config = v1.BaseConfig model_config = v1.BaseConfig
field_info = field_info or v1.FieldInfo() field_info = field_info or v1.FieldInfo()
kwargs.update( kwargs = {
{ "name": name,
"type_": type_, "field_info": field_info,
"class_validators": class_validators, "type_": type_,
"default": default, "class_validators": class_validators,
"required": required, "default": default,
"model_config": model_config, "required": required,
"alias": alias, "model_config": model_config,
} "alias": alias,
) }
try: try:
return v1.ModelField(**kwargs) # type: ignore[arg-type] return v1.ModelField(**kwargs) # type: ignore[arg-type]
except RuntimeError: except RuntimeError:
@ -108,7 +110,7 @@ def create_model_field(
field_info = field_info or FieldInfo( field_info = field_info or FieldInfo(
annotation=type_, default=default, alias=alias annotation=type_, default=default, alias=alias
) )
kwargs.update({"mode": mode}) kwargs = {"mode": mode, "name": name, "field_info": field_info}
try: try:
return v2.ModelField(**kwargs) # type: ignore[arg-type] return v2.ModelField(**kwargs) # type: ignore[arg-type]
except PydanticSchemaGenerationError: except PydanticSchemaGenerationError:
@ -121,7 +123,9 @@ def create_cloned_field(
cloned_types: Optional[MutableMapping[Type[BaseModel], Type[BaseModel]]] = None, cloned_types: Optional[MutableMapping[Type[BaseModel], Type[BaseModel]]] = None,
) -> ModelField: ) -> ModelField:
if PYDANTIC_V2: if PYDANTIC_V2:
return field from ._compat import v2
if isinstance(field, v2.ModelField): # type: ignore[name-defined]
return field
# cloned_types caches already cloned types to support recursive models and improve # cloned_types caches already cloned types to support recursive models and improve
# performance by avoiding unnecessary cloning # performance by avoiding unnecessary cloning
if cloned_types is None: if cloned_types is None:
@ -131,17 +135,17 @@ def create_cloned_field(
if is_dataclass(original_type) and hasattr(original_type, "__pydantic_model__"): if is_dataclass(original_type) and hasattr(original_type, "__pydantic_model__"):
original_type = original_type.__pydantic_model__ original_type = original_type.__pydantic_model__
use_type = original_type use_type = original_type
if lenient_issubclass(original_type, BaseModel): if lenient_issubclass(original_type, v1.BaseModel):
original_type = cast(Type[BaseModel], original_type) original_type = cast(Type[v1.BaseModel], original_type)
use_type = cloned_types.get(original_type) use_type = cloned_types.get(original_type)
if use_type is None: if use_type is None:
use_type = create_model(original_type.__name__, __base__=original_type) use_type = v1.create_model(original_type.__name__, __base__=original_type)
cloned_types[original_type] = use_type cloned_types[original_type] = use_type
for f in original_type.__fields__.values(): for f in original_type.__fields__.values():
use_type.__fields__[f.name] = create_cloned_field( use_type.__fields__[f.name] = create_cloned_field(
f, cloned_types=cloned_types f, cloned_types=cloned_types,
) )
new_field = create_model_field(name=field.name, type_=use_type) new_field = create_model_field(name=field.name, type_=use_type, version="1")
new_field.has_alias = field.has_alias # type: ignore[attr-defined] new_field.has_alias = field.has_alias # type: ignore[attr-defined]
new_field.alias = field.alias # type: ignore[misc] new_field.alias = field.alias # type: ignore[misc]
new_field.class_validators = field.class_validators # type: ignore[attr-defined] new_field.class_validators = field.class_validators # type: ignore[attr-defined]

Loading…
Cancel
Save