use std::time::Duration;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct BackoffPolicy {
pub initial: Duration,
pub max: Duration,
pub multiplier: u32,
}
impl Default for BackoffPolicy {
fn default() -> Self {
Self {
initial: Duration::from_millis(500),
max: Duration::from_secs(30),
multiplier: 2,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum InvalidBackoff {
#[error("backoff.initial must be greater than zero")]
ZeroInitial,
#[error("backoff.multiplier must be at least 2; a lower factor never backs off")]
Multiplier { multiplier: u32 },
#[error("backoff.max ({max:?}) must be at least backoff.initial ({initial:?})")]
Ceiling { initial: Duration, max: Duration },
}
impl BackoffPolicy {
pub const fn validate(&self) -> Result<(), InvalidBackoff> {
if self.initial.is_zero() {
return Err(InvalidBackoff::ZeroInitial);
}
if self.multiplier < 2 {
return Err(InvalidBackoff::Multiplier {
multiplier: self.multiplier,
});
}
if self.max.as_nanos() < self.initial.as_nanos() {
return Err(InvalidBackoff::Ceiling {
initial: self.initial,
max: self.max,
});
}
Ok(())
}
pub fn delay(&self, failures: u32) -> Duration {
let Some(exponent) = failures.checked_sub(1) else {
return Duration::ZERO;
};
let factor = u64::from(self.multiplier)
.checked_pow(exponent.min(u32::from(u8::MAX)))
.unwrap_or(u64::MAX);
self.initial
.checked_mul(u32::try_from(factor).unwrap_or(u32::MAX))
.unwrap_or(self.max)
.min(self.max)
}
}
#[derive(Debug, Clone, Copy)]
pub struct Backoff {
policy: BackoffPolicy,
failures: u32,
}
impl Backoff {
pub const fn new(policy: BackoffPolicy) -> Self {
Self {
policy,
failures: 0,
}
}
pub fn fail(&mut self) -> Duration {
self.failures = self.failures.saturating_add(1);
self.policy.delay(self.failures)
}
pub const fn succeed(&mut self) {
self.failures = 0;
}
pub const fn failures(&self) -> u32 {
self.failures
}
pub fn delay(&self) -> Duration {
self.policy.delay(self.failures)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn policy() -> BackoffPolicy {
BackoffPolicy {
initial: Duration::from_millis(100),
max: Duration::from_secs(1),
multiplier: 2,
}
}
#[test]
fn consecutive_failures_double_the_delay_up_to_the_ceiling_and_stop_there() {
let mut backoff = Backoff::new(policy());
assert_eq!(backoff.delay(), Duration::ZERO);
let observed: Vec<Duration> = (0..6).map(|_| backoff.fail()).collect();
assert_eq!(
observed,
vec![
Duration::from_millis(100),
Duration::from_millis(200),
Duration::from_millis(400),
Duration::from_millis(800),
Duration::from_secs(1),
Duration::from_secs(1),
]
);
assert_eq!(backoff.failures(), 6);
}
#[test]
fn an_unbounded_failure_count_saturates_at_the_ceiling() {
let policy = policy();
for failures in [32u32, 1_000, u32::MAX] {
assert_eq!(policy.delay(failures), policy.max, "{failures} failures");
}
}
#[test]
fn one_success_clears_the_outage() {
let mut backoff = Backoff::new(policy());
backoff.fail();
backoff.fail();
backoff.succeed();
assert_eq!(backoff.failures(), 0);
assert_eq!(backoff.delay(), Duration::ZERO);
assert_eq!(backoff.fail(), policy().initial);
}
#[test]
fn a_policy_that_would_hot_loop_or_never_back_off_is_refused() {
assert_eq!(
BackoffPolicy {
initial: Duration::ZERO,
..policy()
}
.validate(),
Err(InvalidBackoff::ZeroInitial)
);
assert_eq!(
BackoffPolicy {
multiplier: 1,
..policy()
}
.validate(),
Err(InvalidBackoff::Multiplier { multiplier: 1 })
);
assert_eq!(
BackoffPolicy {
initial: Duration::from_secs(5),
max: Duration::from_secs(1),
..policy()
}
.validate(),
Err(InvalidBackoff::Ceiling {
initial: Duration::from_secs(5),
max: Duration::from_secs(1),
})
);
assert_eq!(BackoffPolicy::default().validate(), Ok(()));
}
}