use std::time::Duration;
#[derive(Debug, Clone, Copy)]
pub(crate) struct RetryConfig {
pub max_retries: usize,
pub base_delay: Duration,
pub max_delay: Duration,
}
pub(crate) const DEFAULT_RETRY: RetryConfig = RetryConfig {
max_retries: 3,
base_delay: Duration::from_secs(1),
max_delay: Duration::from_secs(30),
};
pub(crate) async fn post_json_with_retry(
client: &reqwest::Client,
url: &str,
bearer_token: &str,
body: &serde_json::Value,
retry: &RetryConfig,
) -> Result<reqwest::Response, reqwest::Error> {
let mut attempt = 0usize;
loop {
let response = client
.post(url)
.header("Authorization", format!("Bearer {bearer_token}"))
.header("Content-Type", "application/json")
.json(body)
.send()
.await?;
let status = response.status();
if is_transient(&status) && attempt < retry.max_retries {
let retry_after_secs = response
.headers()
.get(reqwest::header::RETRY_AFTER)
.and_then(|v| v.to_str().ok())
.and_then(|v| v.trim().parse::<u64>().ok());
let delay = next_backoff(attempt, retry_after_secs, retry, entropy());
log::warn!(
"embedding HTTP {} (attempt {}), retrying in {:?}",
status,
attempt + 1,
delay
);
tokio::time::sleep(delay).await;
attempt += 1;
continue;
}
return Ok(response);
}
}
fn next_backoff(
attempt: usize,
retry_after_secs: Option<u64>,
retry: &RetryConfig,
entropy: u64,
) -> Duration {
let shift = 1u32.checked_shl(attempt as u32).unwrap_or(u32::MAX);
let base = retry.base_delay.saturating_mul(shift).min(retry.max_delay);
let jitter = base.mul_f64((entropy % 25) as f64 / 100.0);
let backoff = (base + jitter).min(retry.max_delay);
let server_delay = retry_after_secs
.map(|s| Duration::from_secs(s).min(retry.max_delay))
.unwrap_or(Duration::ZERO);
backoff.max(server_delay)
}
fn entropy() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.subsec_nanos() as u64)
.unwrap_or(0)
}
fn is_transient(status: &reqwest::StatusCode) -> bool {
status.as_u16() == 429 || status.as_u16() >= 500
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_support::spawn_status_stub;
use std::sync::atomic::Ordering;
fn cfg() -> RetryConfig {
RetryConfig {
max_retries: 3,
base_delay: Duration::from_millis(1),
max_delay: Duration::from_millis(5),
}
}
#[tokio::test]
async fn retry_succeeds_after_transient_429s() {
let (base_url, requests) = spawn_status_stub(429, 2, 200, "{\"ok\":true}").await;
let client = reqwest::Client::new();
let body = serde_json::json!({"model": "m", "input": ["a", "b"]});
let resp = post_json_with_retry(&client, &base_url, "test-key", &body, &cfg()).await;
assert!(
resp.is_ok(),
"should retry successfully after 429: {:?}",
resp.err()
);
assert_eq!(resp.unwrap().status().as_u16(), 200);
assert_eq!(requests.load(Ordering::SeqCst), 3, "1 initial + 2 retries");
}
#[tokio::test]
async fn retry_succeeds_after_transient_5xx() {
let (base_url, _requests) = spawn_status_stub(503, 1, 200, "{\"ok\":true}").await;
let client = reqwest::Client::new();
let body = serde_json::json!({"model": "m", "input": ["a"]});
let resp = post_json_with_retry(&client, &base_url, "test-key", &body, &cfg()).await;
assert!(resp.is_ok());
assert_eq!(resp.unwrap().status().as_u16(), 200);
}
#[tokio::test]
async fn retry_exhausts_and_returns_last_transient_response() {
let (base_url, requests) = spawn_status_stub(429, 100, 200, "{\"ok\":true}").await;
let client = reqwest::Client::new();
let body = serde_json::json!({"model": "m", "input": ["a"]});
let retry = RetryConfig {
max_retries: 2,
base_delay: Duration::from_millis(1),
max_delay: Duration::from_millis(5),
};
let resp = post_json_with_retry(&client, &base_url, "test-key", &body, &retry).await;
assert!(resp.is_ok(), "after retries are exhausted, should return the last response rather than a transport error");
assert_eq!(resp.unwrap().status().as_u16(), 429);
assert_eq!(requests.load(Ordering::SeqCst), 3, "1 initial + 2 retries");
}
#[tokio::test]
async fn does_not_retry_permanent_4xx() {
let (base_url, requests) = spawn_status_stub(400, 100, 200, "{\"ok\":true}").await;
let client = reqwest::Client::new();
let body = serde_json::json!({"model": "m", "input": ["a"]});
let resp = post_json_with_retry(&client, &base_url, "test-key", &body, &cfg()).await;
assert!(resp.is_ok());
assert_eq!(resp.unwrap().status().as_u16(), 400);
assert_eq!(
requests.load(Ordering::SeqCst),
1,
"4xx is a permanent failure, should not be retried"
);
}
#[test]
fn transient_status_classification() {
assert!(is_transient(&reqwest::StatusCode::TOO_MANY_REQUESTS));
assert!(is_transient(&reqwest::StatusCode::INTERNAL_SERVER_ERROR));
assert!(is_transient(&reqwest::StatusCode::SERVICE_UNAVAILABLE));
assert!(is_transient(&reqwest::StatusCode::BAD_GATEWAY));
assert!(!is_transient(&reqwest::StatusCode::BAD_REQUEST));
assert!(!is_transient(&reqwest::StatusCode::UNAUTHORIZED));
assert!(!is_transient(&reqwest::StatusCode::OK));
}
#[test]
fn next_backoff_zero_entropy_is_plain_base() {
let cfg = RetryConfig {
max_retries: 3,
base_delay: Duration::from_millis(1000),
max_delay: Duration::from_secs(30),
};
assert_eq!(next_backoff(0, None, &cfg, 0), Duration::from_millis(1000));
assert_eq!(next_backoff(2, None, &cfg, 0), Duration::from_millis(4000));
}
#[test]
fn next_backoff_jitter_is_bounded() {
let cfg = RetryConfig {
max_retries: 3,
base_delay: Duration::from_millis(1000),
max_delay: Duration::from_secs(30),
};
assert_eq!(next_backoff(0, None, &cfg, 24), Duration::from_millis(1240));
assert_eq!(next_backoff(0, None, &cfg, 25), Duration::from_millis(1000));
for entropy in [0u64, 1, 7, 13, 24, 999, u64::MAX] {
let delay = next_backoff(0, None, &cfg, entropy);
assert!(delay >= Duration::from_millis(1000));
assert!(delay < Duration::from_millis(1250));
}
}
#[test]
fn next_backoff_caps_at_max_delay() {
let cfg = RetryConfig {
max_retries: 3,
base_delay: Duration::from_secs(1),
max_delay: Duration::from_secs(5),
};
assert_eq!(next_backoff(10, None, &cfg, 0), Duration::from_secs(5));
}
#[test]
fn next_backoff_honors_larger_retry_after() {
let cfg = RetryConfig {
max_retries: 3,
base_delay: Duration::from_millis(1),
max_delay: Duration::from_millis(5),
};
assert_eq!(
next_backoff(0, Some(3600), &cfg, 0),
Duration::from_millis(5)
);
}
#[test]
fn next_backoff_ignores_smaller_retry_after() {
let cfg = RetryConfig {
max_retries: 3,
base_delay: Duration::from_millis(1000),
max_delay: Duration::from_secs(30),
};
assert_eq!(
next_backoff(0, Some(0), &cfg, 0),
Duration::from_millis(1000)
);
}
#[tokio::test]
async fn retry_after_header_raises_retry_delay() {
use crate::test_support::spawn_retry_after_stub;
use std::time::Instant;
let (base_url, requests) = spawn_retry_after_stub(1, 2).await;
let client = reqwest::Client::new();
let body = serde_json::json!({"model": "m", "input": ["a"]});
let retry = RetryConfig {
max_retries: 2,
base_delay: Duration::from_millis(1),
max_delay: Duration::from_millis(5),
};
let start = Instant::now();
let resp = post_json_with_retry(&client, &base_url, "test-key", &body, &retry).await;
let elapsed = start.elapsed();
assert!(resp.is_ok());
assert_eq!(resp.unwrap().status().as_u16(), 200);
assert_eq!(requests.load(Ordering::SeqCst), 3, "1 initial + 2 retries");
assert!(
elapsed >= Duration::from_millis(12),
"Retry-After should raise each retry delay to the 5ms cap (elapsed {:?})",
elapsed
);
}
}