Browse Source

Merge 5a079219dc into eb75fd078e

pull/15708/merge
k-shimura7617 4 days ago
committed by GitHub
parent
commit
a315bf630f
No known key found for this signature in database GPG Key ID: B5690EEEBB952194
  1. 21
      fastapi/routing.py
  2. 160
      tests/test_sse.py

21
fastapi/routing.py

@ -523,9 +523,26 @@ def get_request_handler(
data_str: str | None = item.raw_data
elif item.data is not None:
if hasattr(item.data, "model_dump_json"):
data_str = item.data.model_dump_json()
data_str = item.data.model_dump_json(
include=response_model_include,
exclude=response_model_exclude,
by_alias=response_model_by_alias,
exclude_unset=response_model_exclude_unset,
exclude_defaults=response_model_exclude_defaults,
exclude_none=response_model_exclude_none,
)
else:
data_str = json.dumps(jsonable_encoder(item.data))
data_str = json.dumps(
jsonable_encoder(
item.data,
include=response_model_include,
exclude=response_model_exclude,
by_alias=response_model_by_alias,
exclude_unset=response_model_exclude_unset,
exclude_defaults=response_model_exclude_defaults,
exclude_none=response_model_exclude_none,
)
)
else:
data_str = None
return format_sse_event(

160
tests/test_sse.py

@ -8,7 +8,7 @@ from fastapi import APIRouter, FastAPI
from fastapi.responses import EventSourceResponse
from fastapi.sse import ServerSentEvent
from fastapi.testclient import TestClient
from pydantic import BaseModel
from pydantic import BaseModel, Field
class Item(BaseModel):
@ -16,6 +16,17 @@ class Item(BaseModel):
description: str | None = None
class AliasedItem(BaseModel):
name: str = Field(serialization_alias="itemName")
class AliasedItemWithDefaults(BaseModel):
name: str = Field(serialization_alias="itemName")
description: str | None = None
hidden: str = "secret"
status: str = "new"
items = [
Item(name="Plumbus", description="A multi-purpose household device."),
Item(name="Portal Gun", description="A portal opening device."),
@ -62,6 +73,83 @@ async def sse_items_event():
yield ServerSentEvent(data="retry-test", retry=5000)
@app.get("/items/stream-sse-event-model-data", response_class=EventSourceResponse)
async def sse_items_event_model_data():
yield ServerSentEvent(data=AliasedItem(name="Portal Gun"))
@app.get(
"/items/stream-sse-event-model-data-no-alias",
response_class=EventSourceResponse,
response_model_by_alias=False,
)
async def sse_items_event_model_data_no_alias():
yield ServerSentEvent(data=AliasedItem(name="Portal Gun"))
@app.get(
"/items/stream-sse-event-model-data-include",
response_class=EventSourceResponse,
response_model_include={"name"},
)
async def sse_items_event_model_data_include():
yield ServerSentEvent(data=AliasedItemWithDefaults(name="Portal Gun"))
@app.get(
"/items/stream-sse-event-model-data-exclude",
response_class=EventSourceResponse,
response_model_exclude={"hidden"},
)
async def sse_items_event_model_data_exclude():
yield ServerSentEvent(data=AliasedItemWithDefaults(name="Portal Gun"))
@app.get(
"/items/stream-sse-event-model-data-exclude-none",
response_class=EventSourceResponse,
response_model_exclude_none=True,
)
async def sse_items_event_model_data_exclude_none():
yield ServerSentEvent(data=AliasedItemWithDefaults(name="Portal Gun"))
@app.get(
"/items/stream-sse-event-model-data-exclude-unset",
response_class=EventSourceResponse,
response_model_exclude_unset=True,
)
async def sse_items_event_model_data_exclude_unset():
yield ServerSentEvent(data=AliasedItemWithDefaults(name="Portal Gun"))
@app.get(
"/items/stream-sse-event-model-data-exclude-defaults",
response_class=EventSourceResponse,
response_model_exclude_defaults=True,
)
async def sse_items_event_model_data_exclude_defaults():
yield ServerSentEvent(data=AliasedItemWithDefaults(name="Portal Gun"))
@app.get(
"/items/stream-sse-event-dict-data-include",
response_class=EventSourceResponse,
response_model_include={"name"},
)
async def sse_items_event_dict_data_include():
yield ServerSentEvent(data={"name": "Portal Gun", "hidden": "secret"})
@app.get(
"/items/stream-sse-event-list-model-data-no-alias",
response_class=EventSourceResponse,
response_model_by_alias=False,
)
async def sse_items_event_list_model_data_no_alias():
yield ServerSentEvent(data=[AliasedItem(name="Portal Gun")])
@app.get("/items/stream-mixed", response_class=EventSourceResponse)
async def sse_items_mixed() -> AsyncIterable[Item]:
yield items[0]
@ -199,6 +287,76 @@ def test_sse_events_with_fields(client: TestClient):
assert 'data: "retry-test"\n' in text
def test_sse_event_model_data_uses_serialization_alias(client: TestClient):
response = client.get("/items/stream-sse-event-model-data")
assert response.status_code == 200
assert 'data: {"itemName":"Portal Gun"}\n' in response.text
assert '"name"' not in response.text
def test_sse_event_model_data_respects_response_model_by_alias_false(
client: TestClient,
):
response = client.get("/items/stream-sse-event-model-data-no-alias")
assert response.status_code == 200
assert 'data: {"name":"Portal Gun"}\n' in response.text
assert '"itemName"' not in response.text
@pytest.mark.parametrize(
("path", "expected_data"),
[
(
"/items/stream-sse-event-model-data-include",
'{"itemName":"Portal Gun"}',
),
(
"/items/stream-sse-event-model-data-exclude",
'{"itemName":"Portal Gun","description":null,"status":"new"}',
),
(
"/items/stream-sse-event-model-data-exclude-none",
'{"itemName":"Portal Gun","hidden":"secret","status":"new"}',
),
(
"/items/stream-sse-event-model-data-exclude-unset",
'{"itemName":"Portal Gun"}',
),
(
"/items/stream-sse-event-model-data-exclude-defaults",
'{"itemName":"Portal Gun"}',
),
],
)
def test_sse_event_model_data_respects_response_model_serialization_options(
client: TestClient, path: str, expected_data: str
):
response = client.get(path)
assert response.status_code == 200
assert response.text == f"data: {expected_data}\n\n"
@pytest.mark.parametrize(
("path", "expected_data"),
[
(
"/items/stream-sse-event-dict-data-include",
'{"name": "Portal Gun"}',
),
(
"/items/stream-sse-event-list-model-data-no-alias",
'[{"name": "Portal Gun"}]',
),
],
)
def test_sse_event_jsonable_encoder_data_respects_response_model_serialization_options(
client: TestClient, path: str, expected_data: str
):
response = client.get(path)
assert response.status_code == 200
assert response.text == f"data: {expected_data}\n\n"
def test_mixed_plain_and_sse_events(client: TestClient):
response = client.get("/items/stream-mixed")
assert response.status_code == 200

Loading…
Cancel
Save