use std::future::Future;
use std::marker::PhantomData;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Duration;
use bytes::Bytes;
use futures::future::BoxFuture;
use futures::stream::BoxStream;
use futures::{Stream, StreamExt};
use reqwest::Response;
use serde::de::DeserializeOwned;
use crate::error::{Error, Result};
use crate::retry::MAX_RECONNECT_BACKOFF;
pub(crate) type Reopen =
Arc<dyn Fn() -> BoxFuture<'static, Result<Response>> + Send + Sync + 'static>;
type ByteStream = BoxStream<'static, std::result::Result<Bytes, reqwest::Error>>;
type SleepFuture = Pin<Box<dyn Future<Output = ()> + Send + 'static>>;
type ReopenFuture = BoxFuture<'static, Result<Response>>;
pub struct EventStream<T: DeserializeOwned> {
state: State,
buf: SseBuffer,
reopen: Option<Reopen>,
reconnect_attempt: u32,
_marker: PhantomData<fn() -> T>,
}
impl<T: DeserializeOwned> std::fmt::Debug for EventStream<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("EventStream")
.field("reconnect_attempt", &self.reconnect_attempt)
.field("buffered", &self.buf.pending_len())
.field("state", &self.state.tag())
.finish()
}
}
enum State {
Reading(ByteStream),
Backoff(SleepFuture),
Reopening(ReopenFuture),
Done,
}
impl State {
fn tag(&self) -> &'static str {
match self {
State::Reading(_) => "reading",
State::Backoff(_) => "backoff",
State::Reopening(_) => "reopening",
State::Done => "done",
}
}
}
impl<T: DeserializeOwned> EventStream<T> {
#[allow(dead_code)] pub(crate) fn new(initial: Response, reopen: Reopen) -> Self {
Self {
state: State::Reading(initial.bytes_stream().boxed()),
buf: SseBuffer::default(),
reopen: Some(reopen),
reconnect_attempt: 0,
_marker: PhantomData,
}
}
#[cfg(test)]
pub(crate) fn from_bytes_stream(bytes: ByteStream) -> Self {
Self {
state: State::Reading(bytes),
buf: SseBuffer::default(),
reopen: None,
reconnect_attempt: 0,
_marker: PhantomData,
}
}
fn reconnect_delay(&self) -> Duration {
let base_ms = 100u64.saturating_mul(1u64 << self.reconnect_attempt.min(8));
let computed = Duration::from_millis(base_ms);
computed.min(MAX_RECONNECT_BACKOFF)
}
}
impl<T: DeserializeOwned> Stream for EventStream<T> {
type Item = Result<T>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
loop {
match self.buf.next_event() {
Some(SseEvent::Data(payload)) => {
let decoded: Result<T> = serde_json::from_slice(&payload)
.map_err(|e| Error::Stream(format!("malformed SSE payload: {e}")));
return Poll::Ready(Some(decoded));
}
Some(SseEvent::Done) => {
self.state = State::Done;
return Poll::Ready(None);
}
None => {}
}
let cur = std::mem::replace(&mut self.state, State::Done);
match cur {
State::Reading(mut s) => match s.poll_next_unpin(cx) {
Poll::Ready(Some(Ok(chunk))) => {
self.reconnect_attempt = 0;
self.buf.push(&chunk);
self.state = State::Reading(s);
continue;
}
Poll::Ready(Some(Err(e))) => {
let err: Error = e.into();
if err.is_transient() && self.reopen.is_some() {
self.reconnect_attempt = self.reconnect_attempt.saturating_add(1);
let delay = self.reconnect_delay();
self.state = State::Backoff(Box::pin(tokio::time::sleep(delay)));
continue;
}
self.state = State::Done;
return Poll::Ready(Some(Err(err)));
}
Poll::Ready(None) => {
if let Some(ev) = self.buf.finish() {
self.state = State::Done;
return match ev {
SseEvent::Data(payload) => {
let decoded: Result<T> = serde_json::from_slice(&payload)
.map_err(|e| {
Error::Stream(format!("malformed SSE payload: {e}"))
});
Poll::Ready(Some(decoded))
}
SseEvent::Done => Poll::Ready(None),
};
}
self.state = State::Done;
return Poll::Ready(None);
}
Poll::Pending => {
self.state = State::Reading(s);
return Poll::Pending;
}
},
State::Backoff(mut fut) => match fut.as_mut().poll(cx) {
Poll::Ready(()) => {
let reopen = self
.reopen
.clone()
.expect("Backoff state requires a reopen closure");
let f = (reopen)();
self.state = State::Reopening(f);
continue;
}
Poll::Pending => {
self.state = State::Backoff(fut);
return Poll::Pending;
}
},
State::Reopening(mut fut) => match fut.as_mut().poll(cx) {
Poll::Ready(Ok(resp)) => {
self.state = State::Reading(resp.bytes_stream().boxed());
continue;
}
Poll::Ready(Err(err)) => {
if err.is_transient() {
self.reconnect_attempt = self.reconnect_attempt.saturating_add(1);
let delay = self.reconnect_delay();
self.state = State::Backoff(Box::pin(tokio::time::sleep(delay)));
continue;
}
self.state = State::Done;
return Poll::Ready(Some(Err(err)));
}
Poll::Pending => {
self.state = State::Reopening(fut);
return Poll::Pending;
}
},
State::Done => {
self.state = State::Done;
return Poll::Ready(None);
}
}
}
}
}
#[derive(Debug, PartialEq, Eq)]
enum SseEvent {
Data(Vec<u8>),
Done,
}
#[derive(Default)]
struct SseBuffer {
pending: Vec<u8>,
current_data: Vec<Vec<u8>>,
has_data: bool,
}
impl SseBuffer {
fn pending_len(&self) -> usize {
self.pending.len()
}
fn push(&mut self, chunk: &[u8]) {
self.pending.extend_from_slice(chunk);
}
fn next_event(&mut self) -> Option<SseEvent> {
loop {
let idx = self.pending.iter().position(|&b| b == b'\n')?;
let mut line: Vec<u8> = self.pending.drain(..=idx).collect();
line.pop(); if line.last() == Some(&b'\r') {
line.pop();
}
if let Some(ev) = self.process_line(line) {
return Some(ev);
}
}
}
fn finish(&mut self) -> Option<SseEvent> {
if !self.pending.is_empty() {
let mut line = std::mem::take(&mut self.pending);
if line.last() == Some(&b'\r') {
line.pop();
}
if let Some(ev) = self.process_line(line) {
return Some(ev);
}
}
self.flush_event()
}
fn process_line(&mut self, line: Vec<u8>) -> Option<SseEvent> {
if line.is_empty() {
return self.flush_event();
}
if line.first() == Some(&b':') {
return None;
}
if let Some(rest) = strip_field(&line, b"data") {
self.current_data.push(rest);
self.has_data = true;
}
None
}
fn flush_event(&mut self) -> Option<SseEvent> {
if !self.has_data {
return None;
}
self.has_data = false;
let lines = std::mem::take(&mut self.current_data);
let mut payload: Vec<u8> = Vec::new();
for (i, l) in lines.iter().enumerate() {
if i > 0 {
payload.push(b'\n');
}
payload.extend_from_slice(l);
}
if payload == b"[DONE]" {
return Some(SseEvent::Done);
}
Some(SseEvent::Data(payload))
}
}
fn strip_field(line: &[u8], field: &[u8]) -> Option<Vec<u8>> {
if line.len() < field.len() + 1 {
return None;
}
if &line[..field.len()] != field {
return None;
}
if line[field.len()] != b':' {
return None;
}
let mut rest = &line[field.len() + 1..];
if rest.first() == Some(&b' ') {
rest = &rest[1..];
}
Some(rest.to_vec())
}
#[cfg(test)]
mod tests {
use super::*;
use futures::stream;
use pretty_assertions::assert_eq;
fn drain_buffer(buf: &mut SseBuffer) -> Vec<SseEvent> {
let mut out = Vec::new();
while let Some(ev) = buf.next_event() {
out.push(ev);
}
if let Some(ev) = buf.finish() {
out.push(ev);
}
out
}
#[test]
fn parses_single_event() {
let mut b = SseBuffer::default();
b.push(b"data: {\"x\":1}\n\n");
let events = drain_buffer(&mut b);
assert_eq!(events, vec![SseEvent::Data(b"{\"x\":1}".to_vec())]);
}
#[test]
fn parses_done_terminator() {
let mut b = SseBuffer::default();
b.push(b"data: [DONE]\n\n");
let events = drain_buffer(&mut b);
assert_eq!(events, vec![SseEvent::Done]);
}
#[test]
fn ignores_comment_lines() {
let mut b = SseBuffer::default();
b.push(b": heartbeat\ndata: {\"a\":1}\n\n");
let events = drain_buffer(&mut b);
assert_eq!(events, vec![SseEvent::Data(b"{\"a\":1}".to_vec())]);
}
#[test]
fn joins_multi_line_data() {
let mut b = SseBuffer::default();
b.push(b"data: line1\ndata: line2\n\n");
let events = drain_buffer(&mut b);
assert_eq!(events, vec![SseEvent::Data(b"line1\nline2".to_vec())]);
}
#[test]
fn handles_crlf_line_endings() {
let mut b = SseBuffer::default();
b.push(b"data: {\"x\":1}\r\n\r\n");
let events = drain_buffer(&mut b);
assert_eq!(events, vec![SseEvent::Data(b"{\"x\":1}".to_vec())]);
}
#[test]
fn handles_chunk_boundaries() {
let mut b = SseBuffer::default();
b.push(b"data: {\"x");
assert!(b.next_event().is_none());
b.push(b"\":1}\n");
assert!(b.next_event().is_none());
b.push(b"\n");
let events = drain_buffer(&mut b);
assert_eq!(events, vec![SseEvent::Data(b"{\"x\":1}".to_vec())]);
}
#[test]
fn ignores_non_data_fields() {
let mut b = SseBuffer::default();
b.push(b"event: ping\nid: 42\nretry: 1000\ndata: {\"x\":1}\n\n");
let events = drain_buffer(&mut b);
assert_eq!(events, vec![SseEvent::Data(b"{\"x\":1}".to_vec())]);
}
#[test]
fn flushes_trailing_event_without_blank_line() {
let mut b = SseBuffer::default();
b.push(b"data: {\"x\":1}\n");
let events = drain_buffer(&mut b);
assert_eq!(events, vec![SseEvent::Data(b"{\"x\":1}".to_vec())]);
}
#[test]
fn handles_empty_data_payload() {
let mut b = SseBuffer::default();
b.push(b"data: \n\n");
let events = drain_buffer(&mut b);
assert_eq!(events, vec![SseEvent::Data(Vec::new())]);
}
#[derive(serde::Deserialize, Debug, PartialEq)]
struct Sample {
x: i32,
}
#[tokio::test]
async fn event_stream_yields_decoded_events_then_done() {
let chunks: Vec<std::result::Result<Bytes, reqwest::Error>> = vec![
Ok(Bytes::from_static(b"data: {\"x\":1}\n\n")),
Ok(Bytes::from_static(b"data: {\"x\":2}\n\n")),
Ok(Bytes::from_static(b"data: [DONE]\n\n")),
];
let body: ByteStream = stream::iter(chunks).boxed();
let mut s: EventStream<Sample> = EventStream::from_bytes_stream(body);
let a = s.next().await.unwrap().unwrap();
let b = s.next().await.unwrap().unwrap();
assert_eq!(a, Sample { x: 1 });
assert_eq!(b, Sample { x: 2 });
assert!(s.next().await.is_none());
}
#[tokio::test]
async fn event_stream_surfaces_malformed_payload_as_error() {
let chunks: Vec<std::result::Result<Bytes, reqwest::Error>> =
vec![Ok(Bytes::from_static(b"data: not-json\n\n"))];
let body: ByteStream = stream::iter(chunks).boxed();
let mut s: EventStream<Sample> = EventStream::from_bytes_stream(body);
let item = s.next().await.unwrap();
assert!(matches!(item, Err(Error::Stream(_))));
}
#[tokio::test]
async fn event_stream_handles_split_event_across_chunks() {
let chunks: Vec<std::result::Result<Bytes, reqwest::Error>> = vec![
Ok(Bytes::from_static(b"data: {\"x")),
Ok(Bytes::from_static(b"\":7}\n\n")),
Ok(Bytes::from_static(b"data: [DONE]\n\n")),
];
let body: ByteStream = stream::iter(chunks).boxed();
let mut s: EventStream<Sample> = EventStream::from_bytes_stream(body);
let a = s.next().await.unwrap().unwrap();
assert_eq!(a, Sample { x: 7 });
assert!(s.next().await.is_none());
}
}