Browse Source

fix functions

pull/14186/head
svlandeg 10 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 functools import lru_cache
from typing import ( from typing import (
Any, Any,
@ -275,14 +276,29 @@ def get_definitions(
], ],
Dict[str, Dict[str, Any]], Dict[str, Dict[str, Any]],
]: ]:
v1_fields = [field for field in fields if isinstance(field, may_v1.ModelField)] if sys.version_info < (3, 14):
v1_field_maps, v1_definitions = may_v1.get_definitions( v1_fields = [field for field in fields if isinstance(field, may_v1.ModelField)]
fields=v1_fields, v1_field_maps, v1_definitions = may_v1.get_definitions(
model_name_map=model_name_map, fields=v1_fields,
separate_input_output_schemas=separate_input_output_schemas, 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 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: else:
from . import v2 from . import v2
@ -292,9 +308,7 @@ def get_definitions(
model_name_map=model_name_map, model_name_map=model_name_map,
separate_input_output_schemas=separate_input_output_schemas, separate_input_output_schemas=separate_input_output_schemas,
) )
all_definitions = {**v1_definitions, **v2_definitions} return v2_field_maps, v2_definitions
all_field_maps = {**v1_field_maps, **v2_field_maps}
return all_field_maps, all_definitions
def get_schema_from_model_field( def get_schema_from_model_field(

44
fastapi/_compat/may_v1.py

@ -1,5 +1,5 @@
import sys 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 from fastapi.types import ModelNameMap
@ -56,6 +56,8 @@ if sys.version_info >= (3, 14):
class Url: class Url:
pass pass
from .v2 import create_model
def get_definitions( def get_definitions(
*, *,
fields: List[ModelField], fields: List[ModelField],
@ -69,14 +71,6 @@ if sys.version_info >= (3, 14):
]: ]:
return {}, {} 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: else:
from .v1 import AnyUrl as AnyUrl from .v1 import AnyUrl as AnyUrl
@ -96,6 +90,34 @@ else:
from .v1 import Undefined as Undefined from .v1 import Undefined as Undefined
from .v1 import UndefinedType as UndefinedType from .v1 import UndefinedType as UndefinedType
from .v1 import Url as Url from .v1 import Url as Url
from .v1 import _normalize_errors as _normalize_errors from .v1 import create_model
from .v1 import _regenerate_error_with_loc as _regenerate_error_with_loc
from .v1 import get_definitions as get_definitions 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 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: def _model_rebuild(model: Type[BaseModel]) -> None:
model.update_forward_refs() model.update_forward_refs()

Loading…
Cancel
Save