use crate::{Error, Result};
use std::future::Future;
use std::time::Duration;
use tokio::time::sleep;
#[derive(Debug, Clone)]
pub struct RetryConfig {
pub max_attempts: u32,
pub initial_delay: Duration,
pub max_delay: Duration,
pub backoff_multiplier: f64,
pub jitter_factor: f64,
}
impl Default for RetryConfig {
fn default() -> Self {
Self {
max_attempts: 3,
initial_delay: Duration::from_secs(1),
max_delay: Duration::from_secs(60),
backoff_multiplier: 2.0,
jitter_factor: 0.1,
}
}
}
impl RetryConfig {
pub fn new() -> Self {
Self::default()
}
pub fn with_max_attempts(mut self, attempts: u32) -> Self {
self.max_attempts = attempts;
self
}
pub fn with_initial_delay(mut self, delay: Duration) -> Self {
self.initial_delay = delay;
self
}
pub fn with_max_delay(mut self, delay: Duration) -> Self {
self.max_delay = delay;
self
}
pub fn with_backoff_multiplier(mut self, multiplier: f64) -> Self {
self.backoff_multiplier = multiplier;
self
}
pub fn with_jitter_factor(mut self, jitter: f64) -> Self {
self.jitter_factor = jitter.clamp(0.0, 1.0);
self
}
fn calculate_delay(&self, attempt: u32) -> Duration {
let base_delay_ms = self.initial_delay.as_millis() as f64;
let exponential_delay = base_delay_ms * self.backoff_multiplier.powi(attempt as i32);
let capped_delay = exponential_delay.min(self.max_delay.as_millis() as f64);
let jitter_range = capped_delay * self.jitter_factor;
let jitter = rand::random::<f64>() * jitter_range;
let jittered_delay = capped_delay + jitter - (jitter_range / 2.0);
let final_delay = jittered_delay.clamp(0.0, self.max_delay.as_millis() as f64);
Duration::from_millis(final_delay as u64)
}
}
pub async fn retry_with_backoff<F, Fut, T>(config: RetryConfig, mut operation: F) -> Result<T>
where
F: FnMut() -> Fut,
Fut: Future<Output = Result<T>>,
{
let mut last_error = None;
for attempt in 0..config.max_attempts {
match operation().await {
Ok(result) => return Ok(result),
Err(err) => {
last_error = Some(err);
if attempt < config.max_attempts - 1 {
let delay = config.calculate_delay(attempt);
sleep(delay).await;
}
}
}
}
Err(last_error.unwrap_or_else(|| Error::other("Retry failed with no error")))
}
const RETRYABLE_STATUS_CODES: &[u16] = &[408, 429, 500, 502, 503, 504, 529];
pub fn is_retryable_error(error: &Error) -> bool {
match error {
Error::Http(_) | Error::Timeout | Error::Stream(_) => true,
_ => error
.status_code()
.is_some_and(|status| RETRYABLE_STATUS_CODES.contains(&status)),
}
}
pub async fn retry_with_backoff_conditional<F, Fut, T>(
config: RetryConfig,
mut operation: F,
) -> Result<T>
where
F: FnMut() -> Fut,
Fut: Future<Output = Result<T>>,
{
let mut last_error = None;
for attempt in 0..config.max_attempts {
match operation().await {
Ok(result) => return Ok(result),
Err(err) => {
if !is_retryable_error(&err) {
return Err(err);
}
last_error = Some(err);
if attempt < config.max_attempts - 1 {
let delay = config.calculate_delay(attempt);
sleep(delay).await;
}
}
}
}
Err(last_error.unwrap_or_else(|| Error::other("Retry failed with no error")))
}
#[cfg(test)]
mod tests {
use super::*;
include!("retry/tests.rs");
}