#![allow(deprecated)]
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::RwLock;
#[derive(Debug, Clone, PartialEq)]
pub enum ReconnectionState {
Connected,
Reconnecting {
attempt: u32,
next_delay: Duration,
},
Reconnected {
total_attempts: u32,
},
Failed {
total_attempts: u32,
},
Disabled,
}
#[deprecated(
since = "0.2.0",
note = "Use `ReconnectionSettings` with the `#[settings]` macro instead."
)]
#[derive(Debug, Clone)]
pub struct ReconnectionConfig {
pub max_attempts: Option<u32>,
pub initial_delay: Duration,
pub max_delay: Duration,
pub backoff_multiplier: f64,
pub jitter_factor: f64,
}
impl Default for ReconnectionConfig {
fn default() -> Self {
Self {
max_attempts: Some(10),
initial_delay: Duration::from_secs(1),
max_delay: Duration::from_secs(300), backoff_multiplier: 2.0,
jitter_factor: 0.1,
}
}
}
impl ReconnectionConfig {
pub fn new(max_attempts: Option<u32>, initial_delay: Duration, max_delay: Duration) -> Self {
Self {
max_attempts,
initial_delay,
max_delay,
backoff_multiplier: 2.0,
jitter_factor: 0.1,
}
}
pub fn with_max_attempts(mut self, max_attempts: u32) -> Self {
self.max_attempts = Some(max_attempts);
self
}
pub fn with_unlimited_attempts(mut self) -> Self {
self.max_attempts = None;
self
}
pub fn with_initial_delay(mut self, delay: Duration) -> Self {
self.initial_delay = delay;
self
}
pub fn with_max_delay(mut self, delay: Duration) -> Self {
self.max_delay = delay;
self
}
pub fn with_backoff_multiplier(mut self, multiplier: f64) -> Self {
self.backoff_multiplier = multiplier;
self
}
pub fn with_jitter_factor(mut self, factor: f64) -> Self {
self.jitter_factor = factor;
self
}
}
pub struct ReconnectionStrategy {
config: ReconnectionConfig,
current_attempt: u32,
current_delay: Duration,
}
impl ReconnectionStrategy {
pub fn new(config: ReconnectionConfig) -> Self {
let current_delay = config.initial_delay;
Self {
config,
current_attempt: 0,
current_delay,
}
}
pub fn attempt_count(&self) -> u32 {
self.current_attempt
}
pub fn next_delay(&mut self) -> Option<Duration> {
if let Some(max) = self.config.max_attempts
&& self.current_attempt >= max
{
return None;
}
let delay = if self.current_attempt == 0 {
self.config.initial_delay
} else {
self.current_delay
};
let jitter = self.apply_jitter(delay);
self.current_attempt += 1;
let max_delay_secs = self.config.max_delay.as_secs_f64();
let next_delay_secs = delay.as_secs_f64() * self.config.backoff_multiplier;
let clamped_secs = if next_delay_secs.is_finite() {
next_delay_secs.min(max_delay_secs)
} else {
max_delay_secs
};
self.current_delay = Duration::from_secs_f64(clamped_secs);
Some(jitter)
}
fn apply_jitter(&self, delay: Duration) -> Duration {
use std::collections::hash_map::RandomState;
use std::hash::BuildHasher;
let hash = RandomState::new().hash_one(self.current_attempt);
let random = (hash % 1000) as f64 / 1000.0;
let jitter_range = delay.as_secs_f64() * self.config.jitter_factor;
let jitter = (random - 0.5) * 2.0 * jitter_range;
let final_delay = (delay.as_secs_f64() + jitter).max(0.0);
Duration::from_secs_f64(final_delay)
}
pub fn reset(&mut self) {
self.current_attempt = 0;
self.current_delay = self.config.initial_delay;
}
pub fn can_reconnect(&self) -> bool {
if let Some(max) = self.config.max_attempts {
self.current_attempt < max
} else {
true
}
}
pub fn config(&self) -> &ReconnectionConfig {
&self.config
}
}
pub type OnReconnectStateChange = Box<dyn Fn(&ReconnectionState) + Send + Sync>;
pub struct AutoReconnectHandler {
strategy: RwLock<ReconnectionStrategy>,
state: Arc<RwLock<ReconnectionState>>,
enabled: bool,
on_state_change: Option<OnReconnectStateChange>,
}
impl AutoReconnectHandler {
pub fn new(config: ReconnectionConfig) -> Self {
Self {
strategy: RwLock::new(ReconnectionStrategy::new(config)),
state: Arc::new(RwLock::new(ReconnectionState::Connected)),
enabled: true,
on_state_change: None,
}
}
pub fn disabled() -> Self {
Self {
strategy: RwLock::new(ReconnectionStrategy::new(ReconnectionConfig::default())),
state: Arc::new(RwLock::new(ReconnectionState::Disabled)),
enabled: false,
on_state_change: None,
}
}
pub fn with_on_state_change(mut self, callback: OnReconnectStateChange) -> Self {
self.on_state_change = Some(callback);
self
}
pub fn is_enabled(&self) -> bool {
self.enabled
}
pub async fn state(&self) -> tokio::sync::RwLockReadGuard<'_, ReconnectionState> {
self.state.read().await
}
pub async fn on_disconnect(&self) -> Option<Duration> {
if !self.enabled {
return None;
}
let mut strategy = self.strategy.write().await;
let delay = strategy.next_delay();
match delay {
Some(d) => {
let new_state = ReconnectionState::Reconnecting {
attempt: strategy.attempt_count(),
next_delay: d,
};
self.set_state(new_state).await;
Some(d)
}
None => {
let new_state = ReconnectionState::Failed {
total_attempts: strategy.attempt_count(),
};
self.set_state(new_state).await;
None
}
}
}
pub async fn on_reconnect_success(&self) {
let mut strategy = self.strategy.write().await;
let total_attempts = strategy.attempt_count();
strategy.reset();
drop(strategy);
self.set_state(ReconnectionState::Reconnected { total_attempts })
.await;
}
pub async fn on_connected(&self) {
self.strategy.write().await.reset();
self.set_state(ReconnectionState::Connected).await;
}
async fn set_state(&self, new_state: ReconnectionState) {
if let Some(cb) = &self.on_state_change {
cb(&new_state);
}
*self.state.write().await = new_state;
}
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use super::*;
#[test]
fn test_default_config() {
let config = ReconnectionConfig::default();
assert_eq!(config.max_attempts, Some(10));
assert_eq!(config.initial_delay, Duration::from_secs(1));
assert_eq!(config.max_delay, Duration::from_secs(300));
assert_eq!(config.backoff_multiplier, 2.0);
assert_eq!(config.jitter_factor, 0.1);
}
#[test]
fn test_config_builder() {
let config = ReconnectionConfig::default()
.with_max_attempts(5)
.with_initial_delay(Duration::from_secs(2))
.with_max_delay(Duration::from_secs(60))
.with_backoff_multiplier(1.5)
.with_jitter_factor(0.2);
assert_eq!(config.max_attempts, Some(5));
assert_eq!(config.initial_delay, Duration::from_secs(2));
assert_eq!(config.max_delay, Duration::from_secs(60));
assert_eq!(config.backoff_multiplier, 1.5);
assert_eq!(config.jitter_factor, 0.2);
}
#[test]
fn test_unlimited_attempts() {
let config = ReconnectionConfig::default().with_unlimited_attempts();
assert_eq!(config.max_attempts, None);
}
#[test]
fn test_reconnection_strategy() {
let config = ReconnectionConfig::default()
.with_max_attempts(3)
.with_initial_delay(Duration::from_secs(1))
.with_jitter_factor(0.0);
let mut strategy = ReconnectionStrategy::new(config);
assert_eq!(strategy.attempt_count(), 0);
assert!(strategy.can_reconnect());
let delay1 = strategy.next_delay().unwrap();
assert_eq!(delay1, Duration::from_secs(1));
assert_eq!(strategy.attempt_count(), 1);
let delay2 = strategy.next_delay().unwrap();
assert_eq!(delay2, Duration::from_secs(2));
assert_eq!(strategy.attempt_count(), 2);
let delay3 = strategy.next_delay().unwrap();
assert_eq!(delay3, Duration::from_secs(4));
assert_eq!(strategy.attempt_count(), 3);
let delay4 = strategy.next_delay();
assert!(delay4.is_none());
assert!(!strategy.can_reconnect());
}
#[test]
fn test_exponential_backoff() {
let config = ReconnectionConfig::default()
.with_unlimited_attempts()
.with_initial_delay(Duration::from_secs(1))
.with_backoff_multiplier(2.0)
.with_max_delay(Duration::from_secs(100))
.with_jitter_factor(0.0);
let mut strategy = ReconnectionStrategy::new(config);
let delay1 = strategy.next_delay().unwrap();
assert_eq!(delay1, Duration::from_secs(1));
let delay2 = strategy.next_delay().unwrap();
assert!(delay2.as_secs() >= 1);
let delay3 = strategy.next_delay().unwrap();
assert!(delay3.as_secs() >= 2);
}
#[test]
fn test_max_delay_cap() {
let config = ReconnectionConfig::default()
.with_unlimited_attempts()
.with_initial_delay(Duration::from_secs(1))
.with_backoff_multiplier(10.0)
.with_max_delay(Duration::from_secs(5))
.with_jitter_factor(0.0);
let mut strategy = ReconnectionStrategy::new(config);
for _ in 0..10 {
if let Some(delay) = strategy.next_delay() {
assert!(delay.as_secs() <= 5);
}
}
}
#[test]
fn test_reset() {
let config = ReconnectionConfig::default().with_max_attempts(5);
let mut strategy = ReconnectionStrategy::new(config);
strategy.next_delay();
strategy.next_delay();
assert_eq!(strategy.attempt_count(), 2);
strategy.reset();
assert_eq!(strategy.attempt_count(), 0);
assert!(strategy.can_reconnect());
}
#[test]
fn test_jitter_applied() {
let config = ReconnectionConfig::default()
.with_initial_delay(Duration::from_secs(1))
.with_jitter_factor(0.1);
let mut strategy = ReconnectionStrategy::new(config);
let delay = strategy.next_delay().unwrap();
let delay_secs = delay.as_secs_f64();
assert!((0.9..=1.1).contains(&delay_secs));
}
#[rstest]
fn test_backoff_does_not_overflow_at_high_retry_counts() {
let config = ReconnectionConfig::default()
.with_unlimited_attempts()
.with_initial_delay(Duration::from_secs(1))
.with_backoff_multiplier(10.0)
.with_max_delay(Duration::from_secs(300))
.with_jitter_factor(0.0);
let mut strategy = ReconnectionStrategy::new(config);
for _ in 0..100 {
if let Some(delay) = strategy.next_delay() {
assert!(delay <= Duration::from_secs(300));
}
}
}
#[rstest]
fn test_strategy_config_accessor() {
let config = ReconnectionConfig::default().with_max_attempts(7);
let strategy = ReconnectionStrategy::new(config);
assert_eq!(strategy.config().max_attempts, Some(7));
}
#[rstest]
fn test_reconnection_state_variants() {
let connected = ReconnectionState::Connected;
assert_eq!(connected, ReconnectionState::Connected);
let reconnecting = ReconnectionState::Reconnecting {
attempt: 1,
next_delay: Duration::from_secs(2),
};
assert_eq!(
reconnecting,
ReconnectionState::Reconnecting {
attempt: 1,
next_delay: Duration::from_secs(2),
}
);
let reconnected = ReconnectionState::Reconnected { total_attempts: 3 };
assert_eq!(
reconnected,
ReconnectionState::Reconnected { total_attempts: 3 }
);
let failed = ReconnectionState::Failed { total_attempts: 5 };
assert_eq!(failed, ReconnectionState::Failed { total_attempts: 5 });
let disabled = ReconnectionState::Disabled;
assert_eq!(disabled, ReconnectionState::Disabled);
}
#[rstest]
#[tokio::test]
async fn test_auto_reconnect_handler_new() {
let handler = AutoReconnectHandler::new(ReconnectionConfig::default());
assert!(handler.is_enabled());
assert_eq!(*handler.state().await, ReconnectionState::Connected);
}
#[rstest]
#[tokio::test]
async fn test_auto_reconnect_handler_disabled() {
let handler = AutoReconnectHandler::disabled();
assert!(!handler.is_enabled());
assert_eq!(*handler.state().await, ReconnectionState::Disabled);
assert!(handler.on_disconnect().await.is_none());
}
#[tokio::test]
async fn test_auto_reconnect_handler_disconnect_and_reconnect() {
let config = ReconnectionConfig::default()
.with_max_attempts(3)
.with_initial_delay(Duration::from_secs(1))
.with_jitter_factor(0.0);
let handler = AutoReconnectHandler::new(config);
let delay = handler.on_disconnect().await;
assert_eq!(delay, Some(Duration::from_secs(1)));
match &*handler.state().await {
ReconnectionState::Reconnecting { attempt, .. } => {
assert_eq!(*attempt, 1);
}
other => panic!("Expected Reconnecting, got {:?}", other),
}
handler.on_reconnect_success().await;
match &*handler.state().await {
ReconnectionState::Reconnected { total_attempts } => {
assert_eq!(*total_attempts, 1);
}
other => panic!("Expected Reconnected, got {:?}", other),
}
}
#[tokio::test]
async fn test_auto_reconnect_handler_exhausted() {
let config = ReconnectionConfig::default()
.with_max_attempts(2)
.with_initial_delay(Duration::from_secs(1))
.with_jitter_factor(0.0);
let handler = AutoReconnectHandler::new(config);
let delay1 = handler.on_disconnect().await;
assert!(delay1.is_some());
let delay2 = handler.on_disconnect().await;
assert!(delay2.is_some());
let delay3 = handler.on_disconnect().await;
assert!(delay3.is_none());
match &*handler.state().await {
ReconnectionState::Failed { total_attempts } => {
assert_eq!(*total_attempts, 2);
}
other => panic!("Expected Failed, got {:?}", other),
}
}
#[tokio::test]
async fn test_auto_reconnect_handler_on_connected_resets() {
let config = ReconnectionConfig::default()
.with_max_attempts(5)
.with_jitter_factor(0.0);
let handler = AutoReconnectHandler::new(config);
handler.on_disconnect().await;
handler.on_reconnect_success().await;
handler.on_connected().await;
assert_eq!(*handler.state().await, ReconnectionState::Connected);
}
#[tokio::test]
async fn test_auto_reconnect_handler_with_callback() {
let callback_fired = Arc::new(std::sync::atomic::AtomicBool::new(false));
let callback_fired_clone = callback_fired.clone();
let config = ReconnectionConfig::default()
.with_max_attempts(3)
.with_jitter_factor(0.0);
let handler =
AutoReconnectHandler::new(config).with_on_state_change(Box::new(move |_state| {
callback_fired_clone.store(true, std::sync::atomic::Ordering::SeqCst);
}));
handler.on_disconnect().await;
assert!(callback_fired.load(std::sync::atomic::Ordering::SeqCst));
}
#[tokio::test]
async fn test_auto_reconnect_handler_exponential_backoff() {
let config = ReconnectionConfig::default()
.with_max_attempts(4)
.with_initial_delay(Duration::from_secs(1))
.with_backoff_multiplier(2.0)
.with_jitter_factor(0.0);
let handler = AutoReconnectHandler::new(config);
let delay1 = handler.on_disconnect().await.unwrap();
assert_eq!(delay1, Duration::from_secs(1));
let delay2 = handler.on_disconnect().await.unwrap();
assert_eq!(delay2, Duration::from_secs(2));
let delay3 = handler.on_disconnect().await.unwrap();
assert_eq!(delay3, Duration::from_secs(4));
let delay4 = handler.on_disconnect().await.unwrap();
assert_eq!(delay4, Duration::from_secs(8));
}
}