1use crate::error::FaucetError;
10use std::future::Future;
11use std::sync::atomic::{AtomicU64, Ordering};
12use std::time::Duration;
13
14const MAX_BACKOFF: Duration = Duration::from_secs(60);
19
20pub async fn execute_with_retry<F, Fut, T>(
25 max_retries: u32,
26 base_backoff: Duration,
27 operation: F,
28) -> Result<T, FaucetError>
29where
30 F: FnMut() -> Fut,
31 Fut: Future<Output = Result<T, FaucetError>>,
32{
33 let policy = crate::resilience::RetryPolicy {
34 max_attempts: max_retries.saturating_add(1),
36 backoff: crate::resilience::BackoffKind::Exponential,
37 base: base_backoff,
38 max: MAX_BACKOFF,
39 jitter: true,
40 retry_on: crate::resilience::RetryClassSet::default(),
41 };
42 crate::resilience::execute_with_policy(&policy, None, operation).await
43}
44
45pub fn backoff_with_jitter(base: Duration, attempt: u32) -> Duration {
53 let exp = base
54 .saturating_mul(2u32.saturating_pow(attempt))
55 .min(MAX_BACKOFF);
56 let nanos = exp.as_nanos() as u64;
57 Duration::from_nanos((nanos as f64 * pseudo_random_factor()) as u64)
58}
59
60fn pseudo_random_factor() -> f64 {
62 static COUNTER: AtomicU64 = AtomicU64::new(0);
63 let nanos = std::time::SystemTime::now()
64 .duration_since(std::time::UNIX_EPOCH)
65 .unwrap_or_default()
66 .subsec_nanos();
67 let counter = COUNTER.fetch_add(1, Ordering::Relaxed);
68 jitter_factor(decorrelate(nanos, counter))
69}
70
71fn decorrelate(nanos: u32, counter: u64) -> u32 {
78 let mut x = (nanos as u64) ^ counter.wrapping_mul(0x9E37_79B9_7F4A_7C15);
79 x ^= x >> 30;
80 x = x.wrapping_mul(0xBF58_476D_1CE4_E5B9);
81 x ^= x >> 27;
82 x = x.wrapping_mul(0x94D0_49BB_1331_11EB);
83 x ^= x >> 31;
84 (x % 1_000_000_000) as u32
85}
86
87fn jitter_factor(nanos: u32) -> f64 {
92 0.5 + (nanos as f64 / 1_000_000_000.0)
93}
94
95pub fn apply_jitter(delay: std::time::Duration) -> std::time::Duration {
106 let nanos = delay.as_nanos() as u64;
107 std::time::Duration::from_nanos((nanos as f64 * pseudo_random_factor()) as u64)
108}
109
110#[cfg(test)]
111mod tests {
112 use super::*;
113 use std::sync::Arc;
114 use std::sync::atomic::{AtomicU32, Ordering};
115
116 #[tokio::test]
117 async fn returns_immediately_on_success() {
118 let calls = Arc::new(AtomicU32::new(0));
119 let c = calls.clone();
120 let r = execute_with_retry(3, Duration::from_millis(1), move || {
121 c.fetch_add(1, Ordering::SeqCst);
122 async { Ok::<_, FaucetError>(7) }
123 })
124 .await;
125 assert_eq!(r.unwrap(), 7);
126 assert_eq!(calls.load(Ordering::SeqCst), 1);
127 }
128
129 #[tokio::test]
130 async fn retries_then_succeeds_on_transient_5xx() {
131 let calls = Arc::new(AtomicU32::new(0));
132 let c = calls.clone();
133 let r = execute_with_retry(3, Duration::from_millis(1), move || {
134 let n = c.fetch_add(1, Ordering::SeqCst);
135 async move {
136 if n < 2 {
137 Err::<i32, _>(FaucetError::HttpStatus {
138 status: 503,
139 url: "http://t".into(),
140 body: "x".into(),
141 })
142 } else {
143 Ok(42)
144 }
145 }
146 })
147 .await;
148 assert_eq!(r.unwrap(), 42);
149 assert_eq!(calls.load(Ordering::SeqCst), 3);
150 }
151
152 #[test]
153 fn jitter_factor_spans_documented_half_to_one_and_a_half_range() {
154 assert_eq!(jitter_factor(0), 0.5);
158 let mid = jitter_factor(500_000_000);
159 assert!((mid - 1.0).abs() < 1e-6, "midpoint factor was {mid}");
160 let hi = jitter_factor(999_999_999);
161 assert!(
162 (1.4..1.5).contains(&hi),
163 "factor at max sub-second nanos was {hi}, expected ~1.5"
164 );
165 }
166
167 #[test]
168 fn backoff_is_capped_for_large_attempt() {
169 let d = backoff_with_jitter(Duration::from_secs(1), 60);
173 assert!(d < Duration::from_secs(90), "backoff not capped: {d:?}");
174 assert!(
176 d >= Duration::from_secs(30),
177 "backoff unexpectedly tiny: {d:?}"
178 );
179 }
180
181 #[test]
182 fn decorrelate_diverges_for_same_nanos_concurrent_calls() {
183 let a = decorrelate(123_456_789, 0);
187 let b = decorrelate(123_456_789, 1);
188 let c = decorrelate(123_456_789, 2);
189 assert_ne!(a, b);
190 assert_ne!(b, c);
191 assert_ne!(a, c);
192 for v in [a, b, c] {
193 assert!(
194 v < 1_000_000_000,
195 "decorrelate out of jitter_factor range: {v}"
196 );
197 }
198 }
199
200 #[tokio::test]
201 async fn non_retriable_fails_immediately() {
202 let calls = Arc::new(AtomicU32::new(0));
203 let c = calls.clone();
204 let r = execute_with_retry(3, Duration::from_millis(1), move || {
205 c.fetch_add(1, Ordering::SeqCst);
206 async { Err::<i32, _>(FaucetError::Auth("nope".into())) }
207 })
208 .await;
209 assert!(r.is_err());
210 assert_eq!(calls.load(Ordering::SeqCst), 1);
211 }
212}
213
214#[derive(
223 Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize, schemars::JsonSchema,
224)]
225#[serde(deny_unknown_fields)]
226pub struct PartialRetrySpec {
227 #[serde(default = "default_partial_retry_max_attempts")]
230 pub max_attempts: usize,
231 #[serde(default = "default_partial_retry_initial_backoff_ms")]
234 pub initial_backoff_ms: u64,
235 #[serde(default = "default_partial_retry_max_backoff_ms")]
237 pub max_backoff_ms: u64,
238}
239
240fn default_partial_retry_max_attempts() -> usize {
241 5
242}
243fn default_partial_retry_initial_backoff_ms() -> u64 {
244 100
245}
246fn default_partial_retry_max_backoff_ms() -> u64 {
247 30_000
248}
249
250impl Default for PartialRetrySpec {
251 fn default() -> Self {
252 Self {
253 max_attempts: default_partial_retry_max_attempts(),
254 initial_backoff_ms: default_partial_retry_initial_backoff_ms(),
255 max_backoff_ms: default_partial_retry_max_backoff_ms(),
256 }
257 }
258}
259
260#[cfg(test)]
261mod partial_retry_tests {
262 use super::PartialRetrySpec;
263
264 #[test]
265 fn defaults_match_the_documented_values() {
266 let s = PartialRetrySpec::default();
267 assert_eq!(s.max_attempts, 5);
268 assert_eq!(s.initial_backoff_ms, 100);
269 assert_eq!(s.max_backoff_ms, 30_000);
270 let from_json: PartialRetrySpec = serde_json::from_str("{}").unwrap();
273 assert_eq!(from_json, s);
274 }
275
276 #[test]
277 fn an_unknown_key_is_rejected() {
278 let err = serde_json::from_str::<PartialRetrySpec>(r#"{"max_attempt": 3}"#)
279 .expect_err("a typo'd key must not be silently ignored");
280 assert!(err.to_string().contains("max_attempt"), "{err}");
281 }
282}