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;