toolkit_contract/runtime/sse.rs
1//! Server-Sent Events (SSE) parser used by streaming clients.
2//!
3//! Translates a byte stream into a stream of typed events. Recognises:
4//! - `data: <json>` — emits `Ok(T)` after JSON-deserializing into `T`.
5//! - `event: error` — the next `data:` is parsed as a `ProblemDetails`
6//! wrapped in [`TransportError::Problem`].
7//! - `event: done` — terminates the stream.
8//!
9//! All other event types are ignored. Comments (lines starting with `:`) and
10//! blank lines are stripped per the SSE spec.
11//!
12//! Accumulated per-line and per-event buffers are bounded by
13//! [`MAX_ACCUMULATED_BYTES`] to protect against a peer that streams an
14//! unbounded line (no terminating `\n`) or an unbounded run of `data:` lines
15//! with no dispatching blank line — otherwise the buffer would grow without
16//! limit for the lifetime of a self-healing, indefinitely-reconnecting client.
17
18use std::collections::VecDeque;
19use std::pin::Pin;
20use std::sync::Arc;
21use std::sync::atomic::{AtomicU64, Ordering};
22use std::task::{Context, Poll};
23
24use bytes::{Bytes, BytesMut};
25use futures_core::Stream;
26use parking_lot::RwLock;
27use serde::de::DeserializeOwned;
28
29use toolkit_canonical_errors::Problem;
30
31use crate::runtime::transport_error::TransportError;
32
33/// Adapter that lifts a `Display`-only error into an `Error + Send + Sync + 'static`
34/// so it can be boxed into [`TransportError::Network`] without losing the
35/// original message.
36#[derive(Debug)]
37struct DisplayError(String);
38impl std::fmt::Display for DisplayError {
39 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
40 f.write_str(&self.0)
41 }
42}
43impl std::error::Error for DisplayError {}
44
45/// Shared cell holding the latest seen SSE `id:` field. The streaming
46/// client clones this handle before constructing the parser; on stream
47/// interruption it reads the latest ID and re-issues the request with a
48/// `Last-Event-ID` header (per HTML5 `EventSource` spec).
49///
50/// Wraps `Arc<RwLock<Option<String>>>` as a newtype so the underlying lock
51/// implementation isn't part of the public surface — the `parking_lot` vs
52/// `tokio::sync` choice can change without a breaking SDK release.
53#[derive(Clone, Debug, Default)]
54pub struct LastEventId(Arc<RwLock<Option<String>>>);
55
56impl LastEventId {
57 /// Create an empty [`LastEventId`] cell.
58 #[must_use]
59 pub fn empty() -> Self {
60 Self::default()
61 }
62
63 /// Snapshot the latest ID value, if any.
64 #[must_use]
65 pub fn current(&self) -> Option<String> {
66 self.0.read().clone()
67 }
68
69 /// Replace the latest ID. `None` clears the cell (per HTML5 spec, an
70 /// empty `id:` field resets the saved value).
71 pub fn set(&self, value: Option<String>) {
72 *self.0.write() = value;
73 }
74}
75
76/// Maximum bytes the parser accumulates for a not-yet-terminated line ([`SseStream::buf`])
77/// or for a not-yet-dispatched event's `data:` payload ([`SseStream::event_data`])
78/// before treating the peer as protocol-violating and terminating the stream
79/// with [`TransportError::sse`]. Generous enough for any realistic single
80/// event; guards against unbounded memory growth from a misbehaving peer.
81const MAX_ACCUMULATED_BYTES: usize = 16 * 1024 * 1024;
82
83/// Shared monotonic counter bumped every time the parser receives a byte chunk
84/// from the wire — **including** chunks that carry only keepalive comments or
85/// other non-dispatching frames. Streaming clients snapshot it around an idle
86/// wait so a low-data-rate stream kept alive purely by `:keepalive` comments is
87/// recognised as *active* rather than *idle* (the idle timeout must be idle).
88///
89/// Wraps `Arc<AtomicU64>` as a newtype so the atomic choice isn't part of the
90/// public surface.
91#[derive(Clone, Debug, Default)]
92pub struct SseActivity(Arc<AtomicU64>);
93
94impl SseActivity {
95 /// Create a fresh activity counter at generation 0.
96 #[must_use]
97 pub fn new() -> Self {
98 Self::default()
99 }
100
101 /// Current activity generation. Compare two snapshots to detect whether any
102 /// wire chunk arrived in between.
103 #[must_use]
104 pub fn generation(&self) -> u64 {
105 self.0.load(Ordering::Relaxed)
106 }
107
108 /// Record that a wire chunk arrived.
109 fn bump(&self) {
110 self.0.fetch_add(1, Ordering::Relaxed);
111 }
112}
113
114/// Parse an SSE byte stream into a stream of typed events.
115///
116/// `bytes` is typically the byte-stream view of
117/// `toolkit_http::HttpResponse::into_body()` (adapted via
118/// [`crate::runtime::http::body_to_byte_stream`]). Errors from the inner
119/// stream are surfaced as [`TransportError::Network`].
120///
121/// To capture `id:` fields for `Last-Event-ID` reconnect, use
122/// [`parse_sse_stream_with_id`] and pass in a shared cell that the
123/// streaming client can read from.
124pub fn parse_sse_stream<T, S, E>(bytes: S) -> SseStream<T, S>
125where
126 T: DeserializeOwned + 'static,
127 S: Stream<Item = Result<Bytes, E>> + Unpin + 'static,
128 E: std::fmt::Display,
129{
130 parse_sse_stream_with_id(bytes, LastEventId::empty())
131}
132
133/// Same as [`parse_sse_stream`] but accepts a [`LastEventId`] cell that the
134/// parser updates whenever it encounters an `id:` field. Streaming clients
135/// hand the cell into the request-factory closure on reconnect to populate
136/// the `Last-Event-ID` header — per HTML5 `EventSource` spec.
137pub fn parse_sse_stream_with_id<T, S, E>(bytes: S, last_event_id: LastEventId) -> SseStream<T, S>
138where
139 T: DeserializeOwned + 'static,
140 S: Stream<Item = Result<Bytes, E>> + Unpin + 'static,
141 E: std::fmt::Display,
142{
143 SseStream {
144 inner: bytes,
145 buf: BytesMut::with_capacity(4 * 1024),
146 scan_from: 0,
147 pending: VecDeque::new(),
148 event_kind: None,
149 event_data: String::new(),
150 event_id: None,
151 done: false,
152 explicit_done: false,
153 last_event_id,
154 activity: SseActivity::new(),
155 _marker: std::marker::PhantomData,
156 }
157}
158
159/// Iterator yielded by [`parse_sse_stream`].
160pub struct SseStream<T, S> {
161 inner: S,
162 buf: BytesMut,
163 /// Prefix length of `buf`, in bytes, already confirmed to contain no
164 /// unconsumed `\n`. Lets [`find_line_end`] resume scanning from here
165 /// instead of rescanning the whole buffer on every poll — without this, a
166 /// single long unterminated line delivered over many small chunks costs
167 /// `O(total_length²)` instead of `O(total_length)`.
168 scan_from: usize,
169 pending: VecDeque<Result<T, TransportError>>,
170 /// Last `event:` value seen since the previous dispatch. `None` means
171 /// the implicit default `"message"`.
172 event_kind: Option<String>,
173 /// `data:` payload accumulated for the current event, with multiple
174 /// `data:` lines joined by `\n` (per W3C SSE spec).
175 event_data: String,
176 /// Last `id:` value seen for the current event. Per spec, the
177 /// last-event-id persists across dispatches; this field is just the
178 /// per-event scratch used to update [`LastEventId`] on dispatch.
179 event_id: Option<String>,
180 done: bool,
181 /// `true` only when the stream ended because an `event: done` frame was
182 /// actually dispatched — as opposed to `done` being set because the
183 /// underlying byte stream simply closed (peer disconnect, proxy timeout,
184 /// etc.). The streaming client uses this to tell "the peer said it's
185 /// finished" apart from "the connection just ended", so it can treat the
186 /// latter as reconnect-eligible instead of a silent success.
187 explicit_done: bool,
188 last_event_id: LastEventId,
189 activity: SseActivity,
190 _marker: std::marker::PhantomData<fn() -> T>,
191}
192
193impl<T, S> SseStream<T, S> {
194 /// Returns a clone of the shared cell that captures the latest `id:`
195 /// field seen on the stream. The streaming client uses this to populate
196 /// the `Last-Event-ID` header on reconnect.
197 #[must_use]
198 pub fn last_event_id_handle(&self) -> LastEventId {
199 self.last_event_id.clone()
200 }
201
202 /// `true` iff the stream ended because the peer explicitly dispatched an
203 /// `event: done` frame. `false` if the stream ended for any other reason
204 /// (byte stream closed, error) — including while still in progress, so
205 /// callers should only consult this after the stream has yielded `None`.
206 #[must_use]
207 pub fn saw_done_event(&self) -> bool {
208 self.explicit_done
209 }
210
211 /// Returns a clone of the shared wire-activity counter. The streaming client
212 /// snapshots it around an idle-timeout wait so keepalive-only traffic keeps
213 /// the stream alive (see [`SseActivity`]).
214 #[must_use]
215 pub fn activity_handle(&self) -> SseActivity {
216 self.activity.clone()
217 }
218}
219
220// `inner` is bounded by `Unpin` at construction; the rest of the fields are
221// trivially `Unpin`. Implement `Unpin` unconditionally so callers can poll
222// `Pin<&mut SseStream<...>>` without pinning the type itself.
223impl<T, S: Unpin> Unpin for SseStream<T, S> {}
224
225impl<T, S, E> Stream for SseStream<T, S>
226where
227 T: DeserializeOwned + 'static,
228 S: Stream<Item = Result<Bytes, E>> + Unpin + 'static,
229 E: std::fmt::Display,
230{
231 type Item = Result<T, TransportError>;
232
233 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
234 let this = self.get_mut();
235
236 loop {
237 if let Some(item) = this.pending.pop_front() {
238 return Poll::Ready(Some(item));
239 }
240 if this.done {
241 return Poll::Ready(None);
242 }
243
244 match Pin::new(&mut this.inner).poll_next(cx) {
245 Poll::Pending => return Poll::Pending,
246 Poll::Ready(None) => {
247 this.done = true;
248 // Flush any trailing partial buffer as a final line.
249 drain_remaining(
250 &mut this.buf,
251 &mut this.event_kind,
252 &mut this.event_data,
253 &mut this.event_id,
254 &mut this.pending,
255 &this.last_event_id,
256 );
257 }
258 Poll::Ready(Some(Err(e))) => {
259 this.done = true;
260 // `E: Display` only — wrap in a small Display->Error
261 // adapter so the source chain stays intact through
262 // `TransportError::network`.
263 return Poll::Ready(Some(Err(TransportError::network(DisplayError(
264 e.to_string(),
265 )))));
266 }
267 Poll::Ready(Some(Ok(chunk))) => {
268 // Any wire chunk counts as activity — even one carrying only
269 // keepalive comments that dispatch no item — so the idle
270 // timeout in the streaming driver can tell "quiet but alive"
271 // from "truly idle".
272 this.activity.bump();
273 this.buf.extend_from_slice(&chunk);
274 let saw_done = drain_buffer(
275 &mut this.buf,
276 &mut this.scan_from,
277 &mut this.event_kind,
278 &mut this.event_data,
279 &mut this.event_id,
280 &mut this.pending,
281 &this.last_event_id,
282 );
283 if saw_done {
284 this.done = true;
285 this.explicit_done = true;
286 }
287 // Guard against unbounded memory growth: a peer streaming a
288 // line with no terminating `\n`, or a run of `data:` lines
289 // with no dispatching blank line, would otherwise grow
290 // `buf`/`event_data` without limit for the life of a
291 // self-healing, indefinitely-reconnecting client.
292 if this.buf.len() > MAX_ACCUMULATED_BYTES
293 || this.event_data.len() > MAX_ACCUMULATED_BYTES
294 {
295 this.done = true;
296 return Poll::Ready(Some(Err(TransportError::sse(
297 "SSE frame exceeds maximum accumulated size; aborting stream",
298 ))));
299 }
300 }
301 }
302 }
303 }
304}
305
306/// Drain complete lines from `buf`, mutating the per-event accumulator
307/// fields and pushing dispatched events to `out`. Returns `true` iff a
308/// `done` event was dispatched (caller terminates the stream).
309fn drain_buffer<T: DeserializeOwned + 'static>(
310 buf: &mut BytesMut,
311 scan_from: &mut usize,
312 event_kind: &mut Option<String>,
313 event_data: &mut String,
314 event_id: &mut Option<String>,
315 out: &mut VecDeque<Result<T, TransportError>>,
316 last_event_id: &LastEventId,
317) -> bool {
318 let mut saw_done = false;
319 while let Some(line_end) = find_line_end(buf, *scan_from) {
320 let line_bytes = buf.split_to(line_end.consumed);
321 // Bytes after the found `\n` were never examined by this call (the
322 // scan stops at the first match) — resume from scratch for them.
323 *scan_from = 0;
324 // SSE wire format mandates UTF-8 (RFC 8259 § 8.1, EventSource spec).
325 // Surface non-conforming server output as a typed transport error
326 // instead of silently dropping the line — invisible data loss is
327 // worse than a propagated error.
328 match std::str::from_utf8(&line_bytes[..line_end.line_len]) {
329 Ok(line) => {
330 if process_line(line, event_kind, event_data, event_id, out, last_event_id) {
331 saw_done = true;
332 }
333 }
334 Err(e) => {
335 out.push_back(Err(TransportError::sse(format!(
336 "invalid UTF-8 in SSE frame: {e}"
337 ))));
338 }
339 }
340 }
341 // Nothing left to find: everything currently in `buf` is confirmed
342 // newline-free. Remember that so the next poll's `extend_from_slice` only
343 // needs `find_line_end` to scan the newly appended tail.
344 *scan_from = buf.len();
345 saw_done
346}
347
348/// Flush any trailing bytes (without a final `\n`) as one last line, then
349/// — since end-of-stream implies an event boundary — dispatch any
350/// accumulated event.
351fn drain_remaining<T: DeserializeOwned + 'static>(
352 buf: &mut BytesMut,
353 event_kind: &mut Option<String>,
354 event_data: &mut String,
355 _event_id: &mut Option<String>,
356 _out: &mut VecDeque<Result<T, TransportError>>,
357 _last_event_id: &LastEventId,
358) {
359 // Per the W3C EventSource spec, an event that is not terminated by a blank
360 // line before end-of-stream is INCOMPLETE and MUST be discarded. Forcing a
361 // dispatch here would surface a truncated final frame as a (non-transient)
362 // `Serialization` error — both a spec violation and a defeat of reconnect.
363 // A properly framed final event was already dispatched on its blank line,
364 // so anything left in the buffers is a partial event: drop it.
365 buf.clear();
366 event_kind.take();
367 event_data.clear();
368 // Note: `scan_from` is intentionally left untouched here — the stream is
369 // marked `done` by the caller right after this returns, so it is never
370 // consulted again.
371}
372
373/// Process a single (trailing-CR/LF-stripped) SSE line. Returns `true` iff
374/// the line caused a `done` event to be dispatched.
375fn process_line<T: DeserializeOwned + 'static>(
376 raw: &str,
377 event_kind: &mut Option<String>,
378 event_data: &mut String,
379 event_id: &mut Option<String>,
380 out: &mut VecDeque<Result<T, TransportError>>,
381 last_event_id: &LastEventId,
382) -> bool {
383 let line = raw.trim_end_matches(['\r', '\n']);
384
385 // Blank line — dispatch boundary.
386 if line.is_empty() {
387 // Per spec, suppress dispatch when no fields were set since the
388 // last dispatch (e.g. stray blank lines / keepalives).
389 if event_data.is_empty() && event_kind.is_none() {
390 return false;
391 }
392 return dispatch_event(event_kind, event_data, event_id, out, last_event_id);
393 }
394
395 // Comment — ignore.
396 if line.starts_with(':') {
397 return false;
398 }
399
400 if let Some(value) = line.strip_prefix("event:") {
401 *event_kind = Some(value.trim().to_owned());
402 return false;
403 }
404
405 if let Some(value) = line.strip_prefix("data:") {
406 let payload = value.strip_prefix(' ').unwrap_or(value);
407 if !event_data.is_empty() {
408 event_data.push('\n');
409 }
410 event_data.push_str(payload);
411 return false;
412 }
413
414 // SSE `id:` field — capture per event; also propagate to the
415 // connection-level last-event-id (per HTML5 EventSource spec). Empty
416 // `id:` clears the saved value.
417 if let Some(value) = line.strip_prefix("id:") {
418 let id = value.trim().to_owned();
419 if id.is_empty() {
420 *event_id = None;
421 last_event_id.set(None);
422 } else {
423 *event_id = Some(id.clone());
424 last_event_id.set(Some(id));
425 }
426 return false;
427 }
428
429 // `retry:` and other unknown fields — ignore per spec.
430 false
431}
432
433/// Drain the accumulated per-event state into `out`. Returns `true` iff
434/// the dispatched event was a `done` sentinel.
435fn dispatch_event<T: DeserializeOwned + 'static>(
436 event_kind: &mut Option<String>,
437 event_data: &mut String,
438 event_id: &mut Option<String>,
439 out: &mut VecDeque<Result<T, TransportError>>,
440 _last_event_id: &LastEventId,
441) -> bool {
442 let kind = event_kind.take().unwrap_or_else(|| "message".to_owned());
443 let payload = std::mem::take(event_data);
444 // Per spec, last-event-id persists across events — do NOT clear
445 // `event_id` here. The connection-level `LastEventId` cell was already
446 // updated when the `id:` line was parsed.
447 let _ = event_id;
448
449 match kind.as_str() {
450 "done" => true,
451 "error" => {
452 out.push_back(Err(parse_problem(&payload)));
453 false
454 }
455 // The implicit default channel carries the typed payload. An event with
456 // no `event:` field defaults to `"message"` here.
457 "message" => {
458 match serde_json::from_str::<T>(&payload) {
459 Ok(v) => out.push_back(Ok(v)),
460 Err(e) => out.push_back(Err(TransportError::serialization(e))),
461 }
462 false
463 }
464 // Named control events (`heartbeat`, `ping`, custom kinds) are not the
465 // typed data channel — ignore them rather than trying to decode `T`
466 // (which would yield spurious items or serialization errors).
467 _ => false,
468 }
469}
470
471fn parse_problem(payload: &str) -> TransportError {
472 match serde_json::from_str::<Problem>(payload) {
473 Ok(p) => TransportError::problem(p),
474 Err(e) => TransportError::sse(format!("malformed error event: {e}")),
475 }
476}
477
478struct LineEnd {
479 consumed: usize,
480 line_len: usize,
481}
482
483/// Scan for the next `\n`, starting from byte offset `start` (bytes before
484/// `start` are assumed already confirmed newline-free by the caller — see
485/// [`SseStream::scan_from`]).
486fn find_line_end(buf: &[u8], start: usize) -> Option<LineEnd> {
487 for (i, b) in buf[start..].iter().enumerate() {
488 if *b == b'\n' {
489 let abs = start + i;
490 return Some(LineEnd {
491 consumed: abs + 1,
492 line_len: abs,
493 });
494 }
495 }
496 None
497}
498
499#[cfg(test)]
500#[cfg_attr(coverage_nightly, coverage(off))]
501#[allow(clippy::unwrap_used)]
502#[path = "sse_tests.rs"]
503mod tests;