use std::future::Future;
use std::time::Duration;
use quicknode_sdk::errors::{HttpKind, SdkError};
const BASE_DELAY_MS: u64 = 500;
const MAX_DELAY_MS: u64 = 8_000;
pub async fn retrying<T, F, Fut>(max_retries: u32, f: F) -> Result<T, SdkError>
where
F: Fn() -> Fut,
Fut: Future<Output = Result<T, SdkError>>,
{
let mut attempt: u32 = 0;
loop {
match f().await {
Ok(v) => return Ok(v),
Err(e) if attempt < max_retries && is_retryable(&e) => {
tokio::time::sleep(delay_for(attempt)).await;
attempt += 1;
}
Err(e) => return Err(e),
}
}
}
fn is_retryable(e: &SdkError) -> bool {
match e {
SdkError::Api { status, .. } => matches!(status.as_u16(), 429 | 500 | 502 | 503 | 504),
SdkError::Http(_) => matches!(e.http_kind(), Some(HttpKind::Timeout | HttpKind::Connect)),
_ => false,
}
}
fn delay_for(attempt: u32) -> Duration {
let exp = attempt.min(10); let ceiling_ms = BASE_DELAY_MS.saturating_mul(1 << exp).min(MAX_DELAY_MS);
Duration::from_millis(fastrand::u64(0..=ceiling_ms))
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicU32, Ordering};
fn api_error(status: u16) -> SdkError {
SdkError::Api {
status: reqwest::StatusCode::from_u16(status).unwrap(),
body: String::new(),
}
}
#[test]
fn retryable_statuses() {
for s in [429, 500, 502, 503, 504] {
assert!(is_retryable(&api_error(s)), "{s} should be retryable");
}
for s in [400, 401, 403, 404, 422] {
assert!(!is_retryable(&api_error(s)), "{s} should not be retryable");
}
}
#[test]
fn delay_never_exceeds_cap() {
for attempt in 0..64 {
assert!(delay_for(attempt) <= Duration::from_millis(MAX_DELAY_MS));
}
}
#[tokio::test(start_paused = true)]
async fn retries_until_success() {
let calls = AtomicU32::new(0);
let result = retrying(3, || {
let n = calls.fetch_add(1, Ordering::SeqCst);
async move {
if n < 2 {
Err(api_error(429))
} else {
Ok("ok")
}
}
})
.await;
assert_eq!(result.unwrap(), "ok");
assert_eq!(calls.load(Ordering::SeqCst), 3);
}
#[tokio::test(start_paused = true)]
async fn exhausts_retries_and_returns_last_error() {
let calls = AtomicU32::new(0);
let result: Result<(), _> = retrying(2, || {
calls.fetch_add(1, Ordering::SeqCst);
async { Err(api_error(503)) }
})
.await;
assert!(result.is_err());
assert_eq!(calls.load(Ordering::SeqCst), 3); }
#[tokio::test(start_paused = true)]
async fn non_retryable_fails_on_first_attempt() {
let calls = AtomicU32::new(0);
let result: Result<(), _> = retrying(3, || {
calls.fetch_add(1, Ordering::SeqCst);
async { Err(api_error(404)) }
})
.await;
assert!(result.is_err());
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
#[tokio::test(start_paused = true)]
async fn zero_retries_means_single_attempt() {
let calls = AtomicU32::new(0);
let result: Result<(), _> = retrying(0, || {
calls.fetch_add(1, Ordering::SeqCst);
async { Err(api_error(429)) }
})
.await;
assert!(result.is_err());
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
}