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: def copy_field_info(*, field_info: FieldInfo, annotation: Any) -> FieldInfo:
if isinstance(field_info, v1.FieldInfo): if isinstance(field_info, v1.FieldInfo):
return v1.copy_field_info(field_info=field_info, annotation=annotation) return v1.copy_field_info(field_info=field_info, annotation=annotation)
elif PYDANTIC_V2: else:
assert PYDANTIC_V2
from . import v2 from . import v2
return v2.copy_field_info(field_info=field_info, annotation=annotation) 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( def create_body_model(
@ -111,11 +111,11 @@ def create_body_model(
) -> Type[BaseModel]: ) -> Type[BaseModel]:
if fields and isinstance(fields[0], v1.ModelField): if fields and isinstance(fields[0], v1.ModelField):
return v1.create_body_model(fields=fields, model_name=model_name) # type: ignore[return-value] 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 from . import v2
return v2.create_body_model(fields=fields, model_name=model_name) # type: ignore[return-value] 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( def get_annotation_from_field_info(
@ -125,73 +125,73 @@ def get_annotation_from_field_info(
return v1.get_annotation_from_field_info( return v1.get_annotation_from_field_info(
annotation=annotation, field_info=field_info, field_name=field_name annotation=annotation, field_info=field_info, field_name=field_name
) )
elif PYDANTIC_V2: else:
assert PYDANTIC_V2
from . import v2 from . import v2
return v2.get_annotation_from_field_info( return v2.get_annotation_from_field_info(
annotation=annotation, field_info=field_info, field_name=field_name 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: def is_bytes_field(field: ModelField) -> bool:
if isinstance(field, v1.ModelField): if isinstance(field, v1.ModelField):
return v1.is_bytes_field(field) return v1.is_bytes_field(field)
elif PYDANTIC_V2: else:
assert PYDANTIC_V2
from . import v2 from . import v2
return v2.is_bytes_field(field) # type: ignore[return-value] 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: def is_bytes_sequence_field(field: ModelField) -> bool:
if isinstance(field, v1.ModelField): if isinstance(field, v1.ModelField):
return v1.is_bytes_sequence_field(field) return v1.is_bytes_sequence_field(field)
elif PYDANTIC_V2: else:
assert PYDANTIC_V2
from . import v2 from . import v2
return v2.is_bytes_sequence_field(field) # type: ignore[return-value] 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: def is_scalar_field(field: ModelField) -> bool:
if isinstance(field, v1.ModelField): if isinstance(field, v1.ModelField):
return v1.is_scalar_field(field) return v1.is_scalar_field(field)
elif PYDANTIC_V2: else:
assert PYDANTIC_V2
from . import v2 from . import v2
return v2.is_scalar_field(field) # type: ignore[return-value] 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: def is_scalar_sequence_field(field: ModelField) -> bool:
if isinstance(field, v1.ModelField): if isinstance(field, v1.ModelField):
return v1.is_scalar_sequence_field(field) return v1.is_scalar_sequence_field(field)
elif PYDANTIC_V2: else:
assert PYDANTIC_V2
from . import v2 from . import v2
return v2.is_scalar_sequence_field(field) # type: ignore[return-value] 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: def is_sequence_field(field: ModelField) -> bool:
if isinstance(field, v1.ModelField): if isinstance(field, v1.ModelField):
return v1.is_sequence_field(field) return v1.is_sequence_field(field)
elif PYDANTIC_V2: else:
assert PYDANTIC_V2
from . import v2 from . import v2
return v2.is_sequence_field(field) # type: ignore[return-value] 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]: def serialize_sequence_value(*, field: ModelField, value: Any) -> Sequence[Any]:
if isinstance(field, v1.ModelField): if isinstance(field, v1.ModelField):
return v1.serialize_sequence_value(field=field, value=value) return v1.serialize_sequence_value(field=field, value=value)
elif PYDANTIC_V2: else:
assert PYDANTIC_V2
from . import v2 from . import v2
return v2.serialize_sequence_value(field=field, value=value) # type: ignore[return-value] 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: def _model_rebuild(model: Type[BaseModel]) -> None:
@ -281,16 +281,6 @@ def _is_model_field(value: Any) -> bool:
return False 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: def _is_model_class(value: Any) -> bool:
if lenient_issubclass(value, v1.BaseModel): if lenient_issubclass(value, v1.BaseModel):
return True return True

56
fastapi/_compat/model_field.py

@ -1,4 +1,3 @@
from dataclasses import dataclass
from typing import ( from typing import (
Any, Any,
Dict, Dict,
@ -9,49 +8,28 @@ from typing import (
from fastapi.types import IncEx from fastapi.types import IncEx
from pydantic.fields import FieldInfo from pydantic.fields import FieldInfo
from typing_extensions import Literal from typing_extensions import Literal, Protocol
@dataclass class ModelField(Protocol):
class ModelField:
field_info: "FieldInfo" field_info: "FieldInfo"
name: str name: str
mode: Literal["validation", "serialization"] = "validation" mode: Literal["validation", "serialization"] = "validation"
_version: Literal["v1", "v2"] = "v1" _version: Literal["v1", "v2"] = "v1"
@property @property
def alias(self) -> str: def alias(self) -> str: ...
return self._model_field.alias
@property @property
def required(self) -> bool: def required(self) -> bool: ...
return self._model_field.required
@property @property
def default(self) -> Any: def default(self) -> Any: ...
return self._model_field.default
@property @property
def type_(self) -> Any: def type_(self) -> Any: ...
return self._model_field.type_
def __post_init__(self) -> None: def get_default(self) -> Any: ...
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 validate( def validate(
self, self,
@ -59,8 +37,7 @@ class ModelField:
values: Dict[str, Any] = {}, # noqa: B006 values: Dict[str, Any] = {}, # noqa: B006
*, *,
loc: Tuple[Union[int, str], ...] = (), loc: Tuple[Union[int, str], ...] = (),
) -> Tuple[Any, Union[List[Dict[str, Any]], None]]: ) -> Tuple[Any, Union[List[Dict[str, Any]], None]]: ...
return self._model_field.validate(value=value, values=values, loc=loc)
def serialize( def serialize(
self, self,
@ -73,19 +50,4 @@ class ModelField:
exclude_unset: bool = False, exclude_unset: bool = False,
exclude_defaults: bool = False, exclude_defaults: bool = False,
exclude_none: bool = False, exclude_none: bool = False,
) -> Any: ) -> 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)

41
tests/test_compat.py

@ -2,7 +2,6 @@ from typing import Any, Dict, List, Union
from fastapi import FastAPI, UploadFile from fastapi import FastAPI, UploadFile
from fastapi._compat import ( from fastapi._compat import (
ModelField,
Undefined, Undefined,
_get_model_config, _get_model_config,
get_cached_model_fields, get_cached_model_fields,
@ -10,12 +9,12 @@ from fastapi._compat import (
is_uploadfile_sequence_annotation, is_uploadfile_sequence_annotation,
v1, 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 fastapi.testclient import TestClient
from pydantic import BaseConfig, BaseModel, ConfigDict from pydantic import BaseModel, ConfigDict
from pydantic.fields import FieldInfo from pydantic.fields import FieldInfo
from .utils import needs_pydanticv1, needs_pydanticv2 from .utils import needs_pydanticv2
@needs_pydanticv2 @needs_pydanticv2
@ -28,29 +27,25 @@ def test_model_field_default_required():
assert field.default is Undefined assert field.default is Undefined
@needs_pydanticv1 def test_v1_plain_validator_function():
def test_upload_file_dummy_with_info_plain_validator_function():
# For coverage # 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 # For coverage
# TODO: there might not be a current valid code path that uses this, it would assert lenient_issubclass(Union[str, int], str) is False
# 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 def test_is_model_field():
# For coverage
field_info = FieldInfo() from fastapi._compat import _is_model_field
field = ModelField(
name="foo", assert not _is_model_field(str)
field_info=field_info,
type_=Union[str, List[int]],
class_validators={},
model_config=BaseConfig,
)
assert not is_pv1_scalar_field(field)
@needs_pydanticv2 @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 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(): def test_list_to_list():
input_items = [ input_items = [
{"title": "Item 1", "size": 10, "sub": {"name": "Sub1"}}, {"title": "Item 1", "size": 10, "sub": {"name": "Sub1"}},
@ -250,6 +262,15 @@ def test_list_to_list_filter():
assert "internal_id" not in item 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(): def test_list_validation_error():
response = client.post( response = client.post(
"/item-list", "/item-list",

18
tests/test_response_model_as_return_annotation.py

@ -2,6 +2,7 @@ from typing import List, Union
import pytest import pytest
from fastapi import FastAPI from fastapi import FastAPI
from fastapi._compat import v1
from fastapi.exceptions import FastAPIError, ResponseValidationError from fastapi.exceptions import FastAPIError, ResponseValidationError
from fastapi.responses import JSONResponse, Response from fastapi.responses import JSONResponse, Response
from fastapi.testclient import TestClient 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] 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(): def test_openapi_schema():
response = client.get("/openapi.json") response = client.get("/openapi.json")
assert response.status_code == 200, response.text assert response.status_code == 200, response.text

Loading…
Cancel
Save