use crate::client::Everruns;
use crate::error::{Error, Result};
use crate::models::Event;
use futures::stream::Stream;
use serde::Deserialize;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::task::{Context, Poll};
use std::time::Duration;
use tokio::time::{Sleep, sleep};
const MAX_RETRY_MS: u64 = 30_000;
const INITIAL_BACKOFF_MS: u64 = 1000;
pub const READ_TIMEOUT_SECS: u64 = 45;
pub const DEFAULT_IDLE_TIMEOUT_SECS: u64 = 45;
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct StreamOptions {
pub types: Vec<String>,
pub exclude: Vec<String>,
pub since_id: Option<String>,
pub max_retries: Option<u32>,
pub idle_timeout: Duration,
}
impl Default for StreamOptions {
fn default() -> Self {
Self {
types: vec![],
exclude: vec![],
since_id: None,
max_retries: None,
idle_timeout: Duration::from_secs(DEFAULT_IDLE_TIMEOUT_SECS),
}
}
}
impl StreamOptions {
pub fn new() -> Self {
Self::default()
}
pub fn exclude_deltas() -> Self {
Self {
exclude: vec![
"output.message.delta".to_string(),
"reason.thinking.delta".to_string(),
],
..Self::default()
}
}
pub fn with_types(mut self, types: Vec<String>) -> Self {
self.types = types;
self
}
pub fn with_exclude(mut self, exclude: Vec<String>) -> Self {
self.exclude = exclude;
self
}
pub fn with_since_id(mut self, since_id: impl Into<String>) -> Self {
self.since_id = Some(since_id.into());
self
}
pub fn with_max_retries(mut self, max_retries: u32) -> Self {
self.max_retries = Some(max_retries);
self
}
pub fn with_idle_timeout(mut self, timeout: Duration) -> Self {
self.idle_timeout = timeout;
self
}
}
#[derive(Debug, Clone, serde::Serialize, Deserialize)]
pub struct DisconnectingData {
pub reason: String,
pub retry_ms: u64,
}
pub struct EventStream {
client: Everruns,
session_id: String,
options: StreamOptions,
inner: Option<Pin<Box<dyn Stream<Item = Result<Event>> + Send>>>,
last_event_id: Option<String>,
server_retry_ms: Option<u64>,
current_backoff_ms: u64,
retry_count: u32,
should_reconnect: bool,
graceful_disconnect: bool,
delay_future: Option<Pin<Box<Sleep>>>,
connected_signal: Arc<AtomicBool>,
sse_http_client: reqwest::Client,
idle_deadline: Option<Pin<Box<Sleep>>>,
idle_timeout: Duration,
}
impl EventStream {
pub(crate) fn new(client: Everruns, session_id: String, options: StreamOptions) -> Self {
let sse_http_client = reqwest::Client::builder()
.read_timeout(Duration::from_secs(READ_TIMEOUT_SECS))
.build()
.unwrap_or_else(|_| reqwest::Client::new());
let idle_timeout = options.idle_timeout;
Self {
client,
session_id,
options,
inner: None,
last_event_id: None,
server_retry_ms: None,
current_backoff_ms: INITIAL_BACKOFF_MS,
retry_count: 0,
should_reconnect: true,
graceful_disconnect: false,
delay_future: None,
connected_signal: Arc::new(AtomicBool::new(false)),
sse_http_client,
idle_deadline: None,
idle_timeout,
}
}
pub fn last_event_id(&self) -> Option<&str> {
self.last_event_id.as_deref()
}
pub fn stop(&mut self) {
self.should_reconnect = false;
self.inner = None;
self.delay_future = None;
self.idle_deadline = None;
}
pub fn retry_count(&self) -> u32 {
self.retry_count
}
fn connect(&mut self) -> Pin<Box<dyn Stream<Item = Result<Event>> + Send>> {
let client = self.client.clone();
let session_id = self.session_id.clone();
let since_id = self
.last_event_id
.clone()
.or_else(|| self.options.since_id.clone());
let types: Vec<String> = self.options.types.clone();
let exclude: Vec<String> = self.options.exclude.clone();
let connected_signal = self.connected_signal.clone();
let http_client = self.sse_http_client.clone();
Box::pin(async_stream::try_stream! {
use reqwest_eventsource::{Event as SseEvent, RequestBuilderExt};
use futures::StreamExt;
let types_refs: Vec<&str> = types.iter().map(|s| s.as_str()).collect();
let exclude_refs: Vec<&str> = exclude.iter().map(|s| s.as_str()).collect();
let url = client.sse_url(&session_id, since_id.as_deref(), &types_refs, &exclude_refs);
tracing::debug!("Connecting to SSE: {}", url);
let mut es = http_client
.get(url.clone())
.header("Authorization", client.auth_header())
.header("Accept", "text/event-stream")
.header("Cache-Control", "no-cache")
.eventsource()
.map_err(|e| Error::Sse(e.to_string()))?;
while let Some(event) = es.next().await {
match event {
Ok(SseEvent::Open) => {
tracing::debug!("SSE connection opened");
}
Ok(SseEvent::Message(msg)) => {
if msg.event == "connected" {
tracing::debug!("SSE connected event received");
connected_signal.store(true, Ordering::Release);
continue;
}
if msg.event == "disconnecting" {
if let Ok(data) = serde_json::from_str::<DisconnectingData>(&msg.data) {
tracing::debug!(
"SSE disconnecting: reason={}, retry_ms={}",
data.reason,
data.retry_ms
);
Err(Error::GracefulDisconnect {
reason: data.reason,
retry_ms: data.retry_ms,
})?;
} else {
tracing::debug!("SSE disconnecting event received (no data)");
Err(Error::GracefulDisconnect {
reason: "unknown".to_string(),
retry_ms: 100,
})?;
}
}
if let Ok(event) = serde_json::from_str::<Event>(&msg.data) {
yield event;
} else {
tracing::debug!("Skipping non-event message: {}", msg.event);
}
}
Err(reqwest_eventsource::Error::StreamEnded) => {
tracing::debug!("SSE stream ended");
break;
}
Err(e) => {
tracing::warn!("SSE error: {}", e);
Err(Error::Sse(e.to_string()))?;
}
}
}
})
}
fn get_retry_delay(&self) -> Duration {
if self.graceful_disconnect {
Duration::from_millis(self.server_retry_ms.unwrap_or(100))
} else {
Duration::from_millis(self.current_backoff_ms)
}
}
fn update_backoff(&mut self) {
if !self.graceful_disconnect {
self.current_backoff_ms = (self.current_backoff_ms * 2).min(MAX_RETRY_MS);
}
}
fn reset_backoff(&mut self) {
self.current_backoff_ms = INITIAL_BACKOFF_MS;
self.retry_count = 0;
}
fn should_retry(&self) -> bool {
if !self.should_reconnect {
return false;
}
match self.options.max_retries {
Some(max) => self.retry_count < max,
None => true,
}
}
fn schedule_reconnect(&mut self, delay: Duration) {
self.delay_future = Some(Box::pin(sleep(delay)));
}
}
impl Stream for EventStream {
type Item = Result<Event>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
loop {
if let Some(ref mut delay) = self.delay_future {
match Pin::new(delay).poll(cx) {
Poll::Ready(()) => {
self.delay_future = None;
self.graceful_disconnect = false;
}
Poll::Pending => {
return Poll::Pending;
}
}
}
if self.connected_signal.swap(false, Ordering::Acquire) {
self.reset_backoff();
}
if self.inner.is_none() {
if !self.should_reconnect {
return Poll::Ready(None);
}
self.inner = Some(self.connect());
self.idle_deadline = Some(Box::pin(sleep(self.idle_timeout)));
}
if let Some(ref mut idle) = self.idle_deadline
&& Pin::new(idle).poll(cx).is_ready()
{
tracing::warn!(
timeout_secs = self.idle_timeout.as_secs(),
"SSE idle timeout, reconnecting"
);
self.inner = None;
self.idle_deadline = None;
if self.should_retry() {
self.retry_count += 1;
let delay = self.get_retry_delay();
self.update_backoff();
self.schedule_reconnect(delay);
continue;
}
return Poll::Ready(None);
}
let inner = self.inner.as_mut().unwrap();
match Pin::new(inner).poll_next(cx) {
Poll::Ready(Some(Ok(event))) => {
self.reset_backoff();
self.last_event_id = Some(event.id.clone());
self.idle_deadline = Some(Box::pin(sleep(self.idle_timeout)));
return Poll::Ready(Some(Ok(event)));
}
Poll::Ready(Some(Err(e))) => {
if let Error::GracefulDisconnect { retry_ms, .. } = &e {
self.server_retry_ms = Some(*retry_ms);
self.graceful_disconnect = true;
self.inner = None;
self.idle_deadline = None;
if self.should_reconnect {
let delay = self.get_retry_delay();
tracing::debug!("Graceful reconnect in {:?}", delay);
self.schedule_reconnect(delay);
continue;
} else {
return Poll::Ready(None);
}
}
self.graceful_disconnect = false;
self.inner = None;
self.idle_deadline = None;
if self.should_retry() {
self.retry_count += 1;
let delay = self.get_retry_delay();
self.update_backoff();
tracing::debug!(
"Reconnecting after error in {:?} (attempt {})",
delay,
self.retry_count
);
self.schedule_reconnect(delay);
continue;
} else {
return Poll::Ready(Some(Err(e)));
}
}
Poll::Ready(None) => {
self.inner = None;
self.idle_deadline = None;
if self.should_retry() {
self.retry_count += 1;
let delay = self.get_retry_delay();
self.update_backoff();
tracing::debug!(
"Stream ended, reconnecting in {:?} (attempt {})",
delay,
self.retry_count
);
self.schedule_reconnect(delay);
continue;
}
return Poll::Ready(None);
}
Poll::Pending => return Poll::Pending,
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_stream_options_default() {
let opts = StreamOptions::default();
assert!(opts.exclude.is_empty());
assert!(opts.since_id.is_none());
assert!(opts.max_retries.is_none());
}
#[test]
fn test_stream_options_exclude_deltas() {
let opts = StreamOptions::exclude_deltas();
assert!(opts.exclude.contains(&"output.message.delta".to_string()));
assert!(opts.exclude.contains(&"reason.thinking.delta".to_string()));
}
#[test]
fn test_stream_options_builder() {
let opts = StreamOptions::default()
.with_since_id("event_123")
.with_max_retries(5)
.with_idle_timeout(Duration::from_secs(60));
assert_eq!(opts.since_id, Some("event_123".to_string()));
assert_eq!(opts.max_retries, Some(5));
assert_eq!(opts.idle_timeout, Duration::from_secs(60));
}
#[test]
fn test_stream_options_default_idle_timeout() {
let opts = StreamOptions::default();
assert_eq!(
opts.idle_timeout,
Duration::from_secs(DEFAULT_IDLE_TIMEOUT_SECS)
);
}
#[test]
fn test_disconnecting_data_parse() {
let json = r#"{"reason":"connection_cycle","retry_ms":100}"#;
let data: DisconnectingData = serde_json::from_str(json).unwrap();
assert_eq!(data.reason, "connection_cycle");
assert_eq!(data.retry_ms, 100);
}
}