use std::time::Duration;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ProviderError {
RateLimit { retry_after: Option<Duration> },
Overloaded,
ServerError,
Auth,
Billing,
ContextOverflow,
Invalid(String),
Transport,
}
impl std::fmt::Display for ProviderError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ProviderError::RateLimit {
retry_after: Some(d),
} => {
write!(f, "rate limited (retry after {}s)", d.as_secs())
}
ProviderError::RateLimit { retry_after: None } => write!(f, "rate limited"),
ProviderError::Overloaded => write!(f, "provider overloaded"),
ProviderError::ServerError => write!(f, "provider server error"),
ProviderError::Auth => write!(f, "authentication failed"),
ProviderError::Billing => write!(f, "billing/credit failure"),
ProviderError::ContextOverflow => write!(f, "prompt exceeds the context window"),
ProviderError::Invalid(detail) => write!(f, "invalid request: {detail}"),
ProviderError::Transport => write!(f, "transport failure"),
}
}
}
impl std::error::Error for ProviderError {}
impl ProviderError {
pub fn transient(&self) -> bool {
matches!(
self,
ProviderError::RateLimit { .. }
| ProviderError::Overloaded
| ProviderError::ServerError
| ProviderError::Transport
)
}
}
pub fn overflow_text(text: &str) -> bool {
let t = text.to_ascii_lowercase();
t.contains("exceed_context_size")
|| t.contains("context_length_exceeded")
|| t.contains("context length")
|| t.contains("context size")
|| t.contains("prompt is too long")
|| t.contains("too many tokens")
|| t.contains("maximum context")
}
pub fn classify_http(status: u16, body: &str, retry_after: Option<Duration>) -> ProviderError {
let lower = body.to_ascii_lowercase();
match status {
401 | 403 => ProviderError::Auth,
402 => ProviderError::Billing,
429 => ProviderError::RateLimit { retry_after },
529 => ProviderError::Overloaded,
503 if lower.contains("overload") => ProviderError::Overloaded,
_ if overflow_text(body) => ProviderError::ContextOverflow,
s if s >= 500 => ProviderError::ServerError,
_ if lower.contains("credit balance") || lower.contains("billing") => {
ProviderError::Billing
}
_ => ProviderError::Invalid(body.chars().take(200).collect()),
}
}
#[derive(Debug, Clone)]
pub struct RetryPolicy {
pub max_retries: u32,
pub retry_after_cap: Duration,
pub base_delay: Duration,
}
impl Default for RetryPolicy {
fn default() -> Self {
RetryPolicy {
max_retries: 3,
retry_after_cap: Duration::from_secs(60),
base_delay: Duration::from_millis(2_500),
}
}
}
impl RetryPolicy {
pub const MAX_DELAY: Duration = Duration::from_secs(30);
pub fn from_config(cfg: &crate::config::ProviderConfig) -> Self {
let d = RetryPolicy::default();
RetryPolicy {
max_retries: cfg.max_retries.unwrap_or(d.max_retries),
retry_after_cap: cfg
.retry_after_cap_secs
.map(Duration::from_secs)
.unwrap_or(d.retry_after_cap),
base_delay: d.base_delay,
}
}
pub fn delay_for(&self, error: &ProviderError, attempt: u32) -> Option<Duration> {
if attempt > self.max_retries || !error.transient() {
return None;
}
match error {
ProviderError::RateLimit {
retry_after: Some(after),
} => {
(*after <= self.retry_after_cap).then_some(*after)
}
_ => {
let exp = self
.base_delay
.saturating_mul(1u32 << (attempt - 1).min(16));
Some(exp.min(Self::MAX_DELAY))
}
}
}
}
#[derive(Debug)]
pub struct RequestFailure {
pub class: ProviderError,
pub status: Option<u16>,
pub detail: String,
}
pub async fn send_with_retry(
make_request: impl Fn() -> reqwest::RequestBuilder,
policy: &RetryPolicy,
) -> Result<reqwest::Response, RequestFailure> {
let mut attempt = 0u32;
loop {
let failure = match make_request().send().await {
Ok(resp) if resp.status().is_success() => return Ok(resp),
Ok(resp) => {
let status = resp.status().as_u16();
let retry_after = resp
.headers()
.get(reqwest::header::RETRY_AFTER)
.and_then(|v| v.to_str().ok())
.and_then(|s| s.trim().parse::<u64>().ok())
.map(Duration::from_secs);
let body = resp.text().await.unwrap_or_default();
RequestFailure {
class: classify_http(status, &body, retry_after),
status: Some(status),
detail: body,
}
}
Err(e) => RequestFailure {
class: ProviderError::Transport,
status: None,
detail: e.to_string(),
},
};
attempt += 1;
match policy.delay_for(&failure.class, attempt) {
Some(delay) => {
tracing::warn!(
error = %failure.class,
attempt,
delay_ms = delay.as_millis() as u64,
"provider request failed; retrying"
);
tokio::time::sleep(delay).await;
}
None => return Err(failure),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn each_class_gets_its_policy() {
let p = RetryPolicy {
base_delay: Duration::from_millis(10),
..Default::default()
};
for err in [
ProviderError::Overloaded,
ProviderError::ServerError,
ProviderError::Transport,
] {
assert_eq!(p.delay_for(&err, 1), Some(Duration::from_millis(10)));
assert_eq!(p.delay_for(&err, 2), Some(Duration::from_millis(20)));
assert_eq!(p.delay_for(&err, 4), None, "exhausted past max_retries");
}
for err in [
ProviderError::Auth,
ProviderError::Billing,
ProviderError::Invalid("x".into()),
ProviderError::ContextOverflow,
] {
assert_eq!(p.delay_for(&err, 1), None);
}
}
#[test]
fn retry_after_is_honoured_when_sane_and_a_failure_when_hostile() {
let p = RetryPolicy::default();
let soon = ProviderError::RateLimit {
retry_after: Some(Duration::from_secs(3)),
};
assert_eq!(p.delay_for(&soon, 1), Some(Duration::from_secs(3)));
let hostile = ProviderError::RateLimit {
retry_after: Some(Duration::from_secs(3_600)),
};
assert_eq!(p.delay_for(&hostile, 1), None);
let unstated = ProviderError::RateLimit { retry_after: None };
assert_eq!(p.delay_for(&unstated, 1), Some(p.base_delay));
}
#[test]
fn zero_max_retries_disables_retrying() {
let p = RetryPolicy {
max_retries: 0,
..Default::default()
};
assert_eq!(p.delay_for(&ProviderError::Transport, 1), None);
}
#[test]
fn the_backoff_never_exceeds_the_ceiling() {
let p = RetryPolicy {
max_retries: 40,
..Default::default()
};
assert_eq!(
p.delay_for(&ProviderError::Transport, 39),
Some(RetryPolicy::MAX_DELAY)
);
}
#[test]
fn classification_reads_status_and_text() {
use ProviderError::*;
assert_eq!(classify_http(401, "", None), Auth);
assert_eq!(classify_http(403, "", None), Auth);
assert_eq!(
classify_http(429, "", Some(Duration::from_secs(2))),
RateLimit {
retry_after: Some(Duration::from_secs(2))
}
);
assert_eq!(classify_http(529, "", None), Overloaded);
assert_eq!(
classify_http(503, "The server is overloaded", None),
Overloaded
);
assert_eq!(classify_http(500, "", None), ServerError);
assert_eq!(classify_http(503, "", None), ServerError);
assert_eq!(
classify_http(400, r#"{"type":"exceed_context_size_error"}"#, None),
ContextOverflow
);
assert_eq!(
classify_http(500, "Context size has been exceeded.", None),
ContextOverflow
);
assert_eq!(
classify_http(400, "Your credit balance is too low", None),
Billing
);
assert!(matches!(classify_http(400, "bad json", None), Invalid(_)));
}
}