use std::time::Duration;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum JitterPolicy {
None,
Full,
Equal,
RandomizedBand,
}
impl Default for JitterPolicy {
fn default() -> Self {
JitterPolicy::None
}
}
impl JitterPolicy {
#[must_use]
pub fn apply(&self, delay: Duration) -> Duration {
match self {
JitterPolicy::None => delay,
JitterPolicy::RandomizedBand | JitterPolicy::Full => self.full_jitter(delay),
JitterPolicy::Equal => self.equal_jitter(delay),
}
}
#[must_use]
pub fn apply_randomized_band(
&self,
lower: Duration,
upper_seed: Duration,
max: Duration,
) -> Duration {
let lower = lower.min(max);
if !matches!(self, JitterPolicy::RandomizedBand) {
return self.apply(lower);
}
let lower_ms = (lower.as_millis().min(u128::from(u64::MAX))) as u64;
let seed_ms = (upper_seed.as_millis().min(u128::from(u64::MAX))) as u64;
let max_ms = (max.as_millis().min(u128::from(u64::MAX))) as u64;
let upper = seed_ms.saturating_mul(3).min(max_ms).max(lower_ms);
if lower_ms >= upper {
return lower;
}
Duration::from_millis(fastrand::u64(lower_ms..=upper))
}
fn full_jitter(&self, delay: Duration) -> Duration {
let ns = (delay.as_nanos().min(u128::from(u64::MAX))) as u64;
if ns == 0 {
return Duration::ZERO;
}
Duration::from_nanos(fastrand::u64(0..=ns))
}
fn equal_jitter(&self, delay: Duration) -> Duration {
let ns = (delay.as_nanos().min(u128::from(u64::MAX))) as u64;
if ns == 0 {
return Duration::ZERO;
}
let half = ns / 2;
let jitter = if half == 0 {
0
} else {
fastrand::u64(0..=half)
};
Duration::from_nanos(half + jitter)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
#[test]
fn full_jitter_zero_duration_returns_zero() {
assert_eq!(JitterPolicy::Full.apply(Duration::ZERO), Duration::ZERO);
}
#[test]
fn equal_jitter_zero_duration_returns_zero() {
assert_eq!(JitterPolicy::Equal.apply(Duration::ZERO), Duration::ZERO);
}
#[test]
fn randomized_band_apply_falls_back_to_full_jitter() {
let delay = Duration::from_millis(500);
for _ in 0..100 {
let d = JitterPolicy::RandomizedBand.apply(delay);
assert!(
d <= delay,
"context-free RandomizedBand apply() must be ≤ delay"
);
}
}
#[test]
fn apply_randomized_band_on_full_falls_back_to_apply() {
let lower = Duration::from_millis(100);
let upper_seed = Duration::from_millis(300);
let max = Duration::from_secs(10);
for _ in 0..100 {
let result = JitterPolicy::Full.apply_randomized_band(lower, upper_seed, max);
assert!(result <= lower, "Full fallback should be in [0, lower]");
}
}
#[test]
fn apply_randomized_band_clamps_lower_to_max() {
let lower = Duration::from_secs(10);
let upper_seed = Duration::from_millis(50);
let max = Duration::from_secs(5);
let result = JitterPolicy::RandomizedBand.apply_randomized_band(lower, upper_seed, max);
assert_eq!(
result, max,
"result must be capped at max, not the out-of-band lower"
);
}
#[test]
fn apply_randomized_band_draws_within_first_to_3x_band() {
let lower = Duration::from_millis(100);
let upper_seed = Duration::from_secs(1);
let max = Duration::from_secs(30);
let upper = Duration::from_secs(3); for _ in 0..500 {
let d = JitterPolicy::RandomizedBand.apply_randomized_band(lower, upper_seed, max);
assert!(
d >= lower && d <= upper,
"draw {d:?} outside [{lower:?}, {upper:?}]"
);
}
}
#[test]
fn full_and_equal_jitter_preserve_sub_millisecond_delays() {
let d = Duration::from_micros(500);
let mut full_nonzero = false;
let mut equal_nonzero = false;
for _ in 0..1000 {
let f = JitterPolicy::Full.apply(d);
assert!(f <= d, "full jitter must stay within [0, delay]");
if f > Duration::ZERO {
full_nonzero = true;
}
let e = JitterPolicy::Equal.apply(d);
assert!(
e >= d / 2 && e <= d,
"equal jitter must stay within [delay/2, delay]"
);
if e > Duration::ZERO {
equal_nonzero = true;
}
}
assert!(
full_nonzero,
"sub-ms full jitter must not be systematically zero"
);
assert!(
equal_nonzero,
"sub-ms equal jitter must not be systematically zero"
);
}
}