use std::sync::Arc;
use std::time::Duration;
use tokio::sync::broadcast;
use weida_core::{Fingerprint, LossCause, PeerIdentity};
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ReconnectPolicy {
pub initial: Duration,
pub max: Duration,
pub jitter: bool,
pub max_attempts: Option<u32>,
pub stop_on_peer_closed: bool,
}
impl Default for ReconnectPolicy {
fn default() -> Self {
ReconnectPolicy {
initial: Duration::from_millis(100),
max: Duration::from_secs(30),
jitter: true,
max_attempts: None,
stop_on_peer_closed: false,
}
}
}
impl ReconnectPolicy {
#[must_use]
pub fn never() -> Self {
ReconnectPolicy {
max_attempts: Some(0),
..ReconnectPolicy::default()
}
}
#[must_use]
pub fn redials(&self) -> bool {
self.max_attempts != Some(0)
}
pub(crate) fn redials_after(&self, cause: LossCause) -> bool {
self.redials()
&& cause != LossCause::LocallyClosed
&& !(cause == LossCause::PeerClosed && self.stop_on_peer_closed)
}
pub(crate) fn allows(&self, attempt: u32) -> bool {
self.max_attempts.is_none_or(|max| attempt <= max)
}
#[must_use]
pub fn delay(&self, attempt: u32) -> Duration {
let doublings = attempt.saturating_sub(1).min(32);
let base = self
.initial
.checked_mul(1u32 << doublings.min(31))
.unwrap_or(self.max)
.min(self.max);
if !self.jitter || base.is_zero() {
return base;
}
let half = base / 2;
half + rand::random_range(Duration::ZERO..=half)
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum OutboxFull {
#[default]
Block,
Drop,
Reject,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum GiveUp {
Policy {
attempts: u32,
},
PeerChanged {
presented: Option<Fingerprint>,
},
Failed(String),
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum PeerEvent {
Connected {
url: Arc<str>,
peer: Option<PeerIdentity>,
},
Lost {
url: Arc<str>,
cause: LossCause,
},
Retrying {
url: Arc<str>,
attempt: u32,
delay: Duration,
},
GaveUp {
url: Arc<str>,
why: GiveUp,
},
Missed(u64),
}
pub(crate) const EVENT_QUEUE: usize = 64;
pub struct PeerEvents {
rx: broadcast::Receiver<PeerEvent>,
}
impl PeerEvents {
pub(crate) fn new(rx: broadcast::Receiver<PeerEvent>) -> PeerEvents {
PeerEvents { rx }
}
pub async fn recv(&mut self) -> Option<PeerEvent> {
match self.rx.recv().await {
Ok(event) => Some(event),
Err(broadcast::error::RecvError::Lagged(n)) => Some(PeerEvent::Missed(n)),
Err(broadcast::error::RecvError::Closed) => None,
}
}
}
impl std::fmt::Debug for PeerEvents {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PeerEvents").finish_non_exhaustive()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_delay_doubles_and_is_capped() {
let policy = ReconnectPolicy {
initial: Duration::from_millis(100),
max: Duration::from_millis(350),
jitter: false,
..ReconnectPolicy::default()
};
assert_eq!(policy.delay(1), Duration::from_millis(100));
assert_eq!(policy.delay(2), Duration::from_millis(200));
assert_eq!(policy.delay(3), Duration::from_millis(350));
assert_eq!(policy.delay(40), Duration::from_millis(350));
}
#[test]
fn jitter_stays_in_the_upper_half() {
let policy = ReconnectPolicy {
initial: Duration::from_millis(100),
..ReconnectPolicy::default()
};
for _ in 0..1000 {
let d = policy.delay(2);
assert!(
d >= Duration::from_millis(100) && d <= Duration::from_millis(200),
"{d:?}"
);
}
}
#[test]
fn never_means_no_attempt() {
assert!(!ReconnectPolicy::never().redials());
assert!(!ReconnectPolicy::never().allows(1));
assert!(ReconnectPolicy::default().allows(1_000_000));
assert!(
ReconnectPolicy {
max_attempts: Some(3),
..ReconnectPolicy::default()
}
.allows(3)
);
assert!(
!ReconnectPolicy {
max_attempts: Some(3),
..ReconnectPolicy::default()
}
.allows(4)
);
}
}