use std::time::Duration;
#[derive(Debug, Clone)]
pub struct Backoff {
base: Duration,
jittered: Duration,
}
impl Backoff {
#[inline]
pub async fn sleep(&self) {
re_async::sleep(self.jittered).await;
}
#[inline]
pub fn base(&self) -> Duration {
self.base
}
#[inline]
pub fn jittered(&self) -> Duration {
self.jittered
}
}
#[derive(Debug)]
pub struct BackoffGenerator {
base: Duration,
max: Duration,
jitter_factor: f64,
iteration: u32,
}
impl BackoffGenerator {
pub const DEFAULT_JITTER_FACTOR: f64 = 1.0;
pub fn new(base: Duration, max: Duration) -> Result<Self, String> {
Self::new_with_custom_jitter(base, max, Self::DEFAULT_JITTER_FACTOR)
}
pub fn new_with_custom_jitter(
base: Duration,
max: Duration,
jitter_factor: f64,
) -> Result<Self, String> {
if base > max {
return Err("base duration must be less than or equal to max duration".to_owned());
}
if jitter_factor < 0.0 || jitter_factor > 1.0 {
return Err("jitter factor must be between 0 and 1".to_owned());
}
Ok(Self {
base,
max,
jitter_factor,
iteration: 0,
})
}
fn jitter(&self, duration: Duration) -> Duration {
let rand = rand::random::<f64>(); let factor = (1.0 - self.jitter_factor) + self.jitter_factor * rand; let jittered_secs = duration.as_secs_f64() * factor;
Duration::try_from_secs_f64(jittered_secs).unwrap_or(duration)
}
pub fn gen_next(&mut self) -> Backoff {
let base = 2u32
.checked_pow(self.iteration)
.and_then(|p| self.base.checked_mul(p))
.unwrap_or(self.max)
.clamp(self.base, self.max);
let jittered = self.jitter(base);
self.iteration += 1;
Backoff { base, jittered }
}
pub fn max_backoff(&self) -> Backoff {
let jittered = self.jitter(self.max);
Backoff {
base: self.max,
jittered,
}
}
pub fn is_reset(&self) -> bool {
self.iteration == 0
}
pub fn reset(&mut self) {
self.iteration = 0;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn full_jitter_stays_within_zero_and_base() {
let mut generator =
BackoffGenerator::new(Duration::from_secs(1), Duration::from_secs(8)).unwrap();
let expected_bases = [1, 2, 4, 8, 8, 8];
for expected in expected_bases {
for _ in 0..100 {
let mut g = BackoffGenerator::new(generator.base, generator.max).unwrap();
g.iteration = generator.iteration;
let b = g.gen_next();
assert_eq!(b.base(), Duration::from_secs(expected));
assert!(b.jittered() <= b.base());
}
generator.gen_next();
}
}
#[test]
fn zero_jitter_factor_yields_exactly_base() {
let mut generator = BackoffGenerator::new_with_custom_jitter(
Duration::from_millis(100),
Duration::from_secs(1),
0.0,
)
.unwrap();
for _ in 0..100 {
let b = generator.gen_next();
assert_eq!(b.jittered(), b.base());
generator.reset();
}
}
}