diff --git a/fastapi/routing.py b/fastapi/routing.py index c442b122b..faa0638a4 100644 --- a/fastapi/routing.py +++ b/fastapi/routing.py @@ -45,6 +45,7 @@ from fastapi._compat import ( ModelField, Undefined, lenient_issubclass, + value_is_sequence, ) from fastapi.datastructures import Default, DefaultPlaceholder from fastapi.dependencies.models import Dependant @@ -290,6 +291,16 @@ def _extract_endpoint_context(func: Any) -> EndpointContext: return ctx +def _wrap_incex_for_sequence(incex: IncEx | None) -> IncEx | None: + # Pydantic interprets top-level include/exclude keys on a sequence type as + # item indices, so field names would be silently ignored. Wrap them in + # "__all__" to apply them to every item, unless indices or "__all__" are + # already used explicitly. + if incex and all(isinstance(key, str) and key != "__all__" for key in incex): + return {"__all__": incex} + return incex + + async def serialize_response( *, field: ModelField | None = None, @@ -318,6 +329,9 @@ async def serialize_response( body=response_content, endpoint_ctx=ctx, ) + if (include or exclude) and value_is_sequence(value): + include = _wrap_incex_for_sequence(include) + exclude = _wrap_incex_for_sequence(exclude) serializer = field.serialize_json if dump_json else field.serialize return serializer( value, diff --git a/tests/test_response_model_include_exclude.py b/tests/test_response_model_include_exclude.py index 73c3591eb..1a19de0b0 100644 --- a/tests/test_response_model_include_exclude.py +++ b/tests/test_response_model_include_exclude.py @@ -1,4 +1,5 @@ from fastapi import FastAPI +from fastapi.responses import JSONResponse from fastapi.testclient import TestClient from pydantic import BaseModel @@ -110,6 +111,73 @@ def mixed_dict(): } +@app.get( + "/list_include", + response_model=list[Model1], + response_model_include={"foo"}, +) +def list_include(): + return [ + Model1(foo="list_include model foo", bar="list_include model bar"), + Model1(foo="list_include model2 foo", bar="list_include model2 bar"), + ] + + +@app.get( + "/list_exclude", + response_model=list[Model1], + response_model_exclude={"bar"}, +) +def list_exclude(): + return [ + Model1(foo="list_exclude model foo", bar="list_exclude model bar"), + Model1(foo="list_exclude model2 foo", bar="list_exclude model2 bar"), + ] + + +@app.get( + "/list_exclude_nested", + response_model=list[Model2], + response_model_exclude={"ref": {"bar"}}, +) +def list_exclude_nested(): + return [ + Model2( + ref=Model1( + foo="list_exclude_nested model foo", + bar="list_exclude_nested model bar", + ), + baz="list_exclude_nested model2 baz", + ) + ] + + +@app.get( + "/list_exclude_all", + response_model=list[Model1], + response_model_exclude={"__all__": {"bar"}}, +) +def list_exclude_all(): + return [ + Model1(foo="list_exclude_all model foo", bar="list_exclude_all model bar"), + ] + + +@app.get( + "/list_exclude_response_class", + response_model=list[Model1], + response_model_exclude={"bar"}, + response_class=JSONResponse, +) +def list_exclude_response_class(): + return [ + Model1( + foo="list_exclude_response_class model foo", + bar="list_exclude_response_class model bar", + ), + ] + + client = TestClient(app) @@ -173,3 +241,44 @@ def test_nested_include_mixed_dict(): "ref": {"foo": "mixed_dict model foo", "bar": "mixed_dict model bar"}, }, } + + +def test_list_include(): + response = client.get("/list_include") + assert response.status_code == 200, response.text + assert response.json() == [ + {"foo": "list_include model foo"}, + {"foo": "list_include model2 foo"}, + ] + + +def test_list_exclude(): + response = client.get("/list_exclude") + assert response.status_code == 200, response.text + assert response.json() == [ + {"foo": "list_exclude model foo"}, + {"foo": "list_exclude model2 foo"}, + ] + + +def test_list_exclude_nested(): + response = client.get("/list_exclude_nested") + assert response.status_code == 200, response.text + assert response.json() == [ + { + "ref": {"foo": "list_exclude_nested model foo"}, + "baz": "list_exclude_nested model2 baz", + } + ] + + +def test_list_exclude_all(): + response = client.get("/list_exclude_all") + assert response.status_code == 200, response.text + assert response.json() == [{"foo": "list_exclude_all model foo"}] + + +def test_list_exclude_response_class(): + response = client.get("/list_exclude_response_class") + assert response.status_code == 200, response.text + assert response.json() == [{"foo": "list_exclude_response_class model foo"}]