use std::future::Future;
use std::time::{Duration, Instant};
use crate::error::LlmError;
use crate::provider::StatusTx;
use crate::usage::UsageTracker;
const BASE_BACKOFF_SECS: u64 = 1;
pub(crate) fn exponential_backoff_delay(attempt: u32) -> Duration {
Duration::from_secs(BASE_BACKOFF_SECS << attempt)
}
pub(crate) fn retry_delay(response: &reqwest::Response, attempt: u32) -> Duration {
if let Some(val) = response.headers().get("retry-after")
&& let Ok(s) = val.to_str()
&& let Ok(secs) = s.parse::<u64>()
{
return Duration::from_secs(secs);
}
exponential_backoff_delay(attempt)
}
pub(crate) async fn send_with_retry<F, Fut>(
provider_name: &str,
max_retries: u32,
status_tx: Option<&StatusTx>,
usage: Option<&UsageTracker>,
mut f: F,
) -> Result<reqwest::Response, LlmError>
where
F: FnMut() -> Fut,
Fut: Future<Output = Result<reqwest::Response, reqwest::Error>>,
{
for attempt in 0..=max_retries {
let send_start = Instant::now();
let response = f().await.map_err(LlmError::Http)?;
if let Some(tracker) = usage {
let ms = u64::try_from(send_start.elapsed().as_millis()).unwrap_or(u64::MAX);
tracker.record_ttft(ms);
}
let status = response.status();
if status == reqwest::StatusCode::TOO_MANY_REQUESTS
|| status == reqwest::StatusCode::SERVICE_UNAVAILABLE
{
if attempt == max_retries {
return Err(if status == reqwest::StatusCode::SERVICE_UNAVAILABLE {
LlmError::Unavailable
} else {
LlmError::RateLimited
});
}
let delay = retry_delay(&response, attempt);
let msg = format!(
"{provider_name} rate limited or unavailable, retrying in {}s ({}/{})",
delay.as_secs(),
attempt + 1,
max_retries
);
if let Some(tx) = status_tx {
let _ = tx.send(msg.clone());
}
tracing::warn!("{msg}");
tokio::time::sleep(delay).await;
continue;
}
return Ok(response);
}
Err(LlmError::RateLimited)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn retry_delay_exponential_backoff() {
assert_eq!(BASE_BACKOFF_SECS, 1);
assert_eq!(BASE_BACKOFF_SECS << 1, 2);
assert_eq!(BASE_BACKOFF_SECS << 2, 4);
}
async fn spawn_mock_server(responses: Vec<&'static str>) -> (u16, tokio::task::JoinHandle<()>) {
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server_loop = async move {
for resp in responses {
let Ok((mut stream, _)) = listener.accept().await else {
break;
};
let conn = async move {
let (reader, mut writer) = stream.split();
let mut buf_reader = BufReader::new(reader);
let mut line = String::new();
loop {
line.clear();
buf_reader.read_line(&mut line).await.unwrap_or(0);
if line == "\r\n" || line == "\n" || line.is_empty() {
break;
}
}
writer.write_all(resp.as_bytes()).await.ok();
};
tokio::spawn(conn); }
};
let handle = tokio::spawn(server_loop);
(port, handle)
}
#[tokio::test]
async fn send_with_retry_success_on_first_attempt() {
let ok_response = "HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok";
let (port, _handle) = spawn_mock_server(vec![ok_response]).await;
let client = reqwest::Client::new();
let url = format!("http://127.0.0.1:{port}/test");
let result = send_with_retry("test", 3, None, None, || {
let req = client.get(&url).build().unwrap();
let c = client.clone();
async move { c.execute(req).await }
})
.await;
assert!(result.is_ok(), "expected Ok, got: {result:?}");
assert_eq!(result.unwrap().status(), 200);
}
#[tokio::test]
async fn send_with_retry_exhausts_retries_returns_rate_limited() {
let rate_limit_response =
"HTTP/1.1 429 Too Many Requests\r\nRetry-After: 0\r\nContent-Length: 0\r\n\r\n";
let (port, _handle) =
spawn_mock_server(vec![rate_limit_response, rate_limit_response]).await;
let client = reqwest::Client::new();
let url = format!("http://127.0.0.1:{port}/test");
let result = send_with_retry("test", 1, None, None, || {
let req = client.get(&url).build().unwrap();
let c = client.clone();
async move { c.execute(req).await }
})
.await;
assert!(
matches!(result, Err(LlmError::RateLimited)),
"expected RateLimited, got: {result:?}"
);
}
#[tokio::test]
async fn send_with_retry_exhausts_retries_returns_unavailable() {
let unavailable_response =
"HTTP/1.1 503 Service Unavailable\r\nRetry-After: 0\r\nContent-Length: 0\r\n\r\n";
let (port, _handle) =
spawn_mock_server(vec![unavailable_response, unavailable_response]).await;
let client = reqwest::Client::new();
let url = format!("http://127.0.0.1:{port}/test");
let result = send_with_retry("test", 1, None, None, || {
let req = client.get(&url).build().unwrap();
let c = client.clone();
async move { c.execute(req).await }
})
.await;
assert!(
matches!(result, Err(LlmError::Unavailable)),
"expected Unavailable, got: {result:?}"
);
}
#[tokio::test]
async fn send_with_retry_mixed_429_then_503_returns_unavailable() {
let rate_limit_response =
"HTTP/1.1 429 Too Many Requests\r\nRetry-After: 0\r\nContent-Length: 0\r\n\r\n";
let unavailable_response =
"HTTP/1.1 503 Service Unavailable\r\nRetry-After: 0\r\nContent-Length: 0\r\n\r\n";
let (port, _handle) = spawn_mock_server(vec![
rate_limit_response,
unavailable_response,
unavailable_response,
])
.await;
let client = reqwest::Client::new();
let url = format!("http://127.0.0.1:{port}/test");
let result = send_with_retry("test", 2, None, None, || {
let req = client.get(&url).build().unwrap();
let c = client.clone();
async move { c.execute(req).await }
})
.await;
assert!(
matches!(result, Err(LlmError::Unavailable)),
"last status (503) must win over the earlier 429, got: {result:?}"
);
}
#[tokio::test]
async fn send_with_retry_mixed_503_then_429_returns_rate_limited() {
let rate_limit_response =
"HTTP/1.1 429 Too Many Requests\r\nRetry-After: 0\r\nContent-Length: 0\r\n\r\n";
let unavailable_response =
"HTTP/1.1 503 Service Unavailable\r\nRetry-After: 0\r\nContent-Length: 0\r\n\r\n";
let (port, _handle) = spawn_mock_server(vec![
unavailable_response,
rate_limit_response,
rate_limit_response,
])
.await;
let client = reqwest::Client::new();
let url = format!("http://127.0.0.1:{port}/test");
let result = send_with_retry("test", 2, None, None, || {
let req = client.get(&url).build().unwrap();
let c = client.clone();
async move { c.execute(req).await }
})
.await;
assert!(
matches!(result, Err(LlmError::RateLimited)),
"last status (429) must win over the earlier 503, got: {result:?}"
);
}
#[tokio::test]
async fn send_with_retry_succeeds_after_one_429() {
let rate_limit_response =
"HTTP/1.1 429 Too Many Requests\r\nRetry-After: 0\r\nContent-Length: 0\r\n\r\n";
let ok_response = "HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok";
let (port, _handle) = spawn_mock_server(vec![rate_limit_response, ok_response]).await;
let client = reqwest::Client::new();
let url = format!("http://127.0.0.1:{port}/test");
let result = send_with_retry("test", 2, None, None, || {
let req = client.get(&url).build().unwrap();
let c = client.clone();
async move { c.execute(req).await }
})
.await;
assert!(
result.is_ok(),
"expected Ok after one retry, got: {result:?}"
);
assert_eq!(result.unwrap().status(), 200);
}
#[tokio::test]
async fn send_with_retry_ttft_reflects_final_attempt_not_backoff_delay() {
let rate_limit_response =
"HTTP/1.1 429 Too Many Requests\r\nRetry-After: 1\r\nContent-Length: 0\r\n\r\n";
let ok_response = "HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok";
let (port, _handle) = spawn_mock_server(vec![rate_limit_response, ok_response]).await;
let client = reqwest::Client::new();
let url = format!("http://127.0.0.1:{port}/test");
let usage = UsageTracker::default();
let result = send_with_retry("test", 1, None, Some(&usage), || {
let req = client.get(&url).build().unwrap();
let c = client.clone();
async move { c.execute(req).await }
})
.await;
assert!(
result.is_ok(),
"expected Ok after one retry, got: {result:?}"
);
let ttft = usage
.last_ttft_ms()
.expect("send_with_retry must record a ttft sample when usage is Some");
assert!(
ttft < 500,
"ttft_ms={ttft} must reflect only the final attempt's send time, not the \
1000ms Retry-After backoff sleep between attempts"
);
}
#[tokio::test]
async fn send_with_retry_records_ttft_on_first_attempt_success_too() {
let ok_response = "HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok";
let (port, _handle) = spawn_mock_server(vec![ok_response]).await;
let client = reqwest::Client::new();
let url = format!("http://127.0.0.1:{port}/test");
let usage = UsageTracker::default();
assert!(usage.last_ttft_ms().is_none());
let result = send_with_retry("test", 3, None, Some(&usage), || {
let req = client.get(&url).build().unwrap();
let c = client.clone();
async move { c.execute(req).await }
})
.await;
assert!(result.is_ok());
assert!(usage.last_ttft_ms().is_some());
}
use proptest::prelude::*;
proptest! {
#[test]
fn retry_delay_range_always_valid(attempt in 0u32..63) {
let delay = Duration::from_secs(BASE_BACKOFF_SECS << attempt);
assert!(delay.as_secs() >= BASE_BACKOFF_SECS, "delay must be at least base backoff");
if attempt > 0 {
let prev = Duration::from_secs(BASE_BACKOFF_SECS << (attempt - 1));
assert_eq!(delay.as_secs(), prev.as_secs() * 2);
}
}
}
}