Browse Source

♻️ Tweak implementation and tests for coverage

pull/14168/head
Sebastián Ramírez 10 months ago
parent
commit
aa584829a0
  1. 46
      fastapi/_compat/main.py
  2. 56
      fastapi/_compat/model_field.py
  3. 41
      tests/test_compat.py
  4. 21
      tests/test_pydantic_v1_v2_list.py
  5. 18
      tests/test_response_model_as_return_annotation.py

46
fastapi/_compat/main.py

@ -99,11 +99,11 @@ def _is_error_wrapper(exc: Exception) -> bool:
def copy_field_info(*, field_info: FieldInfo, annotation: Any) -> FieldInfo:
if isinstance(field_info, v1.FieldInfo):
return v1.copy_field_info(field_info=field_info, annotation=annotation)
elif PYDANTIC_V2:
else:
assert PYDANTIC_V2
from . import v2
return v2.copy_field_info(field_info=field_info, annotation=annotation)
raise TypeError("field_info must be an instance of FieldInfo")
def create_body_model(
@ -111,11 +111,11 @@ def create_body_model(
) -> Type[BaseModel]:
if fields and isinstance(fields[0], v1.ModelField):
return v1.create_body_model(fields=fields, model_name=model_name) # type: ignore[return-value]
elif PYDANTIC_V2:
else:
assert PYDANTIC_V2
from . import v2
return v2.create_body_model(fields=fields, model_name=model_name) # type: ignore[return-value]
raise TypeError("fields must be a sequence of ModelField instances")
def get_annotation_from_field_info(
@ -125,73 +125,73 @@ def get_annotation_from_field_info(
return v1.get_annotation_from_field_info(
annotation=annotation, field_info=field_info, field_name=field_name
)
elif PYDANTIC_V2:
else:
assert PYDANTIC_V2
from . import v2
return v2.get_annotation_from_field_info(
annotation=annotation, field_info=field_info, field_name=field_name
)
raise TypeError("field_info must be an instance of FieldInfo")
def is_bytes_field(field: ModelField) -> bool:
if isinstance(field, v1.ModelField):
return v1.is_bytes_field(field)
elif PYDANTIC_V2:
else:
assert PYDANTIC_V2
from . import v2
return v2.is_bytes_field(field) # type: ignore[return-value]
raise TypeError("field_info must be an instance of FieldInfo")
def is_bytes_sequence_field(field: ModelField) -> bool:
if isinstance(field, v1.ModelField):
return v1.is_bytes_sequence_field(field)
elif PYDANTIC_V2:
else:
assert PYDANTIC_V2
from . import v2
return v2.is_bytes_sequence_field(field) # type: ignore[return-value]
raise TypeError("field_info must be an instance of FieldInfo")
def is_scalar_field(field: ModelField) -> bool:
if isinstance(field, v1.ModelField):
return v1.is_scalar_field(field)
elif PYDANTIC_V2:
else:
assert PYDANTIC_V2
from . import v2
return v2.is_scalar_field(field) # type: ignore[return-value]
raise TypeError("field_info must be an instance of FieldInfo")
def is_scalar_sequence_field(field: ModelField) -> bool:
if isinstance(field, v1.ModelField):
return v1.is_scalar_sequence_field(field)
elif PYDANTIC_V2:
else:
assert PYDANTIC_V2
from . import v2
return v2.is_scalar_sequence_field(field) # type: ignore[return-value]
raise TypeError("field_info must be an instance of FieldInfo")
def is_sequence_field(field: ModelField) -> bool:
if isinstance(field, v1.ModelField):
return v1.is_sequence_field(field)
elif PYDANTIC_V2:
else:
assert PYDANTIC_V2
from . import v2
return v2.is_sequence_field(field) # type: ignore[return-value]
raise TypeError("field_info must be an instance of FieldInfo")
def serialize_sequence_value(*, field: ModelField, value: Any) -> Sequence[Any]:
if isinstance(field, v1.ModelField):
return v1.serialize_sequence_value(field=field, value=value)
elif PYDANTIC_V2:
else:
assert PYDANTIC_V2
from . import v2
return v2.serialize_sequence_value(field=field, value=value) # type: ignore[return-value]
raise TypeError("field_info must be an instance of FieldInfo")
def _model_rebuild(model: Type[BaseModel]) -> None:
@ -281,16 +281,6 @@ def _is_model_field(value: Any) -> bool:
return False
def _is_field_info(value: Any) -> bool:
if isinstance(value, v1.FieldInfo):
return True
elif PYDANTIC_V2:
from . import v2
return isinstance(value, v2.FieldInfo)
return False
def _is_model_class(value: Any) -> bool:
if lenient_issubclass(value, v1.BaseModel):
return True

56
fastapi/_compat/model_field.py

@ -1,4 +1,3 @@
from dataclasses import dataclass
from typing import (
Any,
Dict,
@ -9,49 +8,28 @@ from typing import (
from fastapi.types import IncEx
from pydantic.fields import FieldInfo
from typing_extensions import Literal
from typing_extensions import Literal, Protocol
@dataclass
class ModelField:
class ModelField(Protocol):
field_info: "FieldInfo"
name: str
mode: Literal["validation", "serialization"] = "validation"
_version: Literal["v1", "v2"] = "v1"
@property
def alias(self) -> str:
return self._model_field.alias
def alias(self) -> str: ...
@property
def required(self) -> bool:
return self._model_field.required
def required(self) -> bool: ...
@property
def default(self) -> Any:
return self._model_field.default
def default(self) -> Any: ...
@property
def type_(self) -> Any:
return self._model_field.type_
def type_(self) -> Any: ...
def __post_init__(self) -> None:
if self._version == "v1":
from . import v1
self._model_field = v1.ModelField(
field_info=self.field_info, name=self.name
)
else:
assert self._version == "v2"
from . import v2
self._model_field = v2.ModelField(
field_info=self.field_info, name=self.name, mode=self.mode
)
def get_default(self) -> Any:
return self._model_field.get_default()
def get_default(self) -> Any: ...
def validate(
self,
@ -59,8 +37,7 @@ class ModelField:
values: Dict[str, Any] = {}, # noqa: B006
*,
loc: Tuple[Union[int, str], ...] = (),
) -> Tuple[Any, Union[List[Dict[str, Any]], None]]:
return self._model_field.validate(value=value, values=values, loc=loc)
) -> Tuple[Any, Union[List[Dict[str, Any]], None]]: ...
def serialize(
self,
@ -73,19 +50,4 @@ class ModelField:
exclude_unset: bool = False,
exclude_defaults: bool = False,
exclude_none: bool = False,
) -> Any:
return self._model_field.serialize(
value=value,
mode=mode,
include=include,
exclude=exclude,
by_alias=by_alias,
exclude_unset=exclude_unset,
exclude_defaults=exclude_defaults,
exclude_none=exclude_none,
)
def __hash__(self) -> int:
# Each ModelField is unique for our purposes, to allow making a dict from
# ModelField to its JSON Schema.
return id(self)
) -> Any: ...

41
tests/test_compat.py

@ -2,7 +2,6 @@ from typing import Any, Dict, List, Union
from fastapi import FastAPI, UploadFile
from fastapi._compat import (
ModelField,
Undefined,
_get_model_config,
get_cached_model_fields,
@ -10,12 +9,12 @@ from fastapi._compat import (
is_uploadfile_sequence_annotation,
v1,
)
from fastapi._compat.shared import is_bytes_sequence_annotation
from fastapi._compat.shared import is_bytes_sequence_annotation, lenient_issubclass
from fastapi.testclient import TestClient
from pydantic import BaseConfig, BaseModel, ConfigDict
from pydantic import BaseModel, ConfigDict
from pydantic.fields import FieldInfo
from .utils import needs_pydanticv1, needs_pydanticv2
from .utils import needs_pydanticv2
@needs_pydanticv2
@ -28,29 +27,25 @@ def test_model_field_default_required():
assert field.default is Undefined
@needs_pydanticv1
def test_upload_file_dummy_with_info_plain_validator_function():
def test_v1_plain_validator_function():
# For coverage
assert UploadFile.__get_pydantic_core_schema__(str, lambda x: None) == {}
def func(v): # pragma: no cover
return v
result = v1.with_info_plain_validator_function(func)
assert result == {}
@needs_pydanticv1
def test_union_scalar_list():
def test_lenient_is_subclass():
# For coverage
# TODO: there might not be a current valid code path that uses this, it would
# potentially enable query parameters defined as both a scalar and a list
# but that would require more refactors, also not sure it's really useful
from fastapi._compat.v1 import is_pv1_scalar_field
field_info = FieldInfo()
field = ModelField(
name="foo",
field_info=field_info,
type_=Union[str, List[int]],
class_validators={},
model_config=BaseConfig,
)
assert not is_pv1_scalar_field(field)
assert lenient_issubclass(Union[str, int], str) is False
def test_is_model_field():
# For coverage
from fastapi._compat import _is_model_field
assert not _is_model_field(str)
@needs_pydanticv2

21
tests/test_pydantic_v1_v2_list.py

@ -185,6 +185,18 @@ def test_list_to_item_filter():
assert "internal_id" not in result
def test_list_to_item_filter_no_data():
response = client.post("/item-list-filter", json=[])
assert response.status_code == 200, response.text
assert response.json() == {
"title": "",
"size": 0,
"description": None,
"sub": {"name": ""},
"multi": [],
}
def test_list_to_list():
input_items = [
{"title": "Item 1", "size": 10, "sub": {"name": "Sub1"}},
@ -250,6 +262,15 @@ def test_list_to_list_filter():
assert "internal_id" not in item
def test_list_to_list_filter_no_data():
response = client.post(
"/item-list-to-list-filter",
json=[],
)
assert response.status_code == 200, response.text
assert response.json() == []
def test_list_validation_error():
response = client.post(
"/item-list",

18
tests/test_response_model_as_return_annotation.py

@ -2,6 +2,7 @@ from typing import List, Union
import pytest
from fastapi import FastAPI
from fastapi._compat import v1
from fastapi.exceptions import FastAPIError, ResponseValidationError
from fastapi.responses import JSONResponse, Response
from fastapi.testclient import TestClient
@ -509,6 +510,23 @@ def test_invalid_response_model_field():
assert "parameter response_model=None" in e.value.args[0]
# TODO: remove when dropping Pydantic v1 support
def test_invalid_response_model_field_pv1():
app = FastAPI()
class Model(v1.BaseModel):
foo: str
with pytest.raises(FastAPIError) as e:
@app.get("/")
def read_root() -> Union[Response, Model, None]:
return Response(content="Foo") # pragma: no cover
assert "valid Pydantic field type" in e.value.args[0]
assert "parameter response_model=None" in e.value.args[0]
def test_openapi_schema():
response = client.get("/openapi.json")
assert response.status_code == 200, response.text

Loading…
Cancel
Save