Skip to main content

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;