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)}")
return cast(UploadFile, __input_value)
if not PYDANTIC_V2:
@classmethod
def __modify_schema__(cls, field_schema: Dict[str, Any]) -> None:
field_schema.update({"type": "string", "format": "binary"})
# TODO: remove when deprecating Pydantic v1
@classmethod
def __modify_schema__(cls, field_schema: Dict[str, Any]) -> None:
field_schema.update({"type": "string", "format": "binary"})
@classmethod
def __get_pydantic_json_schema__(

51
fastapi/dependencies/utils.py

@ -48,6 +48,7 @@ from fastapi._compat import (
v1,
value_is_sequence,
)
from fastapi._compat.shared import annotation_is_pydantic_v1
from fastapi.background import BackgroundTasks
from fastapi.concurrency import (
asynccontextmanager,
@ -75,6 +76,8 @@ from starlette.responses import Response
from starlette.websockets import WebSocket
from typing_extensions import Annotated, get_args, get_origin
from .._compat import _params_v1
if sys.version_info >= (3, 13): # pragma: no cover
from inspect import iscoroutinefunction
else: # pragma: no cover
@ -316,7 +319,7 @@ def get_dependant(
)
continue
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)
else:
add_param_to_fields(field=param_details.field, dependant=dependant)
@ -375,21 +378,30 @@ def analyze_param(
fastapi_annotations = [
arg
for arg in annotated_args[1:]
if isinstance(arg, (FieldInfo, params.Depends))
if isinstance(arg, (FieldInfo, v1.FieldInfo, params.Depends))
]
fastapi_specific_annotations = [
arg
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:
fastapi_annotation: Union[FieldInfo, params.Depends, None] = (
fastapi_annotation: Union[FieldInfo, v1.FieldInfo, params.Depends, None] = (
fastapi_specific_annotations[-1]
)
else:
fastapi_annotation = None
# 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.
field_info = copy_field_info(
field_info=fastapi_annotation, annotation=use_annotation
@ -420,14 +432,15 @@ def analyze_param(
)
depends = value
# Get FieldInfo from default value
elif isinstance(value, FieldInfo):
elif isinstance(value, (FieldInfo, v1.FieldInfo)):
assert field_info is None, (
"Cannot specify FastAPI annotations in `Annotated` and default value"
f" together for {param_name!r}"
)
field_info = value
if PYDANTIC_V2:
field_info.annotation = type_annotation
if isinstance(field_info, FieldInfo):
field_info.annotation = type_annotation
# Get Depends from type annotation
if depends is not None and depends.dependency is None:
@ -464,7 +477,14 @@ def analyze_param(
) or is_uploadfile_sequence_annotation(type_annotation):
field_info = params.File(annotation=use_annotation, default=default_value)
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:
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)
for arg in union_args:
if not lenient_issubclass(arg, BaseModel):
if not _is_model_class(arg):
return False
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
# otherwise it has to be embedded, so that the key value pair can be extracted
if (
isinstance(first_field.field_info, params.Form)
and not lenient_issubclass(first_field.type_, BaseModel)
isinstance(first_field.field_info, (params.Form, _params_v1.Form))
and not _is_model_class(first_field.type_)
and not is_union_of_base_models(first_field.type_)
):
return True
@ -926,7 +946,7 @@ async def request_body_to_args(
if (
single_not_embedded_field
and lenient_issubclass(first_field.type_, BaseModel)
and _is_model_class(first_field.type_)
and isinstance(received_body, FormData)
):
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):
BodyFieldInfo = params.Form
else:
BodyFieldInfo = params.Body
if annotation_is_pydantic_v1(BodyModel):
BodyFieldInfo = _params_v1.Body
else:
BodyFieldInfo = params.Body
body_param_media_types = [
f.field_info.media_type
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:
BodyFieldInfo_kwargs["media_type"] = body_param_media_types[0]

4
fastapi/encoders.py

@ -220,10 +220,10 @@ def jsonable_encoder(
include = set(include)
if exclude is not None and not isinstance(exclude, (set, dict)):
exclude = set(exclude)
if isinstance(obj, BaseModel):
if isinstance(obj, (BaseModel, v1.BaseModel)):
# TODO: remove when deprecating Pydantic v1
encoders: Dict[Any, Any] = {}
if not PYDANTIC_V2:
if isinstance(obj, v1.BaseModel):
encoders = getattr(obj.__config__, "json_encoders", {}) # type: ignore[attr-defined]
if 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]]:
if not body_field:
return None
assert isinstance(body_field, ModelField)
assert _is_model_field(body_field)
body_schema = get_schema_from_model_field(
field=body_field,
model_name_map=model_name_map,

42
fastapi/utils.py

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

Loading…
Cancel
Save