1use std::time::Duration;
9
10const FIRST_DELAY: Duration = Duration::from_millis(250);
12
13const MAX_DELAY: Duration = Duration::from_secs(2);
15
16#[derive(Debug, Clone, Copy, PartialEq, Eq)]
18pub struct ReconnectPolicy {
19 timeout: Duration,
20}
21
22impl Default for ReconnectPolicy {
23 fn default() -> Self {
24 Self::from_secs(Self::DEFAULT_TIMEOUT_SECS)
25 }
26}
27
28impl ReconnectPolicy {
29 pub const DEFAULT_TIMEOUT_SECS: u64 = 0;
37
38 pub fn from_secs(secs: u64) -> Self {
41 Self {
42 timeout: Duration::from_secs(secs),
43 }
44 }
45
46 pub fn is_disabled(&self) -> bool {
48 self.timeout.is_zero()
49 }
50
51 pub fn timeout(&self) -> Duration {
53 self.timeout
54 }
55
56 pub fn schedule(&self) -> Backoff {
60 Backoff {
61 remaining: self.timeout,
62 next: FIRST_DELAY,
63 }
64 }
65}
66
67#[derive(Debug)]
72pub struct Backoff {
73 remaining: Duration,
74 next: Duration,
75}
76
77impl Iterator for Backoff {
78 type Item = Duration;
79
80 fn next(&mut self) -> Option<Duration> {
81 if self.remaining.is_zero() {
82 return None;
83 }
84 let delay = self.next.min(self.remaining);
85 self.remaining -= delay;
86 self.next = (self.next * 2).min(MAX_DELAY);
87 Some(delay)
88 }
89}
90
91#[cfg(test)]
92mod tests {
93 use super::*;
94
95 #[test]
96 fn default_is_disabled() {
97 assert_eq!(ReconnectPolicy::default().timeout(), Duration::ZERO);
102 assert!(ReconnectPolicy::default().is_disabled());
103 }
104
105 #[test]
106 fn zero_disables_reconnect() {
107 let policy = ReconnectPolicy::from_secs(0);
108 assert!(policy.is_disabled());
109 assert_eq!(policy.schedule().count(), 0, "no attempts when disabled");
110 }
111
112 #[test]
113 fn delays_grow_then_cap() {
114 let delays: Vec<Duration> = ReconnectPolicy::from_secs(30).schedule().take(6).collect();
115 assert_eq!(
116 delays,
117 vec![
118 Duration::from_millis(250),
119 Duration::from_millis(500),
120 Duration::from_secs(1),
121 Duration::from_secs(2),
122 Duration::from_secs(2),
123 Duration::from_secs(2),
124 ]
125 );
126 }
127
128 #[test]
129 fn total_wait_never_exceeds_the_window() {
130 for secs in [1, 3, 30, 120] {
131 let total: Duration = ReconnectPolicy::from_secs(secs).schedule().sum();
132 assert_eq!(
133 total,
134 Duration::from_secs(secs),
135 "the schedule should use exactly the window for {secs}s"
136 );
137 }
138 }
139
140 #[test]
141 fn short_window_still_gets_an_attempt() {
142 let delays: Vec<Duration> = ReconnectPolicy::from_secs(1).schedule().collect();
143 assert_eq!(delays.first(), Some(&Duration::from_millis(250)));
144 assert!(delays.len() >= 3, "a 1s window should retry a few times");
145 }
146
147 #[test]
148 fn last_delay_is_clipped_to_the_window() {
149 let delays: Vec<Duration> = ReconnectPolicy::from_secs(3).schedule().collect();
151 assert_eq!(delays.last(), Some(&Duration::from_millis(1250)));
152 }
153}