use async_trait::async_trait;
use std::any::Any;
use std::time::Duration;
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum RetryBackoff {
Immediate,
Fixed {
delay: Duration,
},
Exponential {
initial_delay: Duration,
max_delay: Duration,
},
ExponentialWithJitter {
initial_delay: Duration,
max_delay: Duration,
max_jitter: Duration,
},
}
impl RetryBackoff {
pub(crate) fn delay(&self, job_key: &str, retry_number: usize) -> Duration {
match self {
Self::Immediate => Duration::ZERO,
Self::Fixed { delay } => *delay,
Self::Exponential {
initial_delay,
max_delay,
} => exponential_delay(*initial_delay, *max_delay, retry_number),
Self::ExponentialWithJitter {
initial_delay,
max_delay,
max_jitter,
} => exponential_delay(*initial_delay, *max_delay, retry_number)
.saturating_add(jitter(job_key, retry_number, *max_jitter))
.min(*max_delay),
}
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct RetryPolicy {
max_retries: usize,
backoff: RetryBackoff,
}
impl RetryPolicy {
pub fn new(max_retries: usize, backoff: RetryBackoff) -> Self {
Self {
max_retries,
backoff,
}
}
pub fn immediate(max_retries: usize) -> Self {
Self::new(max_retries, RetryBackoff::Immediate)
}
pub fn no_retries() -> Self {
Self::immediate(0)
}
pub fn fixed(max_retries: usize, delay: Duration) -> Self {
Self::new(max_retries, RetryBackoff::Fixed { delay })
}
pub fn exponential(max_retries: usize, initial_delay: Duration, max_delay: Duration) -> Self {
Self::new(
max_retries,
RetryBackoff::Exponential {
initial_delay,
max_delay,
},
)
}
pub fn exponential_with_jitter(
max_retries: usize,
initial_delay: Duration,
max_delay: Duration,
max_jitter: Duration,
) -> Self {
Self::new(
max_retries,
RetryBackoff::ExponentialWithJitter {
initial_delay,
max_delay,
max_jitter,
},
)
}
pub fn max_retries(&self) -> usize {
self.max_retries
}
pub fn backoff(&self) -> &RetryBackoff {
&self.backoff
}
}
impl Default for RetryPolicy {
fn default() -> Self {
Self::exponential_with_jitter(
6,
Duration::from_secs(1),
Duration::from_secs(60),
Duration::from_millis(500),
)
}
}
#[async_trait]
pub trait JobRetryPolicy: Sync {
type Context: Sync + 'static;
async fn retry_policy(&self, context: &Self::Context) -> anyhow::Result<RetryPolicy>;
}
#[doc(hidden)]
pub struct RetryPolicyResolver<'a, T, C> {
message: &'a T,
context: &'a C,
}
impl<'a, T, C> RetryPolicyResolver<'a, T, C> {
#[doc(hidden)]
pub fn new(message: &'a T, context: &'a C) -> Self {
Self { message, context }
}
}
#[doc(hidden)]
#[async_trait]
pub trait ResolveRetryPolicy {
#[doc(hidden)]
async fn resolve_retry_policy(self) -> anyhow::Result<Option<RetryPolicy>>;
}
#[async_trait]
impl<T: Sync, C: Sync> ResolveRetryPolicy for &&RetryPolicyResolver<'_, T, C> {
async fn resolve_retry_policy(self) -> anyhow::Result<Option<RetryPolicy>> {
Ok(None)
}
}
#[async_trait]
impl<T, C> ResolveRetryPolicy for &RetryPolicyResolver<'_, T, C>
where
T: JobRetryPolicy + Sync,
C: Any + Sync,
{
async fn resolve_retry_policy(self) -> anyhow::Result<Option<RetryPolicy>> {
let context = (self.context as &dyn Any)
.downcast_ref::<T::Context>()
.ok_or_else(|| {
anyhow::anyhow!(
"retry policy for {} expects application context {}",
std::any::type_name::<T>(),
std::any::type_name::<T::Context>(),
)
})?;
JobRetryPolicy::retry_policy(self.message, context)
.await
.map(Some)
}
}
fn exponential_delay(initial: Duration, maximum: Duration, retry_number: usize) -> Duration {
let exponent = u32::try_from(retry_number.saturating_sub(1)).unwrap_or(u32::MAX);
let multiplier = 2_u32.checked_pow(exponent).unwrap_or(u32::MAX);
initial.saturating_mul(multiplier).min(maximum)
}
fn jitter(job_key: &str, retry_number: usize, maximum: Duration) -> Duration {
let mut input = job_key.as_bytes().to_vec();
input.extend_from_slice(
&u64::try_from(retry_number)
.unwrap_or(u64::MAX)
.to_le_bytes(),
);
let hash = blake3::hash(&input);
let mut bytes = [0_u8; 8];
bytes.copy_from_slice(&hash.as_bytes()[..8]);
let value = u64::from_le_bytes(bytes);
let maximum_nanos = u64::try_from(maximum.as_nanos()).unwrap_or(u64::MAX);
let nanos = value % maximum_nanos.saturating_add(1);
Duration::from_nanos(nanos)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn backoff_presets_apply_limits_and_bounded_jitter() {
let exponential =
RetryPolicy::exponential(4, Duration::from_secs(2), Duration::from_secs(5));
let jittered = RetryPolicy::exponential_with_jitter(
4,
Duration::from_secs(2),
Duration::from_secs(5),
Duration::from_secs(1),
);
assert_eq!(
vec![
exponential.backoff.delay("job", 1),
exponential.backoff.delay("job", 2),
exponential.backoff.delay("job", 3),
RetryPolicy::fixed(1, Duration::from_secs(3))
.backoff
.delay("job", 1),
RetryPolicy::immediate(1).backoff.delay("job", 1),
],
vec![
Duration::from_secs(2),
Duration::from_secs(4),
Duration::from_secs(5),
Duration::from_secs(3),
Duration::ZERO,
]
);
let first = jittered.backoff.delay("job-a", 1);
assert!(first >= Duration::from_secs(2));
assert!(first <= Duration::from_secs(3));
assert_eq!(first, jittered.backoff.delay("job-a", 1));
assert_ne!(first, jittered.backoff.delay("job-b", 1));
assert!(jittered.backoff.delay("job-a", 20) <= Duration::from_secs(5));
assert_eq!(
exponential.backoff.delay("job", usize::MAX),
Duration::from_secs(5)
);
assert_eq!(RetryPolicy::no_retries().max_retries(), 0);
}
}