use derive_builder::Builder;
use parking_lot::Mutex;
use std::{
sync::Arc,
time::{Duration, SystemTime},
};
use rand::{RngExt as _, rngs::StdRng};
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
pub struct RetryContext {
attempt: usize,
retry_after: Option<Duration>,
}
impl RetryContext {
pub const fn new(attempt: usize) -> Self {
Self {
attempt,
retry_after: None,
}
}
pub const fn with_retry_after(
mut self,
retry_after: Option<Duration>,
) -> Self {
self.retry_after = retry_after;
self
}
pub const fn attempt(&self) -> usize {
self.attempt
}
pub const fn retry_after(&self) -> Option<Duration> {
self.retry_after
}
}
pub fn parse_retry_after(value: &str) -> Option<Duration> {
parse_retry_after_at(value, SystemTime::now())
}
fn parse_retry_after_at(value: &str, now: SystemTime) -> Option<Duration> {
let value = value.trim();
if value.is_empty() {
return None;
}
if let Ok(seconds) = value.parse::<u64>() {
return Some(Duration::from_secs(seconds));
}
httpdate::parse_http_date(value)
.ok()
.map(|retry_at| retry_at.duration_since(now).unwrap_or_default())
}
pub trait RetryStrategy: Send + Sync {
fn should_retry_after(&mut self, context: RetryContext)
-> Option<Duration>;
}
#[derive(Builder)]
#[builder(pattern = "owned", setter(into))]
pub struct JitteredBackoff {
#[builder(default = 5)]
max_retry: usize,
#[builder(default = Arc::new(Mutex::new(rand::make_rng())))]
rng: Arc<Mutex<StdRng>>,
}
impl RetryStrategy for JitteredBackoff {
fn should_retry_after(
&mut self,
context: RetryContext,
) -> Option<Duration> {
if context.attempt() >= self.max_retry {
return None;
}
context
.retry_after()
.or_else(|| self.backoff_delay(context.attempt()))
}
}
impl JitteredBackoff {
fn backoff_delay(&mut self, retry_count: usize) -> Option<Duration> {
if retry_count >= self.max_retry {
return None;
}
let mut guard = self.rng.lock();
let jitter_ms = guard.random_range(0..1000);
let base_ms = 2_u64
.saturating_pow(retry_count as u32)
.saturating_mul(1_000);
let wait_ms = base_ms.saturating_sub(jitter_ms);
Some(Duration::from_millis(wait_ms))
}
}
impl Default for JitteredBackoff {
fn default() -> Self {
JitteredBackoffBuilder::default()
.build()
.expect("Builder defaults are valid")
}
}
pub struct NeverRetry;
impl RetryStrategy for NeverRetry {
fn should_retry_after(&mut self, _: RetryContext) -> Option<Duration> {
None
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use crate::retry_strategy::{
JitteredBackoffBuilder, RetryContext, RetryStrategy as _,
parse_retry_after, parse_retry_after_at,
};
#[test]
fn parses_retry_after_seconds() {
assert_eq!(parse_retry_after(" 42 "), Some(Duration::from_secs(42)));
}
#[test]
fn parses_retry_after_http_date() {
let now = httpdate::parse_http_date("Wed, 21 Oct 2015 07:27:00 GMT")
.expect("valid HTTP date");
assert_eq!(
parse_retry_after_at("Wed, 21 Oct 2015 07:28:00 GMT", now),
Some(Duration::from_secs(60))
);
}
#[test]
fn past_retry_after_http_date_retries_immediately() {
let now = httpdate::parse_http_date("Wed, 21 Oct 2015 07:29:00 GMT")
.expect("valid HTTP date");
assert_eq!(
parse_retry_after_at("Wed, 21 Oct 2015 07:28:00 GMT", now),
Some(Duration::ZERO)
);
}
#[test]
fn invalid_retry_after_is_ignored() {
assert_eq!(parse_retry_after("later"), None);
}
#[test]
fn jittered_backoff_prefers_retry_after_header() {
let mut backoff = JitteredBackoffBuilder::default()
.max_retry(3_usize)
.build()
.expect("valid retry strategy");
assert_eq!(
backoff.should_retry_after(
RetryContext::new(0)
.with_retry_after(Some(Duration::from_secs(25))),
),
Some(Duration::from_secs(25))
);
}
#[test]
fn jittered_backoff_still_honors_retry_limit_with_retry_after() {
let mut backoff = JitteredBackoffBuilder::default()
.max_retry(1_usize)
.build()
.expect("valid retry strategy");
assert_eq!(
backoff.should_retry_after(
RetryContext::new(1)
.with_retry_after(Some(Duration::from_secs(25))),
),
None
);
}
}