use std::future::Future;
use std::pin::Pin;
use std::time::Duration;
use bytes::Bytes;
use eventsource_stream::{Event, EventStreamError, Eventsource};
use futures::{Stream, StreamExt, stream};
use crate::error::AgentLoopError;
use crate::llm_retry::{LlmRetryConfig, RetryMetadata, remaining_retry_time, reserve_retry_wait};
pub type SseItem = Result<Event, EventStreamError<reqwest::Error>>;
pub type SseStream = Pin<Box<dyn Stream<Item = SseItem> + Send>>;
pub type ByteStream = Pin<Box<dyn Stream<Item = reqwest::Result<Bytes>> + Send>>;
const FIRST_STREAM_ITEM_TIMEOUT: Duration = Duration::from_secs(120);
pub fn is_reconnectable_reqwest_error(err: &reqwest::Error) -> bool {
err.is_body() || err.is_decode() || err.is_connect() || err.is_request() || err.is_timeout()
}
pub fn is_reconnectable_stream_error(err: &EventStreamError<reqwest::Error>) -> bool {
match err {
EventStreamError::Transport(e) => is_reconnectable_reqwest_error(e),
EventStreamError::Utf8(_) | EventStreamError::Parser(_) => false,
}
}
pub async fn connect_sse_with_reconnect<C, Fut>(
retry_config: &LlmRetryConfig,
driver_name: &str,
mut connect: C,
) -> Result<(SseStream, RetryMetadata), AgentLoopError>
where
C: FnMut(u32) -> Fut,
Fut: Future<Output = Result<(reqwest::Response, RetryMetadata), AgentLoopError>>,
{
connect_sse_with_reconnect_timeout(
retry_config,
driver_name,
&mut connect,
FIRST_STREAM_ITEM_TIMEOUT,
)
.await
}
async fn connect_sse_with_reconnect_timeout<C, Fut>(
retry_config: &LlmRetryConfig,
driver_name: &str,
connect: &mut C,
first_item_timeout: Duration,
) -> Result<(SseStream, RetryMetadata), AgentLoopError>
where
C: FnMut(u32) -> Fut,
Fut: Future<Output = Result<(reqwest::Response, RetryMetadata), AgentLoopError>>,
{
let mut retry_metadata = RetryMetadata::default();
let mut retry_started_at = None;
loop {
let connected = if let Some(remaining) =
remaining_retry_time(retry_config, retry_started_at)
{
tokio::time::timeout(remaining, connect(retry_metadata.attempts))
.await
.map_err(|_| {
AgentLoopError::llm_kind(
crate::error::LlmErrorKind::Unavailable,
format!(
"{driver_name} stream reconnect budget exhausted after {} retries over {:.1}s",
retry_metadata.attempts,
retry_config.max_retry_elapsed.as_secs_f64()
),
)
.with_retry_metadata(&retry_metadata)
})?
} else {
connect(retry_metadata.attempts).await
};
let (response, metadata) = connected?;
if retry_started_at.is_none() && metadata.had_retries() {
retry_started_at = tokio::time::Instant::now()
.checked_sub(metadata.total_retry_elapsed)
.or_else(|| Some(tokio::time::Instant::now()));
}
retry_metadata.absorb(metadata);
let events = Box::pin(response.bytes_stream().eventsource());
let item_timeout = remaining_retry_time(retry_config, retry_started_at)
.map_or(first_item_timeout, |remaining| {
remaining.min(first_item_timeout)
});
let (first, rest) = tokio::time::timeout(item_timeout, events.into_future())
.await
.map_err(|_| {
let error = AgentLoopError::llm(format!(
"provider stream stall: no first event for {}s",
first_item_timeout.as_secs()
));
if retry_metadata.had_retries() {
error.with_retry_metadata(&retry_metadata)
} else {
error
}
})?;
let reconnectable = matches!(&first, Some(Err(e)) if is_reconnectable_stream_error(e));
if reconnectable && retry_metadata.attempts < retry_config.max_retries {
let proposed_wait = retry_config.calculate_backoff(retry_metadata.attempts);
let Some(wait) = reserve_retry_wait(retry_config, &mut retry_started_at, proposed_wait)
else {
return Err(AgentLoopError::llm_kind(
crate::error::LlmErrorKind::Unavailable,
format!(
"{driver_name} stream reconnect budget exhausted after {} retries over {:.1}s",
retry_metadata.attempts,
retry_config.max_retry_elapsed.as_secs_f64()
),
)
.with_retry_metadata(&retry_metadata));
};
retry_metadata.record_retry(wait, None);
if let Some(Err(e)) = &first {
tracing::warn!(
driver = driver_name,
attempt = retry_metadata.attempts,
max_retries = retry_config.max_retries,
wait_secs = wait.as_secs_f64(),
error = %e,
"streaming transport failed before first event; reconnecting"
);
}
tokio::time::sleep(wait).await;
continue;
}
if reconnectable {
return Err(AgentLoopError::llm_kind(
crate::error::LlmErrorKind::Unavailable,
format!(
"{driver_name} stream transport failed after {} retries; the turn is safe to resume",
retry_metadata.attempts
),
)
.with_retry_metadata(&retry_metadata));
}
if retry_metadata.had_retries() {
tracing::info!(
driver = driver_name,
retries = retry_metadata.attempts,
"streaming reconnect succeeded"
);
}
return Ok((Box::pin(stream::iter(first).chain(rest)), retry_metadata));
}
}
pub async fn connect_bytes_with_reconnect<C, Fut>(
retry_config: &LlmRetryConfig,
driver_name: &str,
mut connect: C,
) -> Result<(ByteStream, RetryMetadata), AgentLoopError>
where
C: FnMut(u32) -> Fut,
Fut: Future<Output = Result<(reqwest::Response, RetryMetadata), AgentLoopError>>,
{
connect_bytes_with_reconnect_timeout(
retry_config,
driver_name,
&mut connect,
FIRST_STREAM_ITEM_TIMEOUT,
)
.await
}
async fn connect_bytes_with_reconnect_timeout<C, Fut>(
retry_config: &LlmRetryConfig,
driver_name: &str,
connect: &mut C,
first_item_timeout: Duration,
) -> Result<(ByteStream, RetryMetadata), AgentLoopError>
where
C: FnMut(u32) -> Fut,
Fut: Future<Output = Result<(reqwest::Response, RetryMetadata), AgentLoopError>>,
{
let mut retry_metadata = RetryMetadata::default();
let mut retry_started_at = None;
loop {
let connected = if let Some(remaining) =
remaining_retry_time(retry_config, retry_started_at)
{
tokio::time::timeout(remaining, connect(retry_metadata.attempts))
.await
.map_err(|_| {
AgentLoopError::llm_kind(
crate::error::LlmErrorKind::Unavailable,
format!(
"{driver_name} stream reconnect budget exhausted after {} retries over {:.1}s",
retry_metadata.attempts,
retry_config.max_retry_elapsed.as_secs_f64()
),
)
.with_retry_metadata(&retry_metadata)
})?
} else {
connect(retry_metadata.attempts).await
};
let (response, metadata) = connected?;
if retry_started_at.is_none() && metadata.had_retries() {
retry_started_at = tokio::time::Instant::now()
.checked_sub(metadata.total_retry_elapsed)
.or_else(|| Some(tokio::time::Instant::now()));
}
retry_metadata.absorb(metadata);
let bytes = Box::pin(response.bytes_stream());
let item_timeout = remaining_retry_time(retry_config, retry_started_at)
.map_or(first_item_timeout, |remaining| {
remaining.min(first_item_timeout)
});
let (first, rest) = tokio::time::timeout(item_timeout, bytes.into_future())
.await
.map_err(|_| {
let error = AgentLoopError::llm(format!(
"provider stream stall: no first chunk for {}s",
first_item_timeout.as_secs()
));
if retry_metadata.had_retries() {
error.with_retry_metadata(&retry_metadata)
} else {
error
}
})?;
let reconnectable = matches!(&first, Some(Err(e)) if is_reconnectable_reqwest_error(e));
if reconnectable && retry_metadata.attempts < retry_config.max_retries {
let proposed_wait = retry_config.calculate_backoff(retry_metadata.attempts);
let Some(wait) = reserve_retry_wait(retry_config, &mut retry_started_at, proposed_wait)
else {
return Err(AgentLoopError::llm_kind(
crate::error::LlmErrorKind::Unavailable,
format!(
"{driver_name} stream reconnect budget exhausted after {} retries over {:.1}s",
retry_metadata.attempts,
retry_config.max_retry_elapsed.as_secs_f64()
),
)
.with_retry_metadata(&retry_metadata));
};
retry_metadata.record_retry(wait, None);
if let Some(Err(e)) = &first {
tracing::warn!(
driver = driver_name,
attempt = retry_metadata.attempts,
max_retries = retry_config.max_retries,
wait_secs = wait.as_secs_f64(),
error = %e,
"streaming transport failed before first chunk; reconnecting"
);
}
tokio::time::sleep(wait).await;
continue;
}
if reconnectable {
return Err(AgentLoopError::llm_kind(
crate::error::LlmErrorKind::Unavailable,
format!(
"{driver_name} stream transport failed after {} retries; the turn is safe to resume",
retry_metadata.attempts
),
)
.with_retry_metadata(&retry_metadata));
}
if retry_metadata.had_retries() {
tracing::info!(
driver = driver_name,
retries = retry_metadata.attempts,
"streaming reconnect succeeded"
);
}
return Ok((Box::pin(stream::iter(first).chain(rest)), retry_metadata));
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
use std::time::Duration;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
#[derive(Clone, Copy)]
enum Behavior {
TruncateBeforeEvent,
EventThenTruncate,
FullSse,
Silent,
}
const CONTENT_EVENT: &str = "data: {\"choices\":[{\"delta\":{\"content\":\"hi\"}}]}\n\n";
const DONE_EVENT: &str = "data: [DONE]\n\n";
fn chunk(body: &str) -> String {
format!("{:x}\r\n{}\r\n", body.len(), body)
}
async fn spawn_scripted_sse_server(behaviors: Vec<Behavior>) -> (String, Arc<AtomicU32>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let counter = Arc::new(AtomicU32::new(0));
let counter_task = Arc::clone(&counter);
tokio::spawn(async move {
loop {
let (mut socket, _) = match listener.accept().await {
Ok(pair) => pair,
Err(_) => break,
};
let idx = counter_task.fetch_add(1, Ordering::SeqCst) as usize;
let behavior = behaviors[idx.min(behaviors.len() - 1)];
let mut buf = [0u8; 1024];
let _ = socket.read(&mut buf).await;
let headers = "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nTransfer-Encoding: chunked\r\n\r\n";
let _ = socket.write_all(headers.as_bytes()).await;
match behavior {
Behavior::TruncateBeforeEvent => {
let _ = socket.write_all(b"5\r\nda").await;
}
Behavior::EventThenTruncate => {
let _ = socket.write_all(chunk(CONTENT_EVENT).as_bytes()).await;
let _ = socket.write_all(b"5\r\nda").await;
}
Behavior::FullSse => {
let _ = socket.write_all(chunk(CONTENT_EVENT).as_bytes()).await;
let _ = socket.write_all(chunk(DONE_EVENT).as_bytes()).await;
let _ = socket.write_all(b"0\r\n\r\n").await;
}
Behavior::Silent => {
let _ = socket.flush().await;
std::future::pending::<()>().await;
}
}
let _ = socket.flush().await;
}
});
(format!("http://{addr}/"), counter)
}
fn fast_config(max_retries: u32) -> LlmRetryConfig {
LlmRetryConfig {
max_retries,
initial_backoff: Duration::from_millis(0),
max_backoff: Duration::from_millis(0),
backoff_multiplier: 1.0,
jitter_factor: 0.0,
..Default::default()
}
}
async fn collect_via_reconnect(
base: &str,
config: &LlmRetryConfig,
) -> Result<Vec<SseItem>, AgentLoopError> {
let client = reqwest::Client::builder().no_proxy().build().unwrap();
let (stream, _meta) = connect_sse_with_reconnect(config, "test", |_attempt| {
let client = client.clone();
let base = base.to_string();
async move {
let resp = client
.get(&base)
.send()
.await
.map_err(|e| AgentLoopError::llm(e.to_string()))?;
Ok((resp, RetryMetadata::default()))
}
})
.await?;
Ok(stream.collect().await)
}
async fn connect_via_reconnect_with_timeout(
base: &str,
config: &LlmRetryConfig,
first_item_timeout: Duration,
) -> Result<(SseStream, RetryMetadata), AgentLoopError> {
let client = reqwest::Client::builder().no_proxy().build().unwrap();
connect_sse_with_reconnect_timeout(
config,
"test",
&mut |_attempt| {
let client = client.clone();
let base = base.to_string();
async move {
let resp = client
.get(&base)
.send()
.await
.map_err(|e| AgentLoopError::llm(e.to_string()))?;
Ok((resp, RetryMetadata::default()))
}
},
first_item_timeout,
)
.await
}
#[tokio::test]
async fn reconnects_on_truncated_first_then_succeeds() {
let (base, count) =
spawn_scripted_sse_server(vec![Behavior::TruncateBeforeEvent, Behavior::FullSse]).await;
let items = collect_via_reconnect(&base, &fast_config(2))
.await
.expect("reconnect succeeds");
assert_eq!(
count.load(Ordering::SeqCst),
2,
"should reconnect exactly once"
);
let texts: Vec<String> = items
.iter()
.filter_map(|i| i.as_ref().ok())
.map(|ev| ev.data.clone())
.collect();
assert!(
texts.iter().any(|d| d.contains("hi")),
"expected content event after reconnect, got {texts:?}"
);
assert!(
items.iter().all(|i| i.is_ok()),
"reconnected stream should carry no transport error"
);
}
#[tokio::test]
async fn header_and_body_retries_share_one_attempt_budget() {
let (base, count) =
spawn_scripted_sse_server(vec![Behavior::TruncateBeforeEvent, Behavior::FullSse]).await;
let client = reqwest::Client::builder().no_proxy().build().unwrap();
let consumed = Arc::new(std::sync::Mutex::new(Vec::new()));
let observed = consumed.clone();
let (stream, metadata) = connect_sse_with_reconnect(&fast_config(2), "test", move |used| {
observed.lock().unwrap().push(used);
let client = client.clone();
let base = base.clone();
async move {
let response = client
.get(base)
.send()
.await
.map_err(|error| AgentLoopError::llm(error.to_string()))?;
let mut metadata = RetryMetadata::default();
if used == 0 {
metadata.record_retry(Duration::ZERO, None);
}
Ok((response, metadata))
}
})
.await
.expect("shared budget should recover");
let items: Vec<_> = stream.collect().await;
assert_eq!(count.load(Ordering::SeqCst), 2);
assert_eq!(*consumed.lock().unwrap(), vec![0, 2]);
assert_eq!(metadata.attempts, 2);
assert!(items.iter().all(Result::is_ok));
}
#[tokio::test]
async fn exhausts_reconnects_and_surfaces_error() {
let (base, count) = spawn_scripted_sse_server(vec![Behavior::TruncateBeforeEvent]).await;
let error = collect_via_reconnect(&base, &fast_config(2))
.await
.expect_err("reconnect budget should terminate");
assert_eq!(
count.load(Ordering::SeqCst),
3,
"should exhaust max_retries"
);
assert!(error.llm_retry_handled());
assert_eq!(error.llm_retry_attempts(), 2);
assert!(error.to_string().contains("safe to resume"));
}
#[tokio::test]
async fn clean_stream_makes_single_connection() {
let (base, count) = spawn_scripted_sse_server(vec![Behavior::FullSse]).await;
let items = collect_via_reconnect(&base, &fast_config(2))
.await
.expect("clean stream succeeds");
assert_eq!(
count.load(Ordering::SeqCst),
1,
"healthy stream must not reconnect"
);
assert!(items.iter().all(|i| i.is_ok()));
assert!(items.iter().filter_map(|i| i.as_ref().ok()).count() >= 1);
}
#[tokio::test]
async fn silent_first_event_is_bounded_without_reconnect() {
let (base, count) = spawn_scripted_sse_server(vec![Behavior::Silent]).await;
let err = match connect_via_reconnect_with_timeout(
&base,
&fast_config(2),
Duration::from_millis(20),
)
.await
{
Ok(_) => panic!("silent first event should fail at the setup-time bound"),
Err(err) => err,
};
assert_eq!(
count.load(Ordering::SeqCst),
1,
"setup-time stall must not retry up to the transport read timeout"
);
assert!(
err.to_string().contains("no first event"),
"expected first-event stall error, got {err}"
);
}
#[tokio::test]
async fn silent_first_chunk_is_bounded_without_reconnect() {
let (base, count) = spawn_scripted_sse_server(vec![Behavior::Silent]).await;
let client = reqwest::Client::builder().no_proxy().build().unwrap();
let err = match connect_bytes_with_reconnect_timeout(
&fast_config(2),
"test",
&mut |_attempt| {
let client = client.clone();
let base = base.to_string();
async move {
let resp = client
.get(&base)
.send()
.await
.map_err(|e| AgentLoopError::llm(e.to_string()))?;
Ok((resp, RetryMetadata::default()))
}
},
Duration::from_millis(20),
)
.await
{
Ok(_) => panic!("silent first chunk should fail at the setup-time bound"),
Err(err) => err,
};
assert_eq!(
count.load(Ordering::SeqCst),
1,
"setup-time stall must not retry up to the transport read timeout"
);
assert!(
err.to_string().contains("no first chunk"),
"expected first-chunk stall error, got {err}"
);
}
#[tokio::test]
async fn error_after_first_event_passes_through_without_reconnect() {
let (base, count) = spawn_scripted_sse_server(vec![Behavior::EventThenTruncate]).await;
let items = collect_via_reconnect(&base, &fast_config(2))
.await
.expect("first event commits stream");
assert_eq!(
count.load(Ordering::SeqCst),
1,
"committed stream must not reconnect after emitting an event"
);
assert!(
items.first().is_some_and(|i| i.is_ok()),
"first item should be the committed event"
);
assert!(
items.last().is_some_and(|i| i.is_err()),
"trailing transport error should pass through"
);
}
#[test]
fn classifier_rejects_non_transport_errors() {
let utf8_err = String::from_utf8(vec![0xff, 0xfe]).unwrap_err();
let err: EventStreamError<reqwest::Error> = EventStreamError::Utf8(utf8_err);
assert!(!is_reconnectable_stream_error(&err));
}
}