use std::future::Future;
use std::time::{Duration, SystemTime};
use crate::{CancellationToken, TransportErrorClass};
use url::Url;
use crate::transport::policy::{NetworkErrorKind, RequestRateLimiter, RetryPolicy};
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RetrySignal {
HttpStatus {
status: u16,
headers: Vec<(String, String)>,
},
Transport {
class: TransportErrorClass,
},
}
#[non_exhaustive]
#[derive(Debug)]
pub enum AttemptOutcome<T, E> {
Success(T),
Failure {
error: E,
signal: RetrySignal,
},
}
#[non_exhaustive]
#[derive(Debug)]
pub enum LimiterKey<'a> {
Global,
PerUrl(&'a Url),
}
pub async fn run_with_retry<T, E, F, Fut>(
policy: &RetryPolicy,
rate_limiter: &RequestRateLimiter,
limiter_key: LimiterKey<'_>,
attempt: F,
) -> Result<T, E>
where
F: FnMut(usize) -> Fut,
Fut: Future<Output = AttemptOutcome<T, E>>,
{
run_with_retry_using(
policy,
rate_limiter,
limiter_key,
attempt,
crate::transport::policy::time::sleep,
crate::transport::policy::time::system_now,
)
.await
}
pub(crate) async fn run_with_retry_using<T, E, F, Fut, S, SFut, C>(
policy: &RetryPolicy,
rate_limiter: &RequestRateLimiter,
limiter_key: LimiterKey<'_>,
mut attempt: F,
mut sleeper: S,
clock: C,
) -> Result<T, E>
where
F: FnMut(usize) -> Fut,
Fut: Future<Output = AttemptOutcome<T, E>>,
S: FnMut(Duration) -> SFut,
SFut: Future<Output = ()>,
C: Fn() -> SystemTime,
{
let max_attempts = policy.max_attempts().max(1);
let mut attempt_index = 1usize;
loop {
let token = CancellationToken::new();
match &limiter_key {
LimiterKey::Global => {
let _ = rate_limiter.acquire_global(&token).await;
}
LimiterKey::PerUrl(url) => {
let _ = rate_limiter.acquire(url, &token).await;
}
}
match attempt(attempt_index).await {
AttemptOutcome::Success(value) => return Ok(value),
AttemptOutcome::Failure { error, signal } => {
let retryable = match &signal {
RetrySignal::HttpStatus { status, .. } => policy.should_retry_status(*status),
RetrySignal::Transport { class } => policy
.should_retry_network(NetworkErrorKind::from_transport_error_class(*class)),
};
if retryable && attempt_index < max_attempts {
let delay = match &signal {
RetrySignal::HttpStatus { status, headers } => {
policy.delay_for_status(attempt_index, *status, headers, clock())
}
RetrySignal::Transport { .. } => policy.delay_for_attempt(attempt_index),
};
emit_retry_event(attempt_index, &signal, delay);
sleeper(delay).await;
attempt_index += 1;
continue;
}
if retryable {
emit_exhausted_event(attempt_index, &signal);
}
return Err(error);
}
}
}
}
#[cfg(feature = "tracing")]
fn emit_retry_event(attempt_index: usize, signal: &RetrySignal, delay: Duration) {
let attempt_index = u64::try_from(attempt_index).unwrap_or(u64::MAX);
let backoff_ms = u64::try_from(delay.as_millis()).unwrap_or(u64::MAX);
match signal {
RetrySignal::HttpStatus { status, .. } => tracing::debug!(
target: "cow_sdk::transport",
attempt_index,
status = u64::from(*status),
backoff_ms,
"retry scheduled after status response"
),
RetrySignal::Transport { class } => tracing::debug!(
target: "cow_sdk::transport",
attempt_index,
transport_error_class = class.as_str(),
backoff_ms,
"retry scheduled after transport error"
),
}
}
#[cfg(not(feature = "tracing"))]
#[inline]
const fn emit_retry_event(_attempt_index: usize, _signal: &RetrySignal, _delay: Duration) {}
#[cfg(feature = "tracing")]
fn emit_exhausted_event(attempt_index: usize, signal: &RetrySignal) {
let attempt_index = u64::try_from(attempt_index).unwrap_or(u64::MAX);
match signal {
RetrySignal::HttpStatus { status, .. } => tracing::warn!(
target: "cow_sdk::transport",
attempt_index,
status = u64::from(*status),
backoff_ms = 0_u64,
"retry attempts exhausted after status response"
),
RetrySignal::Transport { class } => tracing::warn!(
target: "cow_sdk::transport",
attempt_index,
transport_error_class = class.as_str(),
backoff_ms = 0_u64,
"retry attempts exhausted after transport error"
),
}
}
#[cfg(not(feature = "tracing"))]
#[inline]
const fn emit_exhausted_event(_attempt_index: usize, _signal: &RetrySignal) {}
#[cfg(test)]
mod tests {
use std::cell::RefCell;
use std::rc::Rc;
use std::time::{Duration, SystemTime};
use crate::TransportErrorClass;
use super::{AttemptOutcome, LimiterKey, RetrySignal, run_with_retry_using};
use crate::transport::policy::{JitterStrategy, RequestRateLimiter, RetryPolicy};
const NOW_SECS: u64 = 1_000_000;
#[derive(Clone, Copy)]
enum Raw {
Ok,
Status {
status: u16,
retry_after: RetryAfter,
},
Transport(TransportErrorClass),
}
#[derive(Clone, Copy)]
enum RetryAfter {
None,
DeltaSecs(u64),
HttpDateAtSecs(u64),
}
#[derive(Debug, PartialEq, Eq)]
enum Outcome {
Success,
Status(u16),
Transport(TransportErrorClass),
}
struct RunReport {
outcome: Outcome,
sleeps_ms: Vec<u64>,
dispatches: usize,
}
fn fixed_clock() -> SystemTime {
SystemTime::UNIX_EPOCH + Duration::from_secs(NOW_SECS)
}
fn policy(max_attempts: usize) -> RetryPolicy {
RetryPolicy::builder()
.max_attempts(max_attempts)
.jitter(JitterStrategy::none())
.build()
}
fn header_pairs(retry_after: RetryAfter) -> Vec<(String, String)> {
match retry_after {
RetryAfter::None => Vec::new(),
RetryAfter::DeltaSecs(secs) => vec![("retry-after".to_owned(), secs.to_string())],
RetryAfter::HttpDateAtSecs(at) => {
let when = SystemTime::UNIX_EPOCH + Duration::from_secs(at);
vec![("retry-after".to_owned(), httpdate::fmt_http_date(when))]
}
}
}
async fn run(script: &[Raw], max_attempts: usize) -> RunReport {
let sleeps: Rc<RefCell<Vec<u64>>> = Rc::new(RefCell::new(Vec::new()));
let dispatches = Rc::new(RefCell::new(0usize));
let policy = policy(max_attempts);
let limiter = RequestRateLimiter::unlimited();
let script = script.to_vec();
let sleeps_for_sleeper = Rc::clone(&sleeps);
let dispatches_for_attempt = Rc::clone(&dispatches);
let outcome = run_with_retry_using::<Outcome, Outcome, _, _, _, _, _>(
&policy,
&limiter,
LimiterKey::Global,
move |attempt_index| {
let idx = (attempt_index - 1).min(script.len() - 1);
let raw = script[idx];
*dispatches_for_attempt.borrow_mut() += 1;
async move {
match raw {
Raw::Ok => AttemptOutcome::Success(Outcome::Success),
Raw::Status {
status,
retry_after,
} => AttemptOutcome::Failure {
error: Outcome::Status(status),
signal: RetrySignal::HttpStatus {
status,
headers: header_pairs(retry_after),
},
},
Raw::Transport(class) => AttemptOutcome::Failure {
error: Outcome::Transport(class),
signal: RetrySignal::Transport { class },
},
}
}
},
move |delay: Duration| {
sleeps_for_sleeper
.borrow_mut()
.push(u64::try_from(delay.as_millis()).unwrap_or(u64::MAX));
async {}
},
fixed_clock,
)
.await;
let outcome = match outcome {
Ok(value) | Err(value) => value,
};
RunReport {
outcome,
sleeps_ms: Rc::try_unwrap(sleeps).unwrap().into_inner(),
dispatches: *dispatches.borrow(),
}
}
#[tokio::test]
async fn immediate_success_does_not_sleep() {
let report = run(&[Raw::Ok], 10).await;
assert_eq!(report.outcome, Outcome::Success);
assert_eq!(report.dispatches, 1);
assert!(report.sleeps_ms.is_empty());
}
#[tokio::test]
async fn retryable_status_then_success_backs_off_once() {
let report = run(
&[
Raw::Status {
status: 503,
retry_after: RetryAfter::None,
},
Raw::Ok,
],
10,
)
.await;
assert_eq!(report.outcome, Outcome::Success);
assert_eq!(report.dispatches, 2);
assert_eq!(report.sleeps_ms, vec![50]);
}
#[tokio::test]
async fn delta_retry_after_overrides_backoff_floor() {
let report = run(
&[
Raw::Status {
status: 429,
retry_after: RetryAfter::DeltaSecs(2),
},
Raw::Ok,
],
10,
)
.await;
assert_eq!(report.outcome, Outcome::Success);
assert_eq!(report.sleeps_ms, vec![2000]);
}
#[tokio::test]
async fn http_date_retry_after_uses_the_injected_clock() {
let report = run(
&[
Raw::Status {
status: 503,
retry_after: RetryAfter::HttpDateAtSecs(NOW_SECS + 10),
},
Raw::Ok,
],
10,
)
.await;
assert_eq!(report.outcome, Outcome::Success);
assert_eq!(report.sleeps_ms, vec![10_000]);
}
#[tokio::test]
async fn persistent_retryable_status_exhausts_attempts() {
let report = run(
&[Raw::Status {
status: 500,
retry_after: RetryAfter::None,
}],
10,
)
.await;
assert_eq!(report.outcome, Outcome::Status(500));
assert_eq!(report.dispatches, 10);
assert_eq!(
report.sleeps_ms,
vec![50, 100, 200, 400, 800, 1600, 3200, 3200, 3200]
);
}
#[tokio::test]
async fn persistent_transport_error_exhausts_attempts() {
let report = run(&[Raw::Transport(TransportErrorClass::Timeout)], 10).await;
assert_eq!(
report.outcome,
Outcome::Transport(TransportErrorClass::Timeout)
);
assert_eq!(report.dispatches, 10);
assert_eq!(
report.sleeps_ms,
vec![50, 100, 200, 400, 800, 1600, 3200, 3200, 3200]
);
}
#[tokio::test]
async fn non_retryable_status_returns_immediately() {
let report = run(
&[Raw::Status {
status: 400,
retry_after: RetryAfter::None,
}],
10,
)
.await;
assert_eq!(report.outcome, Outcome::Status(400));
assert_eq!(report.dispatches, 1);
assert!(report.sleeps_ms.is_empty());
}
#[tokio::test]
async fn non_retryable_transport_returns_without_redispatch() {
for class in [TransportErrorClass::Decode, TransportErrorClass::Builder] {
let report = run(&[Raw::Transport(class)], 10).await;
assert_eq!(report.outcome, Outcome::Transport(class));
assert_eq!(report.dispatches, 1, "class {class:?} must dispatch once");
assert!(report.sleeps_ms.is_empty());
}
}
#[tokio::test]
async fn mixed_transport_then_status_then_success() {
let report = run(
&[
Raw::Transport(TransportErrorClass::Timeout),
Raw::Status {
status: 429,
retry_after: RetryAfter::DeltaSecs(1),
},
Raw::Ok,
],
10,
)
.await;
assert_eq!(report.outcome, Outcome::Success);
assert_eq!(report.dispatches, 3);
assert_eq!(report.sleeps_ms, vec![50, 1000]);
}
#[tokio::test]
async fn no_retry_policy_makes_one_attempt() {
let report = run(
&[Raw::Status {
status: 503,
retry_after: RetryAfter::None,
}],
1,
)
.await;
assert_eq!(report.outcome, Outcome::Status(503));
assert_eq!(report.dispatches, 1);
assert!(report.sleeps_ms.is_empty());
}
}