use std::time::Duration;
use tokio_util::sync::CancellationToken;
use crate::error::AiError;
const DEFAULT_MAX_RETRY_DELAY: Duration = Duration::from_secs(60);
#[derive(Debug, Clone)]
pub struct ProviderRequestError {
pub status: Option<u16>,
pub message: String,
pub retry_after: Option<String>,
pub retry_after_ms: Option<String>,
pub x_should_retry: Option<String>,
}
impl ProviderRequestError {
pub fn is_retryable(&self) -> bool {
match self.x_should_retry.as_deref() {
Some("true") => return true,
Some("false") => return false,
_ => {}
}
match self.status {
None => true,
Some(s) => matches!(s, 408 | 409 | 429) || s >= 500,
}
}
}
impl From<&AiError> for ProviderRequestError {
fn from(err: &AiError) -> Self {
match err {
AiError::Http { status, message } => ProviderRequestError {
status: *status,
message: message.clone(),
retry_after: None,
retry_after_ms: None,
x_should_retry: None,
},
other => ProviderRequestError {
status: None,
message: other.to_string(),
retry_after: None,
retry_after_ms: None,
x_should_retry: None,
},
}
}
}
pub fn retry_delay(
err: &ProviderRequestError,
retry_index: u32,
max_retry_delay: Option<Duration>,
) -> Result<Duration, AiError> {
let cap = max_retry_delay.unwrap_or(DEFAULT_MAX_RETRY_DELAY);
if let Some(raw) = err.retry_after_ms.as_deref() {
if let Ok(value) = raw.trim().parse::<f64>() {
let delay = Duration::from_millis(value as u64);
return validate_server_delay(delay, cap, &err.message);
}
}
if let Some(raw) = err.retry_after.as_deref() {
let trimmed = raw.trim();
if let Ok(seconds) = trimmed.parse::<f64>() {
let delay = Duration::from_millis((seconds * 1000.0) as u64);
return validate_server_delay(delay, cap, &err.message);
}
}
let exp = (0.5f64 * 2f64.powi(retry_index as i32)).min(8.0);
let delay_ms = (exp * 1000.0) as u64;
Ok(Duration::from_millis(delay_ms))
}
fn validate_server_delay(
delay: Duration,
cap: Duration,
provider_message: &str,
) -> Result<Duration, AiError> {
if cap > Duration::ZERO && delay > cap {
return Err(AiError::Provider {
code: "retry_after_exceeds_cap".to_string(),
message: format!(
"Server requested {}s retry delay (max: {}s). {}",
delay.as_secs(),
cap.as_secs(),
provider_message
),
});
}
Ok(delay)
}
async fn abortable_sleep(delay: Duration, token: &CancellationToken) -> Result<(), AiError> {
if token.is_cancelled() {
return Err(AiError::Abort {
message: "Request aborted".to_string(),
});
}
tokio::select! {
_ = tokio::time::sleep(delay) => Ok(()),
_ = token.cancelled() => Err(AiError::Abort {
message: "Request aborted".to_string(),
}),
}
}
pub async fn retry_provider_request<T, F, Fut>(
request: F,
max_retries: Option<u32>,
max_retry_delay: Option<Duration>,
token: &CancellationToken,
) -> Result<T, AiError>
where
F: Fn() -> Fut,
Fut: std::future::Future<Output = Result<T, AiError>>,
{
let max_retries = max_retries.unwrap_or(0);
let mut retries_remaining = max_retries;
loop {
match request().await {
Ok(v) => return Ok(v),
Err(err) => {
if token.is_cancelled() {
return Err(AiError::Abort {
message: "Request aborted".to_string(),
});
}
if retries_remaining == 0 {
return Err(err);
}
let req_err = ProviderRequestError::from(&err);
if !req_err.is_retryable() {
return Err(err);
}
let retry_index = max_retries - retries_remaining;
retries_remaining -= 1;
let delay = retry_delay(&req_err, retry_index, max_retry_delay)?;
abortable_sleep(delay, token).await?;
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicU32, Ordering};
use std::sync::Arc;
fn http_err(status: Option<u16>) -> AiError {
AiError::Http {
status,
message: "boom".to_string(),
}
}
#[tokio::test]
async fn succeeds_first_try() {
let token = CancellationToken::new();
let v: u32 = retry_provider_request(|| async { Ok(7u32) }, Some(3), None, &token)
.await
.unwrap();
assert_eq!(v, 7);
}
#[tokio::test]
async fn retries_on_429_then_succeeds() {
let token = CancellationToken::new();
let attempts = Arc::new(AtomicU32::new(0));
let attempts_clone = Arc::clone(&attempts);
let v: u32 = retry_provider_request(
move || {
let attempts = Arc::clone(&attempts_clone);
async move {
let n = attempts.fetch_add(1, Ordering::SeqCst);
if n < 2 {
Err(http_err(Some(429)))
} else {
Ok(42u32)
}
}
},
Some(5),
None,
&token,
)
.await
.unwrap();
assert_eq!(v, 42);
assert_eq!(attempts.load(Ordering::SeqCst), 3);
}
#[tokio::test]
async fn does_not_retry_non_retryable_status() {
let token = CancellationToken::new();
let attempts = Arc::new(AtomicU32::new(0));
let attempts_clone = Arc::clone(&attempts);
let err = retry_provider_request::<u32, _, _>(
move || {
let attempts = Arc::clone(&attempts_clone);
async move {
attempts.fetch_add(1, Ordering::SeqCst);
Err(http_err(Some(400)))
}
},
Some(5),
None,
&token,
)
.await
.unwrap_err();
assert!(matches!(
err,
AiError::Http {
status: Some(400),
..
}
));
assert_eq!(attempts.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn abort_between_retries_surfaces_abort() {
let token = CancellationToken::new();
let token_clone = token.clone();
let attempts = Arc::new(AtomicU32::new(0));
let attempts_clone = Arc::clone(&attempts);
let handle = tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(5)).await;
token_clone.cancel();
});
let err = retry_provider_request::<u32, _, _>(
move || {
let attempts = Arc::clone(&attempts_clone);
async move {
attempts.fetch_add(1, Ordering::SeqCst);
Err(http_err(Some(503)))
}
},
Some(5),
None,
&token,
)
.await
.unwrap_err();
handle.await.unwrap();
assert!(matches!(err, AiError::Abort { .. }));
}
#[test]
fn retry_after_ms_header_honored() {
let err = ProviderRequestError {
status: Some(429),
message: "rate limited".into(),
retry_after: None,
retry_after_ms: Some("123".into()),
x_should_retry: None,
};
let delay = retry_delay(&err, 0, None).unwrap();
assert_eq!(delay, Duration::from_millis(123));
}
#[test]
fn retry_after_seconds_header_honored() {
let err = ProviderRequestError {
status: Some(429),
message: "rate limited".into(),
retry_after: Some("2".into()),
retry_after_ms: None,
x_should_retry: None,
};
let delay = retry_delay(&err, 0, None).unwrap();
assert_eq!(delay, Duration::from_secs(2));
}
#[test]
fn server_delay_above_cap_errors() {
let err = ProviderRequestError {
status: Some(429),
message: "rate limited".into(),
retry_after: Some("120".into()),
retry_after_ms: None,
x_should_retry: None,
};
let result = retry_delay(&err, 0, Some(Duration::from_secs(60)));
assert!(matches!(result, Err(AiError::Provider { .. })));
}
#[test]
fn x_should_retry_header_overrides_status() {
let retryable = ProviderRequestError {
status: Some(400),
message: "bad".into(),
retry_after: None,
retry_after_ms: None,
x_should_retry: Some("true".into()),
};
assert!(retryable.is_retryable());
let not_retryable = ProviderRequestError {
status: Some(503),
message: "boom".into(),
retry_after: None,
retry_after_ms: None,
x_should_retry: Some("false".into()),
};
assert!(!not_retryable.is_retryable());
}
#[test]
fn exponential_backoff_capped() {
let err = ProviderRequestError {
status: Some(503),
message: "boom".into(),
retry_after: None,
retry_after_ms: None,
x_should_retry: None,
};
assert_eq!(
retry_delay(&err, 0, None).unwrap(),
Duration::from_millis(500)
);
assert_eq!(
retry_delay(&err, 1, None).unwrap(),
Duration::from_millis(1000)
);
assert_eq!(
retry_delay(&err, 4, None).unwrap(),
Duration::from_millis(8000)
);
assert_eq!(
retry_delay(&err, 5, None).unwrap(),
Duration::from_millis(8000)
);
}
}