use crate::transport::{TransportError, TransportSink, TransportStream};
use alloc::boxed::Box;
use core::future::Future;
use core::pin::Pin;
use core::time::Duration;
pub trait Reconnector: Send + Sync + 'static {
#[allow(clippy::type_complexity)]
fn connect<'a>(
&'a self,
) -> Pin<
Box<
dyn Future<
Output = Result<
(Box<dyn TransportSink>, Box<dyn TransportStream>),
TransportError,
>,
> + Send
+ 'a,
>,
>;
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ReconnectPolicy {
pub initial_delay: Duration,
pub max_delay: Duration,
pub multiplier: u32,
pub jitter: bool,
}
impl Default for ReconnectPolicy {
fn default() -> Self {
Self {
initial_delay: Duration::from_secs(1),
max_delay: Duration::from_secs(60),
multiplier: 2,
jitter: true,
}
}
}
impl ReconnectPolicy {
pub(crate) fn delay_for(&self, attempt: u32) -> Duration {
let mut delay = self.initial_delay;
for _ in 0..attempt {
delay = match delay.checked_mul(self.multiplier) {
Some(d) if d < self.max_delay => d,
_ => return self.max_delay,
};
}
delay
}
pub(crate) fn jittered_delay_for(&self, attempt: u32) -> Duration {
let delay = self.delay_for(attempt);
if !self.jitter {
return delay;
}
let half = delay / 2;
let fraction = uuid::Uuid::new_v4().as_u128() as u32;
let extra = (half.as_nanos() * fraction as u128) / u32::MAX as u128;
half + Duration::from_nanos(extra as u64)
}
}
#[derive(Debug, Clone, Copy)]
pub enum ReconnectBehavior {
Enabled(ReconnectPolicy),
Disabled,
}
impl Default for ReconnectBehavior {
fn default() -> Self {
Self::Enabled(ReconnectPolicy::default())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn delay_doubles_and_caps() {
let policy = ReconnectPolicy {
initial_delay: Duration::from_secs(1),
max_delay: Duration::from_secs(10),
multiplier: 2,
jitter: false,
};
assert_eq!(policy.delay_for(0), Duration::from_secs(1));
assert_eq!(policy.delay_for(1), Duration::from_secs(2));
assert_eq!(policy.delay_for(2), Duration::from_secs(4));
assert_eq!(policy.delay_for(3), Duration::from_secs(8));
assert_eq!(policy.delay_for(4), Duration::from_secs(10));
assert_eq!(policy.delay_for(10), Duration::from_secs(10));
}
#[test]
fn jitter_off_is_exact() {
let policy = ReconnectPolicy {
initial_delay: Duration::from_secs(4),
jitter: false,
..ReconnectPolicy::default()
};
for attempt in 0..6 {
assert_eq!(
policy.jittered_delay_for(attempt),
policy.delay_for(attempt)
);
}
}
#[test]
fn jitter_stays_within_half_the_delay_and_the_full_delay() {
let policy = ReconnectPolicy {
initial_delay: Duration::from_secs(8),
max_delay: Duration::from_secs(64),
multiplier: 2,
jitter: true,
};
for attempt in 0..6 {
let full = policy.delay_for(attempt);
for _ in 0..200 {
let jittered = policy.jittered_delay_for(attempt);
assert!(
jittered >= full / 2 && jittered <= full,
"attempt {attempt}: {jittered:?} outside [{:?}, {full:?}]",
full / 2
);
}
}
}
#[test]
fn jitter_actually_varies() {
let policy = ReconnectPolicy::default();
let first = policy.jittered_delay_for(3);
let varies = (0..50).any(|_| policy.jittered_delay_for(3) != first);
assert!(varies, "jittered delays should not all be identical");
}
#[test]
fn a_zero_delay_survives_jittering() {
let policy = ReconnectPolicy {
initial_delay: Duration::ZERO,
max_delay: Duration::ZERO,
multiplier: 2,
jitter: true,
};
assert_eq!(policy.jittered_delay_for(0), Duration::ZERO);
}
}