use std::time::Duration;
use crate::models::{BackendError, ModelError, Result};
pub const DEFAULT_MAX_ATTEMPTS: usize = 3;
const DEFAULT_INITIAL_DELAY_MS: u64 = 500;
const MAX_DELAY_MS: u64 = 3_000;
const MAX_RETRY_AFTER_MS: u64 = 60_000;
const RATE_LIMIT_DELAYS_MS: [u64; 2] = [2_000, 5_000];
pub async fn retry_transient_http<F, Fut>(mut build_and_send: F) -> Result<reqwest::Response>
where
F: FnMut() -> Fut,
Fut: std::future::Future<Output = Result<reqwest::Response>>,
{
retry_transient_http_with(
RetryPolicy {
max_attempts: DEFAULT_MAX_ATTEMPTS,
},
&mut build_and_send,
)
.await
}
async fn retry_transient_http_with<F, Fut>(
policy: RetryPolicy,
build_and_send: &mut F,
) -> Result<reqwest::Response>
where
F: FnMut() -> Fut,
Fut: std::future::Future<Output = Result<reqwest::Response>>,
{
let mut attempt: usize = 1;
let mut delay_ms = DEFAULT_INITIAL_DELAY_MS;
loop {
let result = build_and_send().await;
let transience = classify(&result);
let retry_after_ms = if transience.is_transient() {
result
.as_ref()
.ok()
.and_then(|r| parse_retry_after_ms(r.headers()))
} else {
None
};
if !transience.is_transient() || attempt >= policy.max_attempts {
if transience.is_transient() {
tracing::warn!(
attempts = attempt,
reason = transience.reason(),
"middleware: transient upstream failure — retries exhausted"
);
if transience.reason() == "http_429" && result.is_ok() {
let message = match result {
Ok(response) => response
.text()
.await
.ok()
.as_deref()
.and_then(extract_provider_error_message),
Err(_) => None,
};
return Err(ModelError::RateLimit {
retry_after: retry_after_ms.map(|ms| ms / 1000),
message,
});
}
}
return result;
}
let base_ms = if transience.reason() == "http_429" {
RATE_LIMIT_DELAYS_MS[(attempt - 1).min(RATE_LIMIT_DELAYS_MS.len() - 1)]
} else {
delay_ms
};
let sleep_ms = crate::utils::jitter(base_ms)
.max(retry_after_ms.unwrap_or(0))
.min(MAX_RETRY_AFTER_MS);
tracing::warn!(
attempt,
max = policy.max_attempts,
sleep_ms,
reason = transience.reason(),
"middleware: retrying transient upstream failure"
);
tokio::time::sleep(Duration::from_millis(sleep_ms)).await;
attempt += 1;
delay_ms = (delay_ms * 2).min(MAX_DELAY_MS);
}
}
fn extract_provider_error_message(body: &str) -> Option<String> {
const MAX_LEN: usize = 300;
let value: serde_json::Value = serde_json::from_str(body).ok()?;
let mut message = value
.pointer("/error/message")
.or_else(|| value.pointer("/errors/0/message"))
.or_else(|| value.pointer("/message"))
.and_then(|m| m.as_str())?
.trim();
while let Some((head, rest)) = message.split_once(": ") {
if head.ends_with("Error") && head.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') {
message = rest.trim_start();
} else {
break;
}
}
if message.is_empty() {
return None;
}
let mut message = message.to_string();
if message.len() > MAX_LEN {
let cut = (0..=MAX_LEN)
.rev()
.find(|&i| message.is_char_boundary(i))
.unwrap_or(0);
message.truncate(cut);
message.push('…');
}
Some(message)
}
fn parse_retry_after_ms(headers: &reqwest::header::HeaderMap) -> Option<u64> {
let raw = headers.get(reqwest::header::RETRY_AFTER)?.to_str().ok()?;
raw.trim()
.parse::<u64>()
.ok()
.map(|secs| secs.saturating_mul(1000))
}
#[derive(Debug, Clone, Copy)]
struct RetryPolicy {
max_attempts: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Transience {
Success,
Terminal,
Retryable(&'static str),
}
impl Transience {
fn is_transient(self) -> bool {
matches!(self, Transience::Retryable(_))
}
fn reason(self) -> &'static str {
match self {
Transience::Success => "success",
Transience::Terminal => "terminal",
Transience::Retryable(r) => r,
}
}
}
fn classify(result: &Result<reqwest::Response>) -> Transience {
match result {
Ok(resp) => {
let status = resp.status().as_u16();
if status == 429 {
Transience::Retryable("http_429")
} else if (500..=599).contains(&status) {
Transience::Retryable("http_5xx")
} else {
Transience::Success
}
},
Err(ModelError::Backend(BackendError::ConnectionFailed { .. })) => {
Transience::Retryable("connection_failed")
},
Err(_) => Transience::Terminal,
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
async fn fake_response(status: u16) -> reqwest::Response {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("local_addr");
tokio::spawn(async move {
if let Ok((mut sock, _)) = listener.accept().await {
let mut buf = [0u8; 1024];
let _ = sock.read(&mut buf).await;
let body = format!(
"HTTP/1.1 {status} X\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
);
let _ = sock.write_all(body.as_bytes()).await;
}
});
let url = format!("http://{}/x", addr);
reqwest::get(url).await.expect("send")
}
async fn fake_response_with_retry_after(
status: u16,
retry_after_secs: u64,
) -> reqwest::Response {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("local_addr");
tokio::spawn(async move {
if let Ok((mut sock, _)) = listener.accept().await {
let mut buf = [0u8; 1024];
let _ = sock.read(&mut buf).await;
let body = format!(
"HTTP/1.1 {status} X\r\nRetry-After: {retry_after_secs}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
);
let _ = sock.write_all(body.as_bytes()).await;
}
});
let url = format!("http://{}/x", addr);
reqwest::get(url).await.expect("send")
}
async fn fake_response_with_body(
status: u16,
body: &'static str,
retry_after_secs: Option<u64>,
) -> reqwest::Response {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("local_addr");
tokio::spawn(async move {
if let Ok((mut sock, _)) = listener.accept().await {
let mut buf = [0u8; 1024];
let _ = sock.read(&mut buf).await;
let retry_after = retry_after_secs
.map(|s| format!("Retry-After: {s}\r\n"))
.unwrap_or_default();
let response = format!(
"HTTP/1.1 {status} X\r\n{retry_after}Content-Length: {}\r\nConnection: close\r\n\r\n{body}",
body.len(),
);
let _ = sock.write_all(response.as_bytes()).await;
}
});
let url = format!("http://{}/x", addr);
reqwest::get(url).await.expect("send")
}
#[test]
fn extract_provider_error_message_handles_known_shapes() {
assert_eq!(
extract_provider_error_message(
r#"{"errors":[{"message":"you have used up your daily free allocation","code":4006}],"success":false}"#
),
Some("you have used up your daily free allocation".to_string())
);
assert_eq!(
extract_provider_error_message(
r#"{"error":{"message":"Rate limit reached for gpt-x","type":"tokens"}}"#
),
Some("Rate limit reached for gpt-x".to_string())
);
assert_eq!(
extract_provider_error_message(r#"{"message":"slow down"}"#),
Some("slow down".to_string())
);
assert_eq!(
extract_provider_error_message(
r#"{"errors":[{"message":"AiError: AiError: you have used up your daily free allocation"}]}"#
),
Some("you have used up your daily free allocation".to_string())
);
assert_eq!(
extract_provider_error_message(r#"{"message":"note: limits reset at midnight"}"#),
Some("note: limits reset at midnight".to_string())
);
assert_eq!(extract_provider_error_message("<html>429</html>"), None);
assert_eq!(extract_provider_error_message(r#"{"detail":"nope"}"#), None);
assert_eq!(extract_provider_error_message(r#"{"message":" "}"#), None);
let long = format!(r#"{{"message":"{}"}}"#, "x".repeat(400));
let extracted = extract_provider_error_message(&long).unwrap();
assert!(extracted.chars().count() <= 301);
assert!(extracted.ends_with('…'));
}
#[test]
fn parse_retry_after_handles_integer_seconds_and_ignores_dates() {
use reqwest::header::{HeaderMap, HeaderValue, RETRY_AFTER};
let mut headers = HeaderMap::new();
headers.insert(RETRY_AFTER, HeaderValue::from_static("2"));
assert_eq!(parse_retry_after_ms(&headers), Some(2_000));
let mut dated = HeaderMap::new();
dated.insert(
RETRY_AFTER,
HeaderValue::from_static("Wed, 21 Oct 2026 07:28:00 GMT"),
);
assert_eq!(parse_retry_after_ms(&dated), None);
assert_eq!(parse_retry_after_ms(&HeaderMap::new()), None);
}
#[tokio::test]
async fn honors_retry_after_on_503() {
let calls = Arc::new(AtomicUsize::new(0));
let cc = Arc::clone(&calls);
let start = std::time::Instant::now();
let result = retry_transient_http_with(RetryPolicy { max_attempts: 2 }, &mut move || {
let c = Arc::clone(&cc);
async move {
let n = c.fetch_add(1, Ordering::SeqCst);
if n == 0 {
Ok(fake_response_with_retry_after(503, 1).await)
} else {
Ok(fake_response(200).await)
}
}
})
.await;
let elapsed = start.elapsed();
assert!(result.is_ok());
assert_eq!(result.unwrap().status().as_u16(), 200);
assert_eq!(calls.load(Ordering::SeqCst), 2);
assert!(
elapsed >= Duration::from_millis(850),
"expected Retry-After (1s) to drive the 503 wait, waited only {:?}",
elapsed
);
}
#[tokio::test]
async fn retries_5xx_then_succeeds() {
let calls = Arc::new(AtomicUsize::new(0));
let cc = Arc::clone(&calls);
let result = retry_transient_http_with(RetryPolicy { max_attempts: 3 }, &mut move || {
let c = Arc::clone(&cc);
async move {
let n = c.fetch_add(1, Ordering::SeqCst);
let status = if n < 2 { 500 } else { 200 };
Ok(fake_response(status).await)
}
})
.await;
assert!(result.is_ok());
assert_eq!(result.unwrap().status().as_u16(), 200);
assert_eq!(calls.load(Ordering::SeqCst), 3);
}
#[tokio::test]
async fn does_not_retry_4xx_client_errors() {
let calls = Arc::new(AtomicUsize::new(0));
let cc = Arc::clone(&calls);
let result = retry_transient_http_with(RetryPolicy { max_attempts: 3 }, &mut move || {
let c = Arc::clone(&cc);
async move {
c.fetch_add(1, Ordering::SeqCst);
Ok(fake_response(400).await)
}
})
.await;
assert!(result.is_ok());
assert_eq!(result.unwrap().status().as_u16(), 400);
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn retries_429_then_surfaces_rate_limit() {
let calls = Arc::new(AtomicUsize::new(0));
let cc = Arc::clone(&calls);
let result = retry_transient_http_with(RetryPolicy { max_attempts: 2 }, &mut move || {
let c = Arc::clone(&cc);
async move {
c.fetch_add(1, Ordering::SeqCst);
Ok(fake_response(429).await)
}
})
.await;
assert!(matches!(result, Err(ModelError::RateLimit { .. })));
assert_eq!(calls.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn exhausted_429_carries_body_message_and_retry_after() {
let result = retry_transient_http_with(RetryPolicy { max_attempts: 1 }, &mut || async {
Ok(fake_response_with_body(
429,
r#"{"errors":[{"message":"you have used up your daily free allocation of 10,000 neurons","code":4006}]}"#,
Some(30),
)
.await)
})
.await;
match result {
Err(ModelError::RateLimit {
retry_after,
message,
}) => {
assert_eq!(retry_after, Some(30));
assert_eq!(
message.as_deref(),
Some("you have used up your daily free allocation of 10,000 neurons")
);
},
other => panic!("expected RateLimit, got {other:?}"),
}
}
#[tokio::test]
async fn rate_limit_backoff_is_slower_than_5xx_schedule() {
let calls = Arc::new(AtomicUsize::new(0));
let cc = Arc::clone(&calls);
let start = std::time::Instant::now();
let result = retry_transient_http_with(RetryPolicy { max_attempts: 2 }, &mut move || {
let c = Arc::clone(&cc);
async move {
let n = c.fetch_add(1, Ordering::SeqCst);
if n == 0 {
Ok(fake_response(429).await)
} else {
Ok(fake_response(200).await)
}
}
})
.await;
let elapsed = start.elapsed();
assert!(result.is_ok());
assert_eq!(result.unwrap().status().as_u16(), 200);
assert_eq!(calls.load(Ordering::SeqCst), 2);
assert!(
elapsed >= Duration::from_millis(1_500),
"expected the 429 schedule (~2s first delay) to drive the wait, waited only {:?}",
elapsed
);
}
#[tokio::test]
async fn retries_connection_failed_error() {
let calls = Arc::new(AtomicUsize::new(0));
let cc = Arc::clone(&calls);
let result = retry_transient_http_with(RetryPolicy { max_attempts: 3 }, &mut move || {
let c = Arc::clone(&cc);
async move {
let n = c.fetch_add(1, Ordering::SeqCst);
if n < 2 {
Err(ModelError::Backend(BackendError::ConnectionFailed {
backend: "test".to_string(),
url: "http://nope".to_string(),
reason: "dns".to_string(),
}))
} else {
Ok(fake_response(200).await)
}
}
})
.await;
assert!(result.is_ok());
assert_eq!(calls.load(Ordering::SeqCst), 3);
}
}