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