Browse Source

fix functions

pull/14186/head
svlandeg 9 months ago
parent
commit
bdbd90c39d
  1. 36
      fastapi/_compat/main.py
  2. 44
      fastapi/_compat/may_v1.py
  3. 26
      fastapi/_compat/v1.py

36
fastapi/_compat/main.py

@ -1,3 +1,4 @@
import sys
from functools import lru_cache
from typing import (
Any,
@ -275,14 +276,29 @@ def get_definitions(
],
Dict[str, Dict[str, Any]],
]:
v1_fields = [field for field in fields if isinstance(field, may_v1.ModelField)]
v1_field_maps, v1_definitions = may_v1.get_definitions(
fields=v1_fields,
model_name_map=model_name_map,
separate_input_output_schemas=separate_input_output_schemas,
)
if not PYDANTIC_V2:
return v1_field_maps, v1_definitions
if sys.version_info < (3, 14):
v1_fields = [field for field in fields if isinstance(field, may_v1.ModelField)]
v1_field_maps, v1_definitions = may_v1.get_definitions(
fields=v1_fields,
model_name_map=model_name_map,
separate_input_output_schemas=separate_input_output_schemas,
)
if not PYDANTIC_V2:
return v1_field_maps, v1_definitions
else:
from . import v2
v2_fields = [field for field in fields if isinstance(field, v2.ModelField)]
v2_field_maps, v2_definitions = v2.get_definitions(
fields=v2_fields,
model_name_map=model_name_map,
separate_input_output_schemas=separate_input_output_schemas,
)
all_definitions = {**v1_definitions, **v2_definitions}
all_field_maps = {**v1_field_maps, **v2_field_maps}
return all_field_maps, all_definitions
# Pydantic v1 is not supported since Python 3.14
else:
from . import v2
@ -292,9 +308,7 @@ def get_definitions(
model_name_map=model_name_map,
separate_input_output_schemas=separate_input_output_schemas,
)
all_definitions = {**v1_definitions, **v2_definitions}
all_field_maps = {**v1_field_maps, **v2_field_maps}
return all_field_maps, all_definitions
return v2_field_maps, v2_definitions
def get_schema_from_model_field(

44
fastapi/_compat/may_v1.py

@ -1,5 +1,5 @@
import sys
from typing import Any, Dict, List, Literal, Sequence, Tuple, Union
from typing import Any, Dict, List, Literal, Sequence, Tuple, Union, Type
from fastapi.types import ModelNameMap
@ -56,6 +56,8 @@ if sys.version_info >= (3, 14):
class Url:
pass
from .v2 import create_model
def get_definitions(
*,
fields: List[ModelField],
@ -69,14 +71,6 @@ if sys.version_info >= (3, 14):
]:
return {}, {}
def _normalize_errors(errors: Sequence[Any]) -> List[Dict[str, Any]]:
return []
def _regenerate_error_with_loc(
*, errors: Sequence[Any], loc_prefix: Tuple[Union[str, int], ...]
) -> List[Dict[str, Any]]:
return []
else:
from .v1 import AnyUrl as AnyUrl
@ -96,6 +90,34 @@ else:
from .v1 import Undefined as Undefined
from .v1 import UndefinedType as UndefinedType
from .v1 import Url as Url
from .v1 import _normalize_errors as _normalize_errors
from .v1 import _regenerate_error_with_loc as _regenerate_error_with_loc
from .v1 import create_model
from .v1 import get_definitions as get_definitions
RequestErrorModel: Type[BaseModel] = create_model("Request")
def _normalize_errors(errors: Sequence[Any]) -> List[Dict[str, Any]]:
use_errors: List[Any] = []
for error in errors:
if isinstance(error, ErrorWrapper):
new_errors = ValidationError( # type: ignore[call-arg]
errors=[error], model=RequestErrorModel
).errors()
use_errors.extend(new_errors)
elif isinstance(error, list):
use_errors.extend(_normalize_errors(error))
else:
use_errors.append(error)
return use_errors
def _regenerate_error_with_loc(
*, errors: Sequence[Any], loc_prefix: Tuple[Union[str, int], ...]
) -> List[Dict[str, Any]]:
updated_loc_errors: List[Any] = [
{**err, "loc": loc_prefix + err.get("loc", ())}
for err in _normalize_errors(errors)
]
return updated_loc_errors

26
fastapi/_compat/v1.py

@ -219,32 +219,6 @@ def is_pv1_scalar_sequence_field(field: ModelField) -> bool:
return False
def _normalize_errors(errors: Sequence[Any]) -> List[Dict[str, Any]]:
use_errors: List[Any] = []
for error in errors:
if isinstance(error, ErrorWrapper):
new_errors = ValidationError( # type: ignore[call-arg]
errors=[error], model=RequestErrorModel
).errors()
use_errors.extend(new_errors)
elif isinstance(error, list):
use_errors.extend(_normalize_errors(error))
else:
use_errors.append(error)
return use_errors
def _regenerate_error_with_loc(
*, errors: Sequence[Any], loc_prefix: Tuple[Union[str, int], ...]
) -> List[Dict[str, Any]]:
updated_loc_errors: List[Any] = [
{**err, "loc": loc_prefix + err.get("loc", ())}
for err in _normalize_errors(errors)
]
return updated_loc_errors
def _model_rebuild(model: Type[BaseModel]) -> None:
model.update_forward_refs()

Loading…
Cancel
Save