Skip to main content

openrouter/
stream.rs

1//! SSE parser + generic [`EventStream<T>`].
2//!
3//! The parser is hand-rolled over a byte stream supplied by reqwest on native
4//! targets and the browser's Fetch/ReadableStream APIs on WebAssembly. It
5//! handles `data:` accumulation, the `data: [DONE]` terminator, comment lines
6//! (`:` prefix), `\r\n` and bare `\r` line endings, and events split across
7//! chunk boundaries.
8//!
9//! Reconnection: when the underlying body stream errors on a transient
10//! failure, the stream re-opens via the caller-supplied closure with
11//! exponential backoff capped at [`MAX_RECONNECT_BACKOFF`]. Reconnects are
12//! counted across the lifetime of the stream, so an intermittent connection
13//! cannot exceed the configured request-replay budget. Non-transient errors
14//! and exhausted budget surface as `Err` and terminate the stream.
15//!
16//! Cancellation: dropping the `EventStream` drops the native response stream
17//! or aborts the browser Fetch request. No explicit `CancellationToken` is
18//! required — combine with `tokio::select!` on native targets or an abortable
19//! local future in the browser.
20
21use std::future::Future;
22use std::marker::PhantomData;
23use std::pin::Pin;
24use std::task::{Context, Poll};
25use std::time::Duration;
26
27use bytes::Bytes;
28use futures::{Stream, StreamExt};
29#[cfg(not(target_arch = "wasm32"))]
30use reqwest::Response;
31use serde::de::DeserializeOwned;
32
33#[cfg(target_arch = "wasm32")]
34use wasm_bindgen::JsCast;
35
36use crate::error::{Error, Result};
37use crate::retry::MAX_RECONNECT_BACKOFF;
38
39#[cfg(not(target_arch = "wasm32"))]
40pub(crate) type StreamResponse = Response;
41
42#[cfg(target_arch = "wasm32")]
43pub(crate) struct StreamResponse {
44    response: gloo_net::http::Response,
45    abort: web_sys::AbortController,
46}
47
48#[cfg(target_arch = "wasm32")]
49impl StreamResponse {
50    pub(crate) fn new(response: gloo_net::http::Response, abort: web_sys::AbortController) -> Self {
51        Self { response, abort }
52    }
53}
54
55/// Async factory that re-opens the underlying HTTP response after a
56/// transient failure. Returned by callers in `crate::client` so the stream
57/// can resume the same request body on reconnect.
58#[cfg(not(target_arch = "wasm32"))]
59pub(crate) type Reopen = std::sync::Arc<
60    dyn Fn() -> futures::future::BoxFuture<'static, Result<StreamResponse>> + Send + Sync + 'static,
61>;
62#[cfg(target_arch = "wasm32")]
63pub(crate) type Reopen = std::rc::Rc<
64    dyn Fn() -> futures::future::LocalBoxFuture<'static, Result<StreamResponse>> + 'static,
65>;
66
67#[cfg(not(target_arch = "wasm32"))]
68type ByteStream = futures::stream::BoxStream<'static, Result<Bytes>>;
69#[cfg(target_arch = "wasm32")]
70type ByteStream = futures::stream::LocalBoxStream<'static, Result<Bytes>>;
71
72#[cfg(not(target_arch = "wasm32"))]
73type SleepFuture = Pin<Box<dyn Future<Output = ()> + Send + 'static>>;
74#[cfg(target_arch = "wasm32")]
75type SleepFuture = Pin<Box<dyn Future<Output = ()> + 'static>>;
76
77#[cfg(not(target_arch = "wasm32"))]
78type ReopenFuture = futures::future::BoxFuture<'static, Result<StreamResponse>>;
79#[cfg(target_arch = "wasm32")]
80type ReopenFuture = futures::future::LocalBoxFuture<'static, Result<StreamResponse>>;
81
82/// A stream of deserialized SSE events.
83///
84/// Implements [`futures::Stream`] with `Item = Result<T>`. Yields `None` on
85/// the `data: [DONE]` terminator or when the underlying body finishes.
86pub struct EventStream<T: DeserializeOwned> {
87    state: State,
88    buf: SseBuffer,
89    reopen: Option<Reopen>,
90    reconnect_attempt: u32,
91    max_reconnects: u32,
92    _marker: PhantomData<fn() -> T>,
93}
94
95impl<T: DeserializeOwned> std::fmt::Debug for EventStream<T> {
96    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
97        f.debug_struct("EventStream")
98            .field("reconnect_attempt", &self.reconnect_attempt)
99            .field("max_reconnects", &self.max_reconnects)
100            .field("buffered", &self.buf.pending_len())
101            .field("state", &self.state.tag())
102            .finish()
103    }
104}
105
106enum State {
107    /// Active body stream — poll it for the next chunk.
108    Reading(ByteStream),
109    /// Sleeping before the next reconnect attempt.
110    Backoff(SleepFuture),
111    /// Re-opening the body via the caller's `reopen` closure.
112    Reopening(ReopenFuture),
113    /// Stream terminated (success or fatal error).
114    Done,
115}
116
117impl State {
118    fn tag(&self) -> &'static str {
119        match self {
120            State::Reading(_) => "reading",
121            State::Backoff(_) => "backoff",
122            State::Reopening(_) => "reopening",
123            State::Done => "done",
124        }
125    }
126}
127
128impl<T: DeserializeOwned> EventStream<T> {
129    /// Build a new event stream from an already-opened `Response` and a
130    /// reconnect closure. The closure is invoked on transient mid-stream
131    /// failures with exponential backoff.
132    #[allow(dead_code)] // Consumed by the streaming endpoints (HRA-123).
133    pub(crate) fn new(initial: StreamResponse, reopen: Reopen, max_reconnects: u32) -> Self {
134        Self {
135            state: State::Reading(box_byte_stream(initial)),
136            buf: SseBuffer::default(),
137            reopen: Some(reopen),
138            reconnect_attempt: 0,
139            max_reconnects,
140            _marker: PhantomData,
141        }
142    }
143
144    /// Build a non-reconnecting stream (used by tests and any caller that
145    /// doesn't want resume semantics).
146    #[cfg(test)]
147    pub(crate) fn from_bytes_stream(bytes: ByteStream) -> Self {
148        Self {
149            state: State::Reading(bytes),
150            buf: SseBuffer::default(),
151            reopen: None,
152            reconnect_attempt: 0,
153            max_reconnects: 0,
154            _marker: PhantomData,
155        }
156    }
157
158    fn reconnect_delay(&self) -> Duration {
159        // Exponential: 100ms, 200ms, 400ms, …, capped at MAX_RECONNECT_BACKOFF.
160        let base_ms = 100u64.saturating_mul(1u64 << self.reconnect_attempt.min(8));
161        let computed = Duration::from_millis(base_ms);
162        computed.min(MAX_RECONNECT_BACKOFF)
163    }
164}
165
166impl<T: DeserializeOwned> Stream for EventStream<T> {
167    type Item = Result<T>;
168
169    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
170        loop {
171            // First drain any complete events already buffered.
172            match self.buf.next_event() {
173                Some(SseEvent::Data(payload)) => {
174                    let decoded: Result<T> = serde_json::from_slice(&payload)
175                        .map_err(|e| Error::Stream(format!("malformed SSE payload: {e}")));
176                    return Poll::Ready(Some(decoded));
177                }
178                Some(SseEvent::Done) => {
179                    self.state = State::Done;
180                    return Poll::Ready(None);
181                }
182                None => {}
183            }
184
185            // Drive the state machine to produce more bytes/events.
186            // Take the current state out so we can replace it after polling.
187            let cur = std::mem::replace(&mut self.state, State::Done);
188            match cur {
189                State::Reading(mut s) => match s.poll_next_unpin(cx) {
190                    Poll::Ready(Some(Ok(chunk))) => {
191                        self.buf.push(&chunk);
192                        self.state = State::Reading(s);
193                        continue;
194                    }
195                    Poll::Ready(Some(Err(err))) => {
196                        if err.is_transient()
197                            && self.reopen.is_some()
198                            && self.reconnect_attempt < self.max_reconnects
199                        {
200                            // Schedule a reconnect.
201                            let delay = self.reconnect_delay();
202                            self.reconnect_attempt = self.reconnect_attempt.saturating_add(1);
203                            self.state = State::Backoff(Box::pin(crate::timer::sleep(delay)));
204                            continue;
205                        }
206                        self.state = State::Done;
207                        return Poll::Ready(Some(Err(err)));
208                    }
209                    Poll::Ready(None) => {
210                        // Body finished without [DONE]. Flush any final event,
211                        // then complete.
212                        if let Some(ev) = self.buf.finish() {
213                            self.state = State::Done;
214                            return match ev {
215                                SseEvent::Data(payload) => {
216                                    let decoded: Result<T> = serde_json::from_slice(&payload)
217                                        .map_err(|e| {
218                                            Error::Stream(format!("malformed SSE payload: {e}"))
219                                        });
220                                    Poll::Ready(Some(decoded))
221                                }
222                                SseEvent::Done => Poll::Ready(None),
223                            };
224                        }
225                        self.state = State::Done;
226                        return Poll::Ready(None);
227                    }
228                    Poll::Pending => {
229                        self.state = State::Reading(s);
230                        return Poll::Pending;
231                    }
232                },
233                State::Backoff(mut fut) => match fut.as_mut().poll(cx) {
234                    Poll::Ready(()) => {
235                        let reopen = self
236                            .reopen
237                            .clone()
238                            .expect("Backoff state requires a reopen closure");
239                        let f = (reopen)();
240                        self.state = State::Reopening(f);
241                        continue;
242                    }
243                    Poll::Pending => {
244                        self.state = State::Backoff(fut);
245                        return Poll::Pending;
246                    }
247                },
248                State::Reopening(mut fut) => match fut.as_mut().poll(cx) {
249                    Poll::Ready(Ok(resp)) => {
250                        self.state = State::Reading(box_byte_stream(resp));
251                        continue;
252                    }
253                    Poll::Ready(Err(err)) => {
254                        if err.is_transient() && self.reconnect_attempt < self.max_reconnects {
255                            // Try again, subject to the cap.
256                            let delay = self.reconnect_delay();
257                            self.reconnect_attempt = self.reconnect_attempt.saturating_add(1);
258                            self.state = State::Backoff(Box::pin(crate::timer::sleep(delay)));
259                            continue;
260                        }
261                        self.state = State::Done;
262                        return Poll::Ready(Some(Err(err)));
263                    }
264                    Poll::Pending => {
265                        self.state = State::Reopening(fut);
266                        return Poll::Pending;
267                    }
268                },
269                State::Done => {
270                    self.state = State::Done;
271                    return Poll::Ready(None);
272                }
273            }
274        }
275    }
276}
277
278#[cfg(not(target_arch = "wasm32"))]
279fn box_byte_stream(response: StreamResponse) -> ByteStream {
280    response
281        .bytes_stream()
282        .map(|result| result.map_err(Error::from))
283        .boxed()
284}
285
286#[cfg(target_arch = "wasm32")]
287fn box_byte_stream(response: StreamResponse) -> ByteStream {
288    let StreamResponse { response, abort } = response;
289    let Some(body) = response.body() else {
290        return futures::stream::empty().boxed_local();
291    };
292    let abort = AbortOnDrop(abort);
293    wasm_streams::ReadableStream::from_raw(body.unchecked_into())
294        .into_stream()
295        .map(move |item| {
296            let _abort = &abort;
297            let value = item.map_err(|error| Error::BrowserTransport(format!("{error:?}")))?;
298            let array = js_sys::Uint8Array::new(&value);
299            let mut bytes = vec![0; array.length() as usize];
300            array.copy_to(&mut bytes);
301            Ok(Bytes::from(bytes))
302        })
303        .boxed_local()
304}
305
306#[cfg(target_arch = "wasm32")]
307struct AbortOnDrop(web_sys::AbortController);
308
309#[cfg(target_arch = "wasm32")]
310impl Drop for AbortOnDrop {
311    fn drop(&mut self) {
312        self.0.abort();
313    }
314}
315
316/// A parsed SSE event.
317#[derive(Debug, PartialEq, Eq)]
318enum SseEvent {
319    /// Concatenated `data:` payload of a single event.
320    Data(Vec<u8>),
321    /// `data: [DONE]` terminator.
322    Done,
323}
324
325/// Incremental SSE parser. Accumulates bytes and yields complete events as
326/// `data:` lines are joined and blank lines flush them.
327#[derive(Default)]
328struct SseBuffer {
329    /// Raw bytes not yet consumed as full lines (no trailing `\n` seen).
330    pending: Vec<u8>,
331    /// Lines that belong to the in-progress event (each is one `data:` payload
332    /// without the `data:` prefix or trailing newline).
333    current_data: Vec<Vec<u8>>,
334    /// Whether the current event contained at least one `data:` line.
335    has_data: bool,
336}
337
338impl SseBuffer {
339    fn pending_len(&self) -> usize {
340        self.pending.len()
341    }
342
343    fn push(&mut self, chunk: &[u8]) {
344        self.pending.extend_from_slice(chunk);
345    }
346
347    /// Pop the next complete event from the buffer, if one is available.
348    fn next_event(&mut self) -> Option<SseEvent> {
349        loop {
350            let idx = self.pending.iter().position(|&b| b == b'\n')?;
351            // Take the line (excluding the `\n`); trim any trailing `\r`.
352            let mut line: Vec<u8> = self.pending.drain(..=idx).collect();
353            line.pop(); // remove the `\n`
354            if line.last() == Some(&b'\r') {
355                line.pop();
356            }
357            if let Some(ev) = self.process_line(line) {
358                return Some(ev);
359            }
360        }
361    }
362
363    /// Called when the upstream byte stream is exhausted. Flushes any
364    /// pending bytes as a final line.
365    fn finish(&mut self) -> Option<SseEvent> {
366        if !self.pending.is_empty() {
367            let mut line = std::mem::take(&mut self.pending);
368            if line.last() == Some(&b'\r') {
369                line.pop();
370            }
371            if let Some(ev) = self.process_line(line) {
372                return Some(ev);
373            }
374        }
375        self.flush_event()
376    }
377
378    fn process_line(&mut self, line: Vec<u8>) -> Option<SseEvent> {
379        if line.is_empty() {
380            return self.flush_event();
381        }
382        // Comment line.
383        if line.first() == Some(&b':') {
384            return None;
385        }
386        // Strip the field name. SSE allows arbitrary fields; we only care about `data`.
387        if let Some(rest) = strip_field(&line, b"data") {
388            self.current_data.push(rest);
389            self.has_data = true;
390        }
391        // Other fields (event:, id:, retry:) are intentionally ignored — the
392        // OpenRouter SSE stream uses only `data:`.
393        None
394    }
395
396    fn flush_event(&mut self) -> Option<SseEvent> {
397        if !self.has_data {
398            return None;
399        }
400        self.has_data = false;
401        let lines = std::mem::take(&mut self.current_data);
402        // Per the SSE spec, multi-line `data:` payloads are joined by `\n`.
403        let mut payload: Vec<u8> = Vec::new();
404        for (i, l) in lines.iter().enumerate() {
405            if i > 0 {
406                payload.push(b'\n');
407            }
408            payload.extend_from_slice(l);
409        }
410        // `[DONE]` terminator is treated specially.
411        if payload == b"[DONE]" {
412            return Some(SseEvent::Done);
413        }
414        Some(SseEvent::Data(payload))
415    }
416}
417
418/// If `line` starts with `field:`, return the value (with at most one leading
419/// space trimmed, per the SSE spec).
420fn strip_field(line: &[u8], field: &[u8]) -> Option<Vec<u8>> {
421    if line.len() < field.len() + 1 {
422        return None;
423    }
424    if &line[..field.len()] != field {
425        return None;
426    }
427    if line[field.len()] != b':' {
428        return None;
429    }
430    let mut rest = &line[field.len() + 1..];
431    if rest.first() == Some(&b' ') {
432        rest = &rest[1..];
433    }
434    Some(rest.to_vec())
435}
436
437#[cfg(test)]
438mod tests {
439    use super::*;
440    use futures::stream;
441    use pretty_assertions::assert_eq;
442    use std::sync::atomic::{AtomicUsize, Ordering};
443    use std::sync::Arc;
444
445    fn drain_buffer(buf: &mut SseBuffer) -> Vec<SseEvent> {
446        let mut out = Vec::new();
447        while let Some(ev) = buf.next_event() {
448            out.push(ev);
449        }
450        if let Some(ev) = buf.finish() {
451            out.push(ev);
452        }
453        out
454    }
455
456    #[test]
457    fn parses_single_event() {
458        let mut b = SseBuffer::default();
459        b.push(b"data: {\"x\":1}\n\n");
460        let events = drain_buffer(&mut b);
461        assert_eq!(events, vec![SseEvent::Data(b"{\"x\":1}".to_vec())]);
462    }
463
464    #[test]
465    fn parses_done_terminator() {
466        let mut b = SseBuffer::default();
467        b.push(b"data: [DONE]\n\n");
468        let events = drain_buffer(&mut b);
469        assert_eq!(events, vec![SseEvent::Done]);
470    }
471
472    #[test]
473    fn ignores_comment_lines() {
474        let mut b = SseBuffer::default();
475        b.push(b": heartbeat\ndata: {\"a\":1}\n\n");
476        let events = drain_buffer(&mut b);
477        assert_eq!(events, vec![SseEvent::Data(b"{\"a\":1}".to_vec())]);
478    }
479
480    #[test]
481    fn joins_multi_line_data() {
482        let mut b = SseBuffer::default();
483        b.push(b"data: line1\ndata: line2\n\n");
484        let events = drain_buffer(&mut b);
485        assert_eq!(events, vec![SseEvent::Data(b"line1\nline2".to_vec())]);
486    }
487
488    #[test]
489    fn handles_crlf_line_endings() {
490        let mut b = SseBuffer::default();
491        b.push(b"data: {\"x\":1}\r\n\r\n");
492        let events = drain_buffer(&mut b);
493        assert_eq!(events, vec![SseEvent::Data(b"{\"x\":1}".to_vec())]);
494    }
495
496    #[test]
497    fn handles_chunk_boundaries() {
498        let mut b = SseBuffer::default();
499        b.push(b"data: {\"x");
500        assert!(b.next_event().is_none());
501        b.push(b"\":1}\n");
502        // No terminating blank line yet — event not flushed.
503        assert!(b.next_event().is_none());
504        b.push(b"\n");
505        let events = drain_buffer(&mut b);
506        assert_eq!(events, vec![SseEvent::Data(b"{\"x\":1}".to_vec())]);
507    }
508
509    #[test]
510    fn ignores_non_data_fields() {
511        let mut b = SseBuffer::default();
512        b.push(b"event: ping\nid: 42\nretry: 1000\ndata: {\"x\":1}\n\n");
513        let events = drain_buffer(&mut b);
514        assert_eq!(events, vec![SseEvent::Data(b"{\"x\":1}".to_vec())]);
515    }
516
517    #[test]
518    fn flushes_trailing_event_without_blank_line() {
519        let mut b = SseBuffer::default();
520        b.push(b"data: {\"x\":1}\n");
521        // No second \n; finish() flushes.
522        let events = drain_buffer(&mut b);
523        assert_eq!(events, vec![SseEvent::Data(b"{\"x\":1}".to_vec())]);
524    }
525
526    #[test]
527    fn handles_empty_data_payload() {
528        let mut b = SseBuffer::default();
529        b.push(b"data: \n\n");
530        let events = drain_buffer(&mut b);
531        assert_eq!(events, vec![SseEvent::Data(Vec::new())]);
532    }
533
534    #[derive(serde::Deserialize, Debug, PartialEq)]
535    struct Sample {
536        x: i32,
537    }
538
539    #[tokio::test]
540    async fn event_stream_yields_decoded_events_then_done() {
541        let chunks: Vec<Result<Bytes>> = vec![
542            Ok(Bytes::from_static(b"data: {\"x\":1}\n\n")),
543            Ok(Bytes::from_static(b"data: {\"x\":2}\n\n")),
544            Ok(Bytes::from_static(b"data: [DONE]\n\n")),
545        ];
546        let body: ByteStream = stream::iter(chunks).boxed();
547        let mut s: EventStream<Sample> = EventStream::from_bytes_stream(body);
548        let a = s.next().await.unwrap().unwrap();
549        let b = s.next().await.unwrap().unwrap();
550        assert_eq!(a, Sample { x: 1 });
551        assert_eq!(b, Sample { x: 2 });
552        assert!(s.next().await.is_none());
553    }
554
555    #[tokio::test]
556    async fn event_stream_surfaces_malformed_payload_as_error() {
557        let chunks: Vec<Result<Bytes>> = vec![Ok(Bytes::from_static(b"data: not-json\n\n"))];
558        let body: ByteStream = stream::iter(chunks).boxed();
559        let mut s: EventStream<Sample> = EventStream::from_bytes_stream(body);
560        let item = s.next().await.unwrap();
561        assert!(matches!(item, Err(Error::Stream(_))));
562    }
563
564    #[tokio::test]
565    async fn event_stream_handles_split_event_across_chunks() {
566        let chunks: Vec<Result<Bytes>> = vec![
567            Ok(Bytes::from_static(b"data: {\"x")),
568            Ok(Bytes::from_static(b"\":7}\n\n")),
569            Ok(Bytes::from_static(b"data: [DONE]\n\n")),
570        ];
571        let body: ByteStream = stream::iter(chunks).boxed();
572        let mut s: EventStream<Sample> = EventStream::from_bytes_stream(body);
573        let a = s.next().await.unwrap().unwrap();
574        assert_eq!(a, Sample { x: 7 });
575        assert!(s.next().await.is_none());
576    }
577
578    #[tokio::test(start_paused = true)]
579    async fn reconnect_budget_is_lifetime_bounded() {
580        let body: ByteStream =
581            stream::iter(vec![Err(Error::BrowserTransport("lost".into()))]).boxed();
582        let calls = Arc::new(AtomicUsize::new(0));
583        let reopen_calls = Arc::clone(&calls);
584        let reopen: Reopen = Arc::new(move || {
585            reopen_calls.fetch_add(1, Ordering::SeqCst);
586            Box::pin(async { Err(Error::BrowserTransport("still lost".into())) })
587        });
588        let mut stream: EventStream<Sample> = EventStream {
589            state: State::Reading(body),
590            buf: SseBuffer::default(),
591            reopen: Some(reopen),
592            reconnect_attempt: 0,
593            max_reconnects: 2,
594            _marker: PhantomData,
595        };
596
597        let error = stream.next().await.unwrap().unwrap_err();
598        assert!(matches!(error, Error::BrowserTransport(_)));
599        assert_eq!(calls.load(Ordering::SeqCst), 2);
600        assert_eq!(stream.reconnect_attempt, 2);
601    }
602}