use super::{ChatRequest, ChatResponse, Provider, ProviderError, StreamEvent};
use async_trait::async_trait;
use futures::stream::BoxStream;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
#[derive(Debug, Clone, PartialEq)]
pub struct RetryPolicy {
pub max_attempts: usize,
pub backoff: Backoff,
pub retryable: Retryable,
pub respect_retry_after: bool,
}
impl Default for RetryPolicy {
fn default() -> Self {
Self {
max_attempts: 3,
backoff: Backoff::Exponential {
initial: Duration::from_millis(500),
factor: 2.0,
max: Duration::from_secs(10),
jitter: true,
},
retryable: Retryable::Default,
respect_retry_after: true,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum Backoff {
Fixed(Duration),
Exponential {
initial: Duration,
factor: f64,
max: Duration,
jitter: bool,
},
}
impl Backoff {
fn delay(&self, attempt: usize) -> Duration {
match self {
Self::Fixed(d) => *d,
Self::Exponential {
initial,
factor,
max,
jitter,
} => {
let base =
(initial.as_secs_f64() * factor.powf(attempt as f64)).min(max.as_secs_f64());
if *jitter {
Duration::from_secs_f64(base * random01())
} else {
Duration::from_secs_f64(base)
}
}
}
}
}
const RETRY_AFTER_CAP: Duration = Duration::from_secs(300);
fn random01() -> f64 {
const A: u64 = 6364136223846793005; const C: u64 = 1442695040888963407;
static STATE: AtomicU64 = AtomicU64::new(0);
loop {
let current = STATE.load(Ordering::Relaxed);
let next = if current == 0 {
let seed = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_nanos() as u64;
seed.wrapping_mul(A).wrapping_add(C)
} else {
current.wrapping_mul(A).wrapping_add(C)
};
if STATE
.compare_exchange_weak(current, next, Ordering::Relaxed, Ordering::Relaxed)
.is_ok()
{
return (next >> 40) as f64 / (1u64 << 24) as f64;
}
}
}
#[derive(Clone)]
pub enum Retryable {
Default,
Custom(Arc<dyn Fn(&ProviderError) -> bool + Send + Sync>),
}
impl std::fmt::Debug for Retryable {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Default => f.write_str("Default"),
Self::Custom(_) => f.write_str("Custom(_)"),
}
}
}
impl PartialEq for Retryable {
fn eq(&self, other: &Self) -> bool {
matches!((self, other), (Self::Default, Self::Default))
}
}
impl Retryable {
fn is_retryable(&self, error: &ProviderError) -> bool {
match self {
Self::Default => {
matches!(
error,
ProviderError::Network(_)
| ProviderError::Timeout(_)
| ProviderError::RateLimited { .. }
) || matches!(error, ProviderError::Api { status, .. } if *status >= 500)
}
Self::Custom(f) => f(error),
}
}
}
pub struct RetryProvider<I: Provider> {
inner: I,
policy: RetryPolicy,
}
impl<I: Provider> std::fmt::Debug for RetryProvider<I> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RetryProvider")
.field("inner", &std::any::type_name::<I>())
.field("policy", &self.policy)
.finish()
}
}
impl<I: Provider> RetryProvider<I> {
pub fn new(inner: I) -> Self {
Self {
inner,
policy: RetryPolicy::default(),
}
}
pub fn with_policy(mut self, policy: RetryPolicy) -> Self {
self.policy = policy;
self
}
}
impl<I: Provider> RetryProvider<I> {
fn retry_decision(&self, error: &ProviderError, attempts: usize) -> Option<Duration> {
if !self.policy.retryable.is_retryable(error) || attempts + 1 >= self.policy.max_attempts {
return None;
}
let delay = match (error, self.policy.respect_retry_after) {
(
ProviderError::RateLimited {
retry_after: Some(d),
},
true,
) => (*d).min(RETRY_AFTER_CAP),
_ => self.policy.backoff.delay(attempts),
};
tracing::warn!(
attempt = attempts + 1,
max_attempts = self.policy.max_attempts,
delay = ?delay,
error = %error,
"provider call failed, retrying",
);
Some(delay)
}
}
#[async_trait]
impl<I: Provider + Send + Sync> Provider for RetryProvider<I> {
async fn chat(&self, request: ChatRequest) -> Result<ChatResponse, ProviderError> {
let mut attempts = 0usize;
loop {
match self.inner.chat(request.clone()).await {
Ok(response) => return Ok(response),
Err(error) => match self.retry_decision(&error, attempts) {
Some(delay) => {
tokio::time::sleep(delay).await;
attempts += 1;
}
None => return Err(error),
},
}
}
}
async fn stream_chat(
&self,
request: ChatRequest,
) -> Result<BoxStream<'static, Result<StreamEvent, ProviderError>>, ProviderError> {
let mut attempts = 0usize;
loop {
match self.inner.stream_chat(request.clone()).await {
Ok(stream) => return Ok(stream),
Err(error) => match self.retry_decision(&error, attempts) {
Some(delay) => {
tokio::time::sleep(delay).await;
attempts += 1;
}
None => return Err(error),
},
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::provider::{FakeProvider, FakeReply};
use std::sync::Arc;
fn rate_limited() -> ProviderError {
ProviderError::RateLimited { retry_after: None }
}
fn api(status: u16) -> ProviderError {
ProviderError::Api {
status,
message: "boom".into(),
}
}
fn network() -> ProviderError {
ProviderError::Network("connection refused".into())
}
#[tokio::test]
async fn retries_then_succeeds() {
let inner = Arc::new(FakeProvider::new([
FakeReply::Error(rate_limited()),
FakeReply::Text("ok".into()),
]));
let provider = RetryProvider::new(inner.clone());
let resp = provider.chat(ChatRequest::default()).await.unwrap();
assert_eq!(resp.message, crate::message::Message::assistant("ok"));
assert_eq!(inner.requests().len(), 2);
}
#[tokio::test]
async fn gives_up_after_max_attempts() {
let inner = Arc::new(FakeProvider::new([
FakeReply::Error(network()),
FakeReply::Error(network()),
FakeReply::Error(network()),
]));
let provider = RetryProvider::new(inner.clone());
let err = provider.chat(ChatRequest::default()).await.unwrap_err();
assert!(matches!(err, ProviderError::Network(_)));
assert_eq!(inner.requests().len(), 3);
}
#[tokio::test]
async fn api_4xx_not_retried() {
let inner = Arc::new(FakeProvider::new([FakeReply::Error(api(400))]));
let provider = RetryProvider::new(inner.clone());
let err = provider.chat(ChatRequest::default()).await.unwrap_err();
assert!(matches!(err, ProviderError::Api { status: 400, .. }));
assert_eq!(inner.requests().len(), 1);
}
#[tokio::test]
async fn api_5xx_retried() {
let inner = Arc::new(FakeProvider::new([
FakeReply::Error(api(503)),
FakeReply::Text("ok".into()),
]));
let provider = RetryProvider::new(inner.clone());
provider.chat(ChatRequest::default()).await.unwrap();
assert_eq!(inner.requests().len(), 2);
}
#[tokio::test]
async fn max_attempts_one_disables_retry() {
let inner = Arc::new(FakeProvider::new([FakeReply::Error(network())]));
let provider = RetryProvider::new(inner.clone()).with_policy(RetryPolicy {
max_attempts: 1,
..Default::default()
});
provider.chat(ChatRequest::default()).await.unwrap_err();
assert_eq!(inner.requests().len(), 1);
}
#[test]
fn retry_after_capped_at_cap() {
let provider = RetryProvider::new(FakeProvider::new([]));
let delay = provider.retry_decision(
&ProviderError::RateLimited {
retry_after: Some(Duration::from_secs(3600)),
},
0,
);
assert_eq!(delay, Some(RETRY_AFTER_CAP));
}
#[tokio::test]
async fn retry_after_respected_when_present() {
let inner = Arc::new(FakeProvider::new([
FakeReply::Error(ProviderError::RateLimited {
retry_after: Some(Duration::from_millis(1)),
}),
FakeReply::Text("ok".into()),
]));
let provider = RetryProvider::new(inner.clone());
provider.chat(ChatRequest::default()).await.unwrap();
assert_eq!(inner.requests().len(), 2);
}
#[tokio::test]
async fn retry_after_delay_overrides_backoff() {
use std::time::Instant;
let inner = Arc::new(FakeProvider::new([
FakeReply::Error(ProviderError::RateLimited {
retry_after: Some(Duration::from_millis(200)),
}),
FakeReply::Text("ok".into()),
]));
let provider = RetryProvider::new(inner.clone()).with_policy(RetryPolicy {
backoff: Backoff::Exponential {
initial: Duration::from_millis(50),
factor: 1.0,
max: Duration::from_secs(1),
jitter: false,
},
..Default::default()
});
let start = Instant::now();
provider.chat(ChatRequest::default()).await.unwrap();
let elapsed = start.elapsed();
assert!(
elapsed >= Duration::from_millis(150) && elapsed < Duration::from_millis(600),
"should wait Retry-After (200ms) rather than 50ms backoff, elapsed {elapsed:?}"
);
}
#[tokio::test]
async fn custom_retryable_predicate() {
let inner = Arc::new(FakeProvider::new([
FakeReply::Error(api(400)),
FakeReply::Text("ok".into()),
]));
let provider = RetryProvider::new(inner.clone()).with_policy(RetryPolicy {
retryable: Retryable::Custom(Arc::new(
|e| matches!(e, ProviderError::Api { status, .. } if *status == 400),
)),
..Default::default()
});
provider.chat(ChatRequest::default()).await.unwrap();
assert_eq!(inner.requests().len(), 2);
}
#[tokio::test]
async fn stream_retries_only_before_first_event() {
let inner = Arc::new(FakeProvider::new([
FakeReply::Error(network()),
FakeReply::Text("ok".into()),
]));
let provider = RetryProvider::new(inner.clone());
let mut stream = provider.stream_chat(ChatRequest::default()).await.unwrap();
use futures::StreamExt;
let first = stream.next().await.unwrap().unwrap();
assert_eq!(first, StreamEvent::Delta("ok".into()));
assert_eq!(inner.requests().len(), 2);
}
#[tokio::test]
async fn stream_failure_after_max_attempts() {
let inner = Arc::new(FakeProvider::new([FakeReply::Error(network())]));
let provider = RetryProvider::new(inner.clone()).with_policy(RetryPolicy {
max_attempts: 1,
..Default::default()
});
let err = match provider.stream_chat(ChatRequest::default()).await {
Err(e) => e,
Ok(_) => panic!("expected error"),
};
assert!(matches!(err, ProviderError::Network(_)));
assert_eq!(inner.requests().len(), 1);
}
#[test]
fn exponential_backoff_grows_and_caps() {
let backoff = Backoff::Exponential {
initial: Duration::from_millis(500),
factor: 2.0,
max: Duration::from_secs(10),
jitter: false,
};
assert_eq!(backoff.delay(0), Duration::from_millis(500));
assert_eq!(backoff.delay(1), Duration::from_secs(1));
assert_eq!(backoff.delay(2), Duration::from_secs(2));
assert_eq!(backoff.delay(10), Duration::from_secs(10));
}
#[test]
fn jitter_stays_within_bounds() {
let backoff = Backoff::Exponential {
initial: Duration::from_secs(1),
factor: 1.0,
max: Duration::from_secs(10),
jitter: true,
};
for attempt in 0..50 {
let d = backoff.delay(attempt);
assert!(d < Duration::from_secs(1), "jitter out of bounds: {d:?}");
}
}
#[test]
fn jitter_first_draw_nonzero_and_dispersed() {
let first = random01();
assert!(
first > 0.0,
"first draw must not be 0 (bootstrap must advance LCG first)"
);
assert!(first < 1.0);
let mut seen = std::collections::HashSet::new();
for _ in 0..64 {
let d = random01();
assert!(d > 0.0 && d < 1.0);
seen.insert(d.to_bits());
}
assert!(
seen.len() > 1,
"LCG draws must be dispersed (state advancing)"
);
}
}