sse.py 6.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241
  1. from typing import Annotated, Any
  2. from annotated_doc import Doc
  3. from pydantic import AfterValidator, BaseModel, Field, model_validator
  4. from starlette.responses import StreamingResponse
  5. # Canonical SSE event schema matching the OpenAPI 3.2 spec
  6. # (Section 4.14.4 "Special Considerations for Server-Sent Events")
  7. _SSE_EVENT_SCHEMA: dict[str, Any] = {
  8. "type": "object",
  9. "properties": {
  10. "data": {"type": "string"},
  11. "event": {"type": "string"},
  12. "id": {"type": "string"},
  13. "retry": {"type": "integer", "minimum": 0},
  14. },
  15. }
  16. class EventSourceResponse(StreamingResponse):
  17. """Streaming response with `text/event-stream` media type.
  18. Use as `response_class=EventSourceResponse` on a *path operation* that uses `yield`
  19. to enable Server Sent Events (SSE) responses.
  20. Works with **any HTTP method** (`GET`, `POST`, etc.), which makes it compatible
  21. with protocols like MCP that stream SSE over `POST`.
  22. The actual encoding logic lives in the FastAPI routing layer. This class
  23. serves mainly as a marker and sets the correct `Content-Type`.
  24. """
  25. media_type = "text/event-stream"
  26. def _check_single_line(v: str | None, field_name: str) -> str | None:
  27. if v is not None and ("\r" in v or "\n" in v):
  28. raise ValueError(f"SSE '{field_name}' must be a single line")
  29. return v
  30. def _check_event_single_line(v: str | None) -> str | None:
  31. return _check_single_line(v, "event")
  32. def _check_id_valid(v: str | None) -> str | None:
  33. if v is not None and "\0" in v:
  34. raise ValueError("SSE 'id' must not contain null characters")
  35. return _check_single_line(v, "id")
  36. class ServerSentEvent(BaseModel):
  37. """Represents a single Server-Sent Event.
  38. When `yield`ed from a *path operation function* that uses
  39. `response_class=EventSourceResponse`, each `ServerSentEvent` is encoded
  40. into the [SSE wire format](https://html.spec.whatwg.org/multipage/server-sent-events.html#parsing-an-event-stream)
  41. (`text/event-stream`).
  42. If you yield a plain object (dict, Pydantic model, etc.) instead, it is
  43. automatically JSON-encoded and sent as the `data:` field.
  44. All `data` values **including plain strings** are JSON-serialized.
  45. For example, `data="hello"` produces `data: "hello"` on the wire (with
  46. quotes).
  47. """
  48. data: Annotated[
  49. Any,
  50. Doc(
  51. """
  52. The event payload.
  53. Can be any JSON-serializable value: a Pydantic model, dict, list,
  54. string, number, etc. It is **always** serialized to JSON: strings
  55. are quoted (`"hello"` becomes `data: "hello"` on the wire).
  56. Mutually exclusive with `raw_data`.
  57. """
  58. ),
  59. ] = None
  60. raw_data: Annotated[
  61. str | None,
  62. Doc(
  63. """
  64. Raw string to send as the `data:` field **without** JSON encoding.
  65. Use this when you need to send pre-formatted text, HTML fragments,
  66. CSV lines, or any non-JSON payload. The string is placed directly
  67. into the `data:` field as-is.
  68. Mutually exclusive with `data`.
  69. """
  70. ),
  71. ] = None
  72. event: Annotated[
  73. str | None,
  74. AfterValidator(_check_event_single_line),
  75. Doc(
  76. """
  77. Optional event type name.
  78. Maps to `addEventListener(event, ...)` on the browser. When omitted,
  79. the browser dispatches on the generic `message` event. Must be a
  80. single line.
  81. """
  82. ),
  83. ] = None
  84. id: Annotated[
  85. str | None,
  86. AfterValidator(_check_id_valid),
  87. Doc(
  88. """
  89. Optional event ID.
  90. The browser sends this value back as the `Last-Event-ID` header on
  91. automatic reconnection. **Must be a single line** and must not contain
  92. null (`\\0`) characters.
  93. """
  94. ),
  95. ] = None
  96. retry: Annotated[
  97. int | None,
  98. Field(ge=0),
  99. Doc(
  100. """
  101. Optional reconnection time in **milliseconds**.
  102. Tells the browser how long to wait before reconnecting after the
  103. connection is lost. Must be a non-negative integer.
  104. """
  105. ),
  106. ] = None
  107. comment: Annotated[
  108. str | None,
  109. Doc(
  110. """
  111. Optional comment line(s).
  112. Comment lines start with `:` in the SSE wire format and are ignored by
  113. `EventSource` clients. Useful for keep-alive pings to prevent
  114. proxy/load-balancer timeouts.
  115. """
  116. ),
  117. ] = None
  118. @model_validator(mode="after")
  119. def _check_data_exclusive(self) -> "ServerSentEvent":
  120. if self.data is not None and self.raw_data is not None:
  121. raise ValueError(
  122. "Cannot set both 'data' and 'raw_data' on the same "
  123. "ServerSentEvent. Use 'data' for JSON-serialized payloads "
  124. "or 'raw_data' for pre-formatted strings."
  125. )
  126. return self
  127. def _split_sse_lines(value: str) -> list[str]:
  128. # Split on SSE-spec line terminators only (\n, \r\n, \r), preserving
  129. # trailing empty strings.
  130. return value.replace("\r\n", "\n").replace("\r", "\n").split("\n")
  131. def format_sse_event(
  132. *,
  133. data_str: Annotated[
  134. str | None,
  135. Doc(
  136. """
  137. Pre-serialized data string to use as the `data:` field.
  138. """
  139. ),
  140. ] = None,
  141. event: Annotated[
  142. str | None,
  143. Doc(
  144. """
  145. Optional event type name (`event:` field).
  146. """
  147. ),
  148. ] = None,
  149. id: Annotated[
  150. str | None,
  151. Doc(
  152. """
  153. Optional event ID (`id:` field).
  154. """
  155. ),
  156. ] = None,
  157. retry: Annotated[
  158. int | None,
  159. Doc(
  160. """
  161. Optional reconnection time in milliseconds (`retry:` field).
  162. """
  163. ),
  164. ] = None,
  165. comment: Annotated[
  166. str | None,
  167. Doc(
  168. """
  169. Optional comment line(s) (`:` prefix).
  170. """
  171. ),
  172. ] = None,
  173. ) -> bytes:
  174. """Build SSE wire-format bytes from **pre-serialized** data.
  175. The result always ends with `\\n\\n` (the event terminator).
  176. """
  177. lines: list[str] = []
  178. if comment is not None:
  179. for line in _split_sse_lines(comment):
  180. lines.append(f": {line}")
  181. if event is not None:
  182. lines.append(f"event: {event}")
  183. if data_str is not None:
  184. for line in _split_sse_lines(data_str):
  185. lines.append(f"data: {line}")
  186. if id is not None:
  187. lines.append(f"id: {id}")
  188. if retry is not None:
  189. lines.append(f"retry: {retry}")
  190. lines.append("")
  191. lines.append("")
  192. return "\n".join(lines).encode("utf-8")
  193. # Keep-alive comment, per the SSE spec recommendation
  194. KEEPALIVE_COMMENT = b": ping\n\n"
  195. # Seconds between keep-alive pings when a generator is idle.
  196. # Private but importable so tests can monkeypatch it.
  197. _PING_INTERVAL: float = 15.0