← oracle_full

fastapi_15030

resolved RESOLVED UNSUBMITTED PASS · None tool calls · 0 s · fastapi/fastapi

Task input

✨ Add support for Server Sent Events

✨ Add support for Server Sent Events

This will be available in FastAPI 0.135.0, released in the next hours.

Tool calls (0)

#ToolArgumentsResult
No trace captured.

Patch

--- a/docs_src/server_sent_events/tutorial001_py310.py
+++ b/docs_src/server_sent_events/tutorial001_py310.py
@@ -0,0 +1,43 @@
+from collections.abc import AsyncIterable, Iterable
+
+from fastapi import FastAPI
+from fastapi.sse import EventSourceResponse
+from pydantic import BaseModel
+
+app = FastAPI()
+
+
+class Item(BaseModel):
+    name: str
+    description: str | None
+
+
+items = [
+    Item(name="Plumbus", description="A multi-purpose household device."),
+    Item(name="Portal Gun", description="A portal opening device."),
+    Item(name="Meeseeks Box", description="A box that summons a Meeseeks."),
+]
+
+
+@app.get("/items/stream", response_class=EventSourceResponse)
+async def sse_items() -> AsyncIterable[Item]:
+    for item in items:
+        yield item
+
+
+@app.get("/items/stream-no-async", response_class=EventSourceResponse)
+def sse_items_no_async() -> Iterable[Item]:
+    for item in items:
+        yield item
+
+
+@app.get("/items/stream-no-annotation", response_class=EventSourceResponse)
+async def sse_items_no_annotation():
+    for item in items:
+        yield item
+
+
+@app.get("/items/stream-no-async-no-annotation", response_class=EventSourceResponse)
+def sse_items_no_async_no_annotation():
+    for item in items:
+        yield item
--- a/docs_src/server_sent_events/tutorial002_py310.py
+++ b/docs_src/server_sent_events/tutorial002_py310.py
@@ -0,0 +1,26 @@
+from collections.abc import AsyncIterable
+
+from fastapi import FastAPI
+from fastapi.sse import EventSourceResponse, ServerSentEvent
+from pydantic import BaseModel
+
+app = FastAPI()
+
+
+class Item(BaseModel):
+    name: str
+    price: float
+
+
+items = [
+    Item(name="Plumbus", price=32.99),
+    Item(name="Portal Gun", price=999.99),
+    Item(name="Meeseeks Box", price=49.99),
+]
+
+
+@app.get("/items/stream", response_class=EventSourceResponse)
+async def stream_items() -> AsyncIterable[ServerSentEvent]:
+    yield ServerSentEvent(comment="stream of item updates")
+    for i, item in enumerate(items):
+        yield ServerSentEvent(data=item, event="item_update", id=str(i + 1), retry=5000)
--- a/docs_src/server_sent_events/tutorial003_py310.py
+++ b/docs_src/server_sent_events/tutorial003_py310.py
@@ -0,0 +1,17 @@
+from collections.abc import AsyncIterable
+
+from fastapi import FastAPI
+from fastapi.sse import EventSourceResponse, ServerSentEvent
+
+app = FastAPI()
+
+
+@app.get("/logs/stream", response_class=EventSourceResponse)
+async def stream_logs() -> AsyncIterable[ServerSentEvent]:
+    logs = [
+        "2025-01-01 INFO  Application started",
+        "2025-01-01 DEBUG Connected to database",
+        "2025-01-01 WARN  High memory usage detected",
+    ]
+    for log_line in logs:
+        yield ServerSentEvent(raw_data=log_line)
--- a/docs_src/server_sent_events/tutorial004_py310.py
+++ b/docs_src/server_sent_events/tutorial004_py310.py
@@ -0,0 +1,31 @@
+from collections.abc import AsyncIterable
+from typing import Annotated
+
+from fastapi import FastAPI, Header
+from fastapi.sse import EventSourceResponse, ServerSentEvent
+from pydantic import BaseModel
+
+app = FastAPI()
+
+
+class Item(BaseModel):
+    name: str
+    price: float
+
+
+items = [
+    Item(name="Plumbus", price=32.99),
+    Item(name="Portal Gun", price=999.99),
+    Item(name="Meeseeks Box", price=49.99),
+]
+
+
+@app.get("/items/stream", response_class=EventSourceResponse)
+async def stream_items(
+    last_event_id: Annotated[int | None, Header()] = None,
+) -> AsyncIterable[ServerSentEvent]:
+    start = last_event_id + 1 if last_event_id is not None else 0
+    for i, item in enumerate(items):
+        if i < start:
+            continue
+        yield ServerSentEvent(data=item, id=str(i))
--- a/docs_src/server_sent_events/tutorial005_py310.py
+++ b/docs_src/server_sent_events/tutorial005_py310.py
@@ -0,0 +1,19 @@
+from collections.abc import AsyncIterable
+
+from fastapi import FastAPI
+from fastapi.sse import EventSourceResponse, ServerSentEvent
+from pydantic import BaseModel
+
+app = FastAPI()
+
+
+class Prompt(BaseModel):
+    text: str
+
+
+@app.post("/chat/stream", response_class=EventSourceResponse)
+async def stream_chat(prompt: Prompt) -> AsyncIterable[ServerSentEvent]:
+    words = prompt.text.split()
+    for word in words:
+        yield ServerSentEvent(data=word, event="token")
+    yield ServerSentEvent(raw_data="[DONE]", event="done")
--- a/fastapi/openapi/utils.py
+++ b/fastapi/openapi/utils.py
@@ -29,6 +29,7 @@
 from fastapi.openapi.models import OpenAPI
 from fastapi.params import Body, ParamTypes
 from fastapi.responses import Response
+from fastapi.sse import _SSE_EVENT_SCHEMA
 from fastapi.types import ModelNameMap
 from fastapi.utils import (
     deep_dict_update,
@@ -372,6 +373,26 @@ def get_openapi_path(
                     operation.setdefault("responses", {}).setdefault(
                         status_code, {}
                     ).setdefault("content", {})["application/jsonl"] = jsonl_content
+                elif route.is_sse_stream:
+                    sse_content: dict[str, Any] = {}
+                    item_schema = copy.deepcopy(_SSE_EVENT_SCHEMA)
+                    if route.stream_item_field:
+                        content_schema = get_schema_from_model_field(
+                            field=route.stream_item_field,
+                            model_name_map=model_name_map,
+                            field_mapping=field_mapping,
+                            separate_input_output_schemas=separate_input_output_schemas,
+                        )
+                        item_schema["required"] = ["data"]
+                        item_schema["properties"]["data"] = {
+                            "type": "string",
+                            "contentMediaType": "application/json",
+                            "contentSchema": content_schema,
+                        }
+                    sse_content["itemSchema"] = item_schema
+                    operation.setdefault("responses", {}).setdefault(
+                        status_code, {}
+                    ).setdefault("content", {})["text/event-stream"] = sse_content
                 elif route_response_media_type:
                     response_schema = {"type": "string"}
                     if lenient_issubclass(current_response_class, JSONResponse):
--- a/fastapi/responses.py
+++ b/fastapi/responses.py
@@ -1,6 +1,7 @@
 from typing import Any
 
 from fastapi.exceptions import FastAPIDeprecationWarning
+from fastapi.sse import EventSourceResponse as EventSourceResponse  # noqa
 from starlette.responses import FileResponse as FileResponse  # noqa
 from starlette.responses import HTMLResponse as HTMLResponse  # noqa
 from starlette.responses import JSONResponse as JSONResponse  # noqa
--- a/fastapi/routing.py
+++ b/fastapi/routing.py
@@ -56,6 +56,13 @@
     ResponseValidationError,
     WebSocketRequestValidationError,
 )
+from fastapi.sse import (
+    _PING_INTERVAL,
+    KEEPALIVE_COMMENT,
+    EventSourceResponse,
+    ServerSentEvent,
+    format_sse_event,
+)
 from fastapi.types import DecoratedCallable, IncEx
 from fastapi.utils import (
     create_model_field,
@@ -66,7 +73,7 @@
 from starlette import routing
 from starlette._exception_handler import wrap_app_handling_exceptions
 from starlette._utils import is_async_callable
-from starlette.concurrency import run_in_threadpool
+from starlette.concurrency import iterate_in_threadpool, run_in_threadpool
 from starlette.exceptions import HTTPException
 from starlette.requests import Request
 from starlette.responses import JSONResponse, Response, StreamingResponse
@@ -361,6 +368,7 @@ def get_request_handler(
         actual_response_class: type[Response] = response_class.value
     else:
         actual_response_class = response_class
+    is_sse_stream = lenient_issubclass(actual_response_class, EventSourceResponse)
     if isinstance(strict_content_type, DefaultPlaceholder):
         actual_strict_content_type: bool = strict_content_type.value
     else:
@@ -452,35 +460,125 @@ async def app(request: Request) -> Response:
         errors = solved_result.errors
         assert dependant.call  # For types
         if not errors:
-            if is_json_stream:
-                # Generator endpoint: stream as JSONL
+            # Shared serializer for stream items (JSONL and SSE).
+            # Validates against stream_item_field when set, then
+            # serializes to JSON bytes.
+            def _serialize_data(data: Any) -> bytes:
+                if stream_item_field:
+                    value, errors_ = stream_item_field.validate(
+                        data, {}, loc=("response",)
+                    )
+                    if errors_:
+                        ctx = endpoint_ctx or EndpointContext()
+                        raise ResponseValidationError(
+                            errors=errors_,
+                            body=data,
+                            endpoint_ctx=ctx,
+                        )
+                    return stream_item_field.serialize_json(
+                        value,
+                        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 = jsonable_encoder(data)
+                    return json.dumps(data).encode("utf-8")
+
+            if is_sse_stream:
+                # Generator endpoint: stream as Server-Sent Events
                 gen = dependant.call(**solved_result.values)
 
-                def _serialize_item(item: Any) -> bytes:
-                    if stream_item_field:
-                        value, errors = stream_item_field.validate(
-                            item, {}, loc=("response",)
-                        )
-                        if errors:
-                            ctx = endpoint_ctx or EndpointContext()
-                            raise ResponseValidationError(
-                                errors=errors,
-                                body=item,
-                                endpoint_ctx=ctx,
-                            )
-                        line = stream_item_field.serialize_json(
-                            value,
-                            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,
+                def _serialize_sse_item(item: Any) -> bytes:
+                    if isinstance(item, ServerSentEvent):
+                        # User controls the event structure.
+                        # Serialize the data payload if present.
+                        # For ServerSentEvent items we skip stream_item_field
+                        # validation (the user may mix types intentionally).
+                        if item.raw_data is not None:
+                            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()
+                            else:
+                                data_str = json.dumps(jsonable_encoder(item.data))
+                        else:
+                            data_str = None
+                        return format_sse_event(
+                            data_str=data_str,
+                            event=item.event,
+                            id=item.id,
+                            retry=item.retry,
+                            comment=item.comment,
                         )
-                        return line + b"\n"
                     else:
-                        data = jsonable_encoder(item)
-                        return json.dumps(data).encode("utf-8") + b"\n"
+                        # Plain object: validate + serialize via
+                        # stream_item_field (if set) and wrap in data field
+                        return format_sse_event(
+                            data_str=_serialize_data(item).decode("utf-8")
+                        )
+
+                if dependant.is_async_gen_callable:
+                    sse_aiter: AsyncIterator[Any] = gen.__aiter__()
+                else:
+                    sse_aiter = iterate_in_threadpool(gen)
+
+                async def _async_stream_sse() -> AsyncIterator[bytes]:
+                    # Use a memory stream to decouple generator iteration
+                    # from the keepalive timer. A producer task pulls items
+                    # from the generator independently, so
+                    # `anyio.fail_after` never wraps the generator's
+                    # `__anext__` directly - avoiding CancelledError that
+                    # would finalize the generator and also working for sync
+                    # generators running in a thread pool.
+                    send_stream, receive_stream = anyio.create_memory_object_stream[
+                        bytes
+                    ](max_buffer_size=1)
+
+                    async def _producer() -> None:
+                        async with send_stream:
+                            async for raw_item in sse_aiter:
+                                await send_stream.send(_serialize_sse_item(raw_item))
+
+                    async with anyio.create_task_group() as tg:
+                        tg.start_soon(_producer)
+                        async with receive_stream:
+                            try:
+                                while True:
+                                    try:
+                                        with anyio.fail_after(_PING_INTERVAL):
+                                            data = await receive_stream.receive()
+                                        yield data
+                                        # To allow for cancellation to trigger
+                                        # Ref: https://github.com/fastapi/fastapi/issues/14680
+                                        await anyio.sleep(0)
+                                    except TimeoutError:
+                                        yield KEEPALIVE_COMMENT
+                            except anyio.EndOfStream:
+                                pass
+
+                sse_stream_content: AsyncIterator[bytes] | Iterator[bytes] = (
+                    _async_stream_sse()
+                )
+
+                response = StreamingResponse(
+                    sse_stream_content,
+                    media_type="text/event-stream",
+                    background=solved_result.background_tasks,
+                )
+                response.headers["Cache-Control"] = "no-cache"
+                # For Nginx proxies to not buffer server sent events
+                response.headers["X-Accel-Buffering"] = "no"
+                response.headers.raw.extend(solved_result.response.headers.raw)
+            elif is_json_stream:
+                # Generator endpoint: stream as JSONL
+                gen = dependant.call(**solved_result.values)
+
+                def _serialize_item(item: Any) -> bytes:
+                    return _serialize_data(item) + b"\n"
 
                 if dependant.is_async_gen_callable:
 
@@ -491,7 +589,7 @@ async def _async_stream_jsonl() -> AsyncIterator[bytes]:
                             # Ref: https://github.com/fastapi/fastapi/issues/14680
                             await anyio.sleep(0)
 
-                    stream_content: AsyncIterator[bytes] | Iterator[bytes] = (
+                    jsonl_stream_content: AsyncIterator[bytes] | Iterator[bytes] = (
                         _async_stream_jsonl()
                     )
                 else:
@@ -500,10 +598,10 @@ def _sync_stream_jsonl() -> Iterator[bytes]:
                         for item in gen:
                             yield _serialize_item(item)
 
-                    stream_content = _sync_stream_jsonl()
+                    jsonl_stream_content = _sync_stream_jsonl()
 
                 response = StreamingResponse(
-                    stream_content,
+                    jsonl_stream_content,
                     media_type="application/jsonl",
                     background=solved_result.background_tasks,
                 )
@@ -709,9 +807,16 @@ def __init__(
             else:
                 stream_item = get_stream_item_type(return_annotation)
                 if stream_item is not None:
-                    # Only extract item type for JSONL streaming when no
-                    # explicit response_class (e.g. StreamingResponse) was set
-                    if isinstance(response_class, DefaultPlaceholder):
+                    # Extract item type for JSONL or SSE streaming when
+                    # response_class is DefaultPlaceholder (JSONL) or
+                    # EventSourceResponse (SSE).
+                    # ServerSentEvent is excluded: it's a transport
+                    # wrapper, not a data model, so it shouldn't feed
+                    # into validation or OpenAPI schema generation.
+                    if (
+                        isinstance(response_class, DefaultPlaceholder)
+                        or lenient_issubclass(response_class, EventSourceResponse)
+                    ) and not lenient_issubclass(stream_item, ServerSentEvent):
                         self.stream_item_type = stream_item
                     response_model = None
                 else:
@@ -814,11 +919,16 @@ def __init__(
             name=self.unique_id,
             embed_body_fields=self._embed_body_fields,
         )
-        # Detect generator endpoints that should stream as JSONL
-        # (only when no explicit response_class like StreamingResponse is set)
-        self.is_json_stream = isinstance(response_class, DefaultPlaceholder) and (
+        # Detect generator endpoints that should stream as JSONL or SSE
+        is_generator = (
             self.dependant.is_async_gen_callable or self.dependant.is_gen_callable
         )
+        self.is_sse_stream = is_generator and lenient_issubclass(
+            response_class, EventSourceResponse
+        )
+        self.is_json_stream = is_generator and isinstance(
+            response_class, DefaultPlaceholder
+        )
         self.app = request_response(self.get_route_handler())
 
     def get_route_handler(self) -> Callable[[Request], Coroutine[Any, Any, Response]]:
--- a/fastapi/sse.py
+++ b/fastapi/sse.py
@@ -0,0 +1,222 @@
+from typing import Annotated, Any
+
+from annotated_doc import Doc
+from pydantic import AfterValidator, BaseModel, Field, model_validator
+from starlette.responses import StreamingResponse
+
+# Canonical SSE event schema matching the OpenAPI 3.2 spec
+# (Section 4.14.4 "Special Considerations for Server-Sent Events")
+_SSE_EVENT_SCHEMA: dict[str, Any] = {
+    "type": "object",
+    "properties": {
+        "data": {"type": "string"},
+        "event": {"type": "string"},
+        "id": {"type": "string"},
+        "retry": {"type": "integer", "minimum": 0},
+    },
+}
+
+
+class EventSourceResponse(StreamingResponse):
+    """Streaming response with `text/event-stream` media type.
+
+    Use as `response_class=EventSourceResponse` on a *path operation* that uses `yield`
+    to enable Server Sent Events (SSE) responses.
+
+    Works with **any HTTP method** (`GET`, `POST`, etc.), which makes it compatible
+    with protocols like MCP that stream SSE over `POST`.
+
+    The 

Test output

show
.................................                                        [100%]
=============================== warnings summary ===============================
../../../../../../../Users/jp/repos/kaggle-gemini-coding-agent-post-training/.envs/overlays/starlette-1.6.0-py3-none-any/starlette/testclient.py:53
  /Users/jp/repos/kaggle-gemini-coding-agent-post-training/.envs/overlays/starlette-1.6.0-py3-none-any/starlette/testclient.py:53: DeprecationWarning: The anyio.abc.BlockingPortal alias is deprecated, use anyio.from_thread.BlockingPortal instead.
    _PortalFactoryType = Callable[[], AbstractContextManager[anyio.abc.BlockingPortal]]

-- Docs: https://docs.pytest.org/en/stable/how-to/capture-warnings.html
33 passed, 1 warning in 1.63s