Browse Source

♻️ Refactor Pydantic v1 FastAPI params

pull/14168/head
Sebastián Ramírez 9 months ago
parent
commit
367d804a62
  1. 48
      fastapi/dependencies/utils.py
  2. 5
      fastapi/routing.py
  3. 4
      fastapi/temp_pydantic_v1_params.py

48
fastapi/dependencies/utils.py

@ -76,7 +76,7 @@ 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 from .. import temp_pydantic_v1_params
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
@ -319,7 +319,9 @@ 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, _params_v1.Body)): if isinstance(
param_details.field.field_info, (params.Body, temp_pydantic_v1_params.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)
@ -387,9 +389,9 @@ def analyze_param(
arg, arg,
( (
params.Param, params.Param,
_params_v1.Param, temp_pydantic_v1_params.Param,
params.Body, params.Body,
_params_v1.Body, temp_pydantic_v1_params.Body,
params.Depends, params.Depends,
), ),
) )
@ -479,7 +481,7 @@ def analyze_param(
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):
if annotation_is_pydantic_v1(use_annotation): if annotation_is_pydantic_v1(use_annotation):
field_info = _params_v1.Body( field_info = temp_pydantic_v1_params.Body(
annotation=use_annotation, default=default_value annotation=use_annotation, default=default_value
) )
else: else:
@ -494,12 +496,14 @@ def analyze_param(
if field_info is not None: if field_info is not None:
# Handle field_info.in_ # Handle field_info.in_
if is_path_param: if is_path_param:
assert isinstance(field_info, (params.Path, _params_v1.Path)), ( assert isinstance(
field_info, (params.Path, temp_pydantic_v1_params.Path)
), (
f"Cannot use `{field_info.__class__.__name__}` for path param" f"Cannot use `{field_info.__class__.__name__}` for path param"
f" {param_name!r}" f" {param_name!r}"
) )
elif ( elif (
isinstance(field_info, (params.Param, _params_v1.Param)) isinstance(field_info, (params.Param, temp_pydantic_v1_params.Param))
and getattr(field_info, "in_", None) is None and getattr(field_info, "in_", None) is None
): ):
field_info.in_ = params.ParamTypes.query field_info.in_ = params.ParamTypes.query
@ -508,7 +512,7 @@ def analyze_param(
field_info, field_info,
param_name, param_name,
) )
if isinstance(field_info, (params.Form, _params_v1.Form)): if isinstance(field_info, (params.Form, temp_pydantic_v1_params.Form)):
ensure_multipart_is_installed() ensure_multipart_is_installed()
if not field_info.alias and getattr(field_info, "convert_underscores", None): if not field_info.alias and getattr(field_info, "convert_underscores", None):
alias = param_name.replace("_", "-") alias = param_name.replace("_", "-")
@ -527,7 +531,7 @@ def analyze_param(
assert is_scalar_field(field=field), ( assert is_scalar_field(field=field), (
"Path params must be of one of the supported types" "Path params must be of one of the supported types"
) )
elif isinstance(field_info, (params.Query, _params_v1.Query)): elif isinstance(field_info, (params.Query, temp_pydantic_v1_params.Query)):
assert ( assert (
is_scalar_field(field) is_scalar_field(field)
or is_scalar_sequence_field(field) or is_scalar_sequence_field(field)
@ -754,7 +758,7 @@ def _get_multidict_value(
if ( if (
value is None value is None
or ( or (
isinstance(field.field_info, (params.Form, _params_v1.Form)) isinstance(field.field_info, (params.Form, temp_pydantic_v1_params.Form))
and isinstance(value, str) # For type checks and isinstance(value, str) # For type checks
and value == "" and value == ""
) )
@ -820,7 +824,7 @@ def request_params_to_args(
if single_not_embedded_field: if single_not_embedded_field:
field_info = first_field.field_info field_info = first_field.field_info
assert isinstance(field_info, (params.Param, _params_v1.Param)), ( assert isinstance(field_info, (params.Param, temp_pydantic_v1_params.Param)), (
"Params must be subclasses of Param" "Params must be subclasses of Param"
) )
loc: Tuple[str, ...] = (field_info.in_.value,) loc: Tuple[str, ...] = (field_info.in_.value,)
@ -832,7 +836,7 @@ def request_params_to_args(
for field in fields: for field in fields:
value = _get_multidict_value(field, received_params) value = _get_multidict_value(field, received_params)
field_info = field.field_info field_info = field.field_info
assert isinstance(field_info, (params.Param, _params_v1.Param)), ( assert isinstance(field_info, (params.Param, temp_pydantic_v1_params.Param)), (
"Params must be subclasses of Param" "Params must be subclasses of Param"
) )
loc = (field_info.in_.value, field.alias) loc = (field_info.in_.value, field.alias)
@ -881,7 +885,7 @@ 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, _params_v1.Form)) isinstance(first_field.field_info, (params.Form, temp_pydantic_v1_params.Form))
and not _is_model_class(first_field.type_) 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_)
): ):
@ -899,14 +903,14 @@ async def _extract_form_body(
value = _get_multidict_value(field, received_body) value = _get_multidict_value(field, received_body)
field_info = field.field_info field_info = field.field_info
if ( if (
isinstance(field_info, (params.File, _params_v1.File)) isinstance(field_info, (params.File, temp_pydantic_v1_params.File))
and is_bytes_field(field) and is_bytes_field(field)
and isinstance(value, UploadFile) and isinstance(value, UploadFile)
): ):
value = await value.read() value = await value.read()
elif ( elif (
is_bytes_sequence_field(field) is_bytes_sequence_field(field)
and isinstance(field_info, (params.File, _params_v1.File)) and isinstance(field_info, (params.File, temp_pydantic_v1_params.File))
and value_is_sequence(value) and value_is_sequence(value)
): ):
# For types # For types
@ -1013,25 +1017,27 @@ def get_body_field(
if any(isinstance(f.field_info, params.File) for f in flat_dependant.body_params): if any(isinstance(f.field_info, params.File) for f in flat_dependant.body_params):
BodyFieldInfo: Type[params.Body] = params.File BodyFieldInfo: Type[params.Body] = params.File
elif any( elif any(
isinstance(f.field_info, _params_v1.File) for f in flat_dependant.body_params isinstance(f.field_info, temp_pydantic_v1_params.File)
for f in flat_dependant.body_params
): ):
BodyFieldInfo: Type[_params_v1.Body] = _params_v1.File # type: ignore[no-redef] BodyFieldInfo: Type[temp_pydantic_v1_params.Body] = temp_pydantic_v1_params.File # type: ignore[no-redef]
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
elif any( elif any(
isinstance(f.field_info, _params_v1.Form) for f in flat_dependant.body_params isinstance(f.field_info, temp_pydantic_v1_params.Form)
for f in flat_dependant.body_params
): ):
BodyFieldInfo = _params_v1.Form # type: ignore[assignment] BodyFieldInfo = temp_pydantic_v1_params.Form # type: ignore[assignment]
else: else:
if annotation_is_pydantic_v1(BodyModel): if annotation_is_pydantic_v1(BodyModel):
BodyFieldInfo = _params_v1.Body # type: ignore[assignment] BodyFieldInfo = temp_pydantic_v1_params.Body # type: ignore[assignment]
else: else:
BodyFieldInfo = params.Body 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, _params_v1.Body)) if isinstance(f.field_info, (params.Body, temp_pydantic_v1_params.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]

5
fastapi/routing.py

@ -24,14 +24,13 @@ from typing import (
Union, Union,
) )
from fastapi import params from fastapi import params, temp_pydantic_v1_params
from fastapi._compat import ( from fastapi._compat import (
ModelField, ModelField,
Undefined, Undefined,
_get_model_config, _get_model_config,
_model_dump, _model_dump,
_normalize_errors, _normalize_errors,
_params_v1,
lenient_issubclass, lenient_issubclass,
) )
from fastapi.datastructures import Default, DefaultPlaceholder from fastapi.datastructures import Default, DefaultPlaceholder
@ -309,7 +308,7 @@ def get_request_handler(
assert dependant.call is not None, "dependant.call must be a function" assert dependant.call is not None, "dependant.call must be a function"
is_coroutine = iscoroutinefunction(dependant.call) is_coroutine = iscoroutinefunction(dependant.call)
is_body_form = body_field and isinstance( is_body_form = body_field and isinstance(
body_field.field_info, (params.Form, _params_v1.Form) body_field.field_info, (params.Form, temp_pydantic_v1_params.Form)
) )
if isinstance(response_class, DefaultPlaceholder): if isinstance(response_class, DefaultPlaceholder):
actual_response_class: Type[Response] = response_class.value actual_response_class: Type[Response] = response_class.value

4
fastapi/_compat/_params_v1.py → fastapi/temp_pydantic_v1_params.py

@ -5,8 +5,8 @@ from fastapi.openapi.models import Example
from fastapi.params import ParamTypes from fastapi.params import ParamTypes
from typing_extensions import Annotated, deprecated from typing_extensions import Annotated, deprecated
from .shared import PYDANTIC_VERSION_MINOR_TUPLE from ._compat.shared import PYDANTIC_VERSION_MINOR_TUPLE
from .v1 import FieldInfo, Undefined from ._compat.v1 import FieldInfo, Undefined
_Unset: Any = Undefined _Unset: Any = Undefined
Loading…
Cancel
Save