use std::time::Duration;
use tokio_retry::Retry;
use tokio_retry::strategy::{ExponentialBackoff, jitter};
use tracing::{debug, warn};
#[derive(Debug, Clone)]
pub struct RetryConfig {
pub initial_delay: Duration,
pub max_delay: Duration,
pub max_retries: usize,
pub backoff_multiplier: f64,
}
impl Default for RetryConfig {
fn default() -> Self {
Self {
initial_delay: Duration::from_millis(100),
max_delay: Duration::from_secs(60),
max_retries: 10,
backoff_multiplier: 2.0,
}
}
}
#[derive(Debug, Clone)]
pub struct BackoffConfig {
pub initial: Duration,
pub max: Duration,
pub multiplier: f64,
}
impl BackoffConfig {
pub fn into_strategy(self) -> impl Iterator<Item = Duration> {
let mut current = self.initial.as_millis() as u64;
let max_ms = self.max.as_millis() as u64;
std::iter::from_fn(move || {
let _delay = Duration::from_millis(current);
current = (current as f64 * self.multiplier) as u64;
if current > max_ms {
current = max_ms;
}
let jitter = (rand::random::<f64>() * 0.1 - 0.05) * current as f64;
Some(Duration::from_millis((current as f64 + jitter) as u64))
})
}
}
pub type RetryResult<T> = Result<T, anyhow::Error>;
#[allow(unused_mut)] pub async fn retry_with_backoff<F, Fut, T>(mut operation: F, config: RetryConfig) -> RetryResult<T>
where
F: FnMut() -> Fut,
Fut: std::future::Future<Output = RetryResult<T>>,
{
let strategy = config.build_strategy();
Retry::spawn(strategy, operation).await
}
impl RetryConfig {
pub fn fast() -> Self {
Self {
initial_delay: Duration::from_millis(50),
max_delay: Duration::from_secs(5),
max_retries: 5,
backoff_multiplier: 2.0,
}
}
pub fn slow() -> Self {
Self {
initial_delay: Duration::from_secs(1),
max_delay: Duration::from_secs(300), max_retries: 15,
backoff_multiplier: 2.0,
}
}
pub fn critical() -> Self {
Self {
initial_delay: Duration::from_millis(100),
max_delay: Duration::from_secs(120), max_retries: 20,
backoff_multiplier: 2.0,
}
}
pub fn build_strategy(&self) -> impl Iterator<Item = Duration> {
let backoff = ExponentialBackoff::from_millis(self.initial_delay.as_millis() as u64)
.max_delay(self.max_delay)
.take(self.max_retries.saturating_sub(1));
backoff.map(jitter)
}
}
pub async fn retry_dial<F, Fut, T, E>(
peer_id: &str,
config: RetryConfig,
mut dial_fn: F,
) -> Result<T, E>
where
F: FnMut() -> Fut,
Fut: std::future::Future<Output = Result<T, E>>,
E: std::fmt::Display,
{
let mut attempt = 0;
let strategy = config.build_strategy();
for delay in strategy {
attempt += 1;
match dial_fn().await {
Ok(result) => {
if attempt > 1 {
debug!("Dial to {} succeeded on attempt {}", peer_id, attempt);
}
return Ok(result);
}
Err(e) => {
debug!(
"Dial to {} failed (attempt {}): {} - retrying in {:?}",
peer_id, attempt, e, delay
);
tokio::time::sleep(delay).await;
}
}
}
attempt += 1;
dial_fn().await.map_err(|e| {
warn!(
"Dial to {} failed after {} attempts: {}",
peer_id, attempt, e
);
e
})
}
pub async fn retry_coordinator_discovery<F, Fut, T>(
config: RetryConfig,
operation: F,
) -> RetryResult<T>
where
F: FnMut() -> Fut,
Fut: std::future::Future<Output = RetryResult<T>>,
{
retry_with_backoff(operation, config).await
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
#[tokio::test]
async fn test_retry_succeeds_eventually() {
let attempts = Arc::new(AtomicUsize::new(0));
let attempts_clone = attempts.clone();
let config = RetryConfig {
initial_delay: Duration::from_millis(10),
max_delay: Duration::from_millis(100),
max_retries: 5,
backoff_multiplier: 2.0,
};
let result = retry_with_backoff(
|| {
let attempts = attempts_clone.clone();
async move {
let count = attempts.fetch_add(1, Ordering::SeqCst);
if count < 2 {
Err(anyhow::anyhow!("Not yet"))
} else {
Ok("Success")
}
}
},
config,
)
.await;
assert!(result.is_ok());
assert_eq!(result.unwrap(), "Success");
assert_eq!(attempts.load(Ordering::SeqCst), 3);
}
#[tokio::test]
async fn test_retry_fails_after_max_attempts() {
let attempts = Arc::new(AtomicUsize::new(0));
let attempts_clone = attempts.clone();
let config = RetryConfig {
initial_delay: Duration::from_millis(10),
max_delay: Duration::from_millis(50),
max_retries: 3,
backoff_multiplier: 2.0,
};
let result = retry_with_backoff(
|| {
let attempts = attempts_clone.clone();
async move {
attempts.fetch_add(1, Ordering::SeqCst);
Err::<(), _>(anyhow::anyhow!("Always fails"))
}
},
config,
)
.await;
assert!(result.is_err());
assert_eq!(attempts.load(Ordering::SeqCst), 3);
}
#[test]
fn test_retry_config_presets() {
let fast = RetryConfig::fast();
assert_eq!(fast.initial_delay, Duration::from_millis(50));
assert_eq!(fast.max_retries, 5);
let slow = RetryConfig::slow();
assert_eq!(slow.initial_delay, Duration::from_secs(1));
assert_eq!(slow.max_retries, 15);
let critical = RetryConfig::critical();
assert_eq!(critical.max_retries, 20);
}
#[tokio::test]
async fn test_retry_dial_with_logging() {
let attempts = Arc::new(AtomicUsize::new(0));
let attempts_clone = attempts.clone();
let config = RetryConfig {
initial_delay: Duration::from_millis(10),
max_delay: Duration::from_millis(50),
max_retries: 3,
backoff_multiplier: 2.0,
};
let result = retry_dial("test-peer", config, || {
let attempts = attempts_clone.clone();
async move {
let count = attempts.fetch_add(1, Ordering::SeqCst);
if count < 1 {
Err(anyhow::anyhow!("Connection refused"))
} else {
Ok(())
}
}
})
.await;
assert!(result.is_ok());
assert_eq!(attempts.load(Ordering::SeqCst), 2);
}
}