k-shimura7617 5 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 data_str: str | None = item.raw_data
elif item.data is not None: elif item.data is not None:
if hasattr(item.data, "model_dump_json"): 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: 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: else:
data_str = None data_str = None
return format_sse_event( 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.responses import EventSourceResponse
from fastapi.sse import ServerSentEvent from fastapi.sse import ServerSentEvent
from fastapi.testclient import TestClient from fastapi.testclient import TestClient
from pydantic import BaseModel from pydantic import BaseModel, Field
class Item(BaseModel): class Item(BaseModel):
@ -16,6 +16,17 @@ class Item(BaseModel):
description: str | None = None 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 = [ items = [
Item(name="Plumbus", description="A multi-purpose household device."), Item(name="Plumbus", description="A multi-purpose household device."),
Item(name="Portal Gun", description="A portal opening 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) 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) @app.get("/items/stream-mixed", response_class=EventSourceResponse)
async def sse_items_mixed() -> AsyncIterable[Item]: async def sse_items_mixed() -> AsyncIterable[Item]:
yield items[0] yield items[0]
@ -199,6 +287,76 @@ def test_sse_events_with_fields(client: TestClient):
assert 'data: "retry-test"\n' in text 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): def test_mixed_plain_and_sse_events(client: TestClient):
response = client.get("/items/stream-mixed") response = client.get("/items/stream-mixed")
assert response.status_code == 200 assert response.status_code == 200

Loading…
Cancel
Save