1use 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#[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
82pub 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 Reading(ByteStream),
109 Backoff(SleepFuture),
111 Reopening(ReopenFuture),
113 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 #[allow(dead_code)] 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 #[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 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 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 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 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 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 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#[derive(Debug, PartialEq, Eq)]
318enum SseEvent {
319 Data(Vec<u8>),
321 Done,
323}
324
325#[derive(Default)]
328struct SseBuffer {
329 pending: Vec<u8>,
331 current_data: Vec<Vec<u8>>,
334 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 fn next_event(&mut self) -> Option<SseEvent> {
349 loop {
350 let idx = self.pending.iter().position(|&b| b == b'\n')?;
351 let mut line: Vec<u8> = self.pending.drain(..=idx).collect();
353 line.pop(); 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 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 if line.first() == Some(&b':') {
384 return None;
385 }
386 if let Some(rest) = strip_field(&line, b"data") {
388 self.current_data.push(rest);
389 self.has_data = true;
390 }
391 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 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 if payload == b"[DONE]" {
412 return Some(SseEvent::Done);
413 }
414 Some(SseEvent::Data(payload))
415 }
416}
417
418fn 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 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 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}