1use std::future::Future;
35
36#[derive(Debug, Clone, PartialEq)]
41pub enum RetryStrategy {
42 NoRetry,
44 FixedDelay {
46 delay_ms: u64,
48 },
49 ExponentialBackoff {
52 initial_ms: u64,
54 multiplier: f64,
56 max_ms: u64,
58 },
59 LinearBackoff {
61 initial_ms: u64,
63 increment_ms: u64,
65 max_ms: u64,
67 },
68 Fibonacci {
71 initial_ms: u64,
73 max_ms: u64,
75 },
76}
77
78#[derive(Debug, Clone, PartialEq)]
82pub enum JitterMode {
83 None,
85 Full,
87 Equal,
89 Decorrelated,
91}
92
93#[derive(Debug, Clone)]
97pub struct RetryPolicy {
98 pub strategy: RetryStrategy,
100 pub jitter: JitterMode,
102 pub max_attempts: u32,
104 pub retryable_errors: Vec<String>,
107}
108
109impl RetryPolicy {
110 pub fn is_retryable(&self, error: &str) -> bool {
114 if self.retryable_errors.is_empty() {
115 return true;
116 }
117 self.retryable_errors.iter().any(|pat| error.contains(pat.as_str()))
118 }
119}
120
121struct Lcg {
126 state: u64,
127}
128
129impl Lcg {
130 fn new(seed: u64) -> Self {
131 Self {
133 state: seed.wrapping_mul(6_364_136_223_846_793_005)
134 .wrapping_add(1_442_695_040_888_963_407),
135 }
136 }
137
138 fn next_bounded(&mut self, max: u64) -> u64 {
140 self.state = self.state
141 .wrapping_mul(6_364_136_223_846_793_005)
142 .wrapping_add(1_442_695_040_888_963_407);
143 if max == 0 {
144 return 0;
145 }
146 self.state % (max + 1)
147 }
148}
149
150#[derive(Debug, Clone)]
154pub struct RetryState {
155 pub attempt: u32,
157 pub last_delay_ms: u64,
160 pub total_delay_ms: u64,
162 fib_prev: u64,
164 fib_curr: u64,
166}
167
168impl RetryState {
169 pub fn new() -> Self {
171 Self {
172 attempt: 0,
173 last_delay_ms: 0,
174 total_delay_ms: 0,
175 fib_prev: 0,
176 fib_curr: 0,
177 }
178 }
179
180 pub fn next_delay(&mut self, policy: &RetryPolicy) -> Option<u64> {
184 self.attempt += 1;
185
186 if self.attempt > policy.max_attempts {
187 return None;
188 }
189
190 if self.attempt == 1 {
192 return Some(0);
193 }
194
195 let retry_num = self.attempt - 1; let base_delay = match &policy.strategy {
198 RetryStrategy::NoRetry => return None,
199
200 RetryStrategy::FixedDelay { delay_ms } => *delay_ms,
201
202 RetryStrategy::ExponentialBackoff {
203 initial_ms,
204 multiplier,
205 max_ms,
206 } => {
207 let exp = (*multiplier).powi((retry_num - 1) as i32);
208 let d = (*initial_ms as f64 * exp) as u64;
209 d.min(*max_ms)
210 }
211
212 RetryStrategy::LinearBackoff {
213 initial_ms,
214 increment_ms,
215 max_ms,
216 } => {
217 let d = initial_ms.saturating_add(increment_ms.saturating_mul((retry_num - 1) as u64));
218 d.min(*max_ms)
219 }
220
221 RetryStrategy::Fibonacci { initial_ms, max_ms } => {
222 if retry_num == 1 {
223 self.fib_prev = 0;
225 self.fib_curr = *initial_ms;
226 self.fib_curr.min(*max_ms)
227 } else {
228 let next = self.fib_prev.saturating_add(self.fib_curr);
229 self.fib_prev = self.fib_curr;
230 self.fib_curr = next;
231 self.fib_curr.min(*max_ms)
232 }
233 }
234 };
235
236 let mut rng = Lcg::new(self.attempt as u64 ^ base_delay);
238 let initial_ms = match &policy.strategy {
239 RetryStrategy::ExponentialBackoff { initial_ms, .. } => *initial_ms,
240 RetryStrategy::LinearBackoff { initial_ms, .. } => *initial_ms,
241 RetryStrategy::Fibonacci { initial_ms, .. } => *initial_ms,
242 RetryStrategy::FixedDelay { delay_ms } => *delay_ms,
243 RetryStrategy::NoRetry => 0,
244 };
245
246 let jittered = match &policy.jitter {
247 JitterMode::None => base_delay,
248
249 JitterMode::Full => {
250 if base_delay == 0 { 0 } else { rng.next_bounded(base_delay) }
251 }
252
253 JitterMode::Equal => {
254 let half = base_delay / 2;
255 let jitter_part = if half == 0 { 0 } else { rng.next_bounded(half) };
256 half + jitter_part
257 }
258
259 JitterMode::Decorrelated => {
260 let prev = self.last_delay_ms.max(initial_ms);
261 let upper = (prev * 3).max(initial_ms);
262 let lower = initial_ms.min(upper);
263 lower + rng.next_bounded(upper.saturating_sub(lower))
264 }
265 };
266
267 self.last_delay_ms = jittered;
268 self.total_delay_ms = self.total_delay_ms.saturating_add(jittered);
269 Some(jittered)
270 }
271}
272
273impl Default for RetryState {
274 fn default() -> Self {
275 Self::new()
276 }
277}
278
279#[derive(Debug, Clone, Default)]
283pub struct RetryMetrics {
284 pub total_attempts: u32,
286 pub successful_on_attempt: Option<u32>,
288 pub total_delay_ms: u64,
290}
291
292pub async fn retry_async<F, Fut, T, E>(policy: &RetryPolicy, f: F) -> Result<T, E>
306where
307 F: Fn() -> Fut,
308 Fut: Future<Output = Result<T, E>>,
309 E: std::fmt::Display,
310{
311 let mut state = RetryState::new();
312 let mut last_err: Option<E> = None;
313
314 while let Some(delay) = state.next_delay(policy) {
316
317 if delay > 0 {
318 tokio::time::sleep(std::time::Duration::from_millis(delay)).await;
319 }
320
321 match f().await {
322 Ok(v) => return Ok(v),
323 Err(e) => {
324 let msg = e.to_string();
325 if !policy.is_retryable(&msg) {
326 return Err(e);
327 }
328 last_err = Some(e);
329 }
330 }
331 }
332
333 match last_err {
336 Some(e) => Err(e),
337 None => f().await,
338 }
339}
340
341#[cfg(test)]
344mod tests {
345 use super::*;
346 use std::sync::{Arc, Mutex};
347
348 fn no_retry_policy() -> RetryPolicy {
349 RetryPolicy {
350 strategy: RetryStrategy::NoRetry,
351 jitter: JitterMode::None,
352 max_attempts: 1,
353 retryable_errors: vec![],
354 }
355 }
356
357 fn fixed_policy(delay_ms: u64, max_attempts: u32) -> RetryPolicy {
358 RetryPolicy {
359 strategy: RetryStrategy::FixedDelay { delay_ms },
360 jitter: JitterMode::None,
361 max_attempts,
362 retryable_errors: vec![],
363 }
364 }
365
366 #[test]
367 fn no_retry_gives_one_attempt() {
368 let policy = no_retry_policy();
369 let mut state = RetryState::new();
370
371 assert_eq!(state.next_delay(&policy), Some(0));
373 assert_eq!(state.next_delay(&policy), None);
375 }
376
377 #[test]
378 fn fixed_delay_sequence() {
379 let policy = fixed_policy(100, 4);
380 let mut state = RetryState::new();
381
382 assert_eq!(state.next_delay(&policy), Some(0)); assert_eq!(state.next_delay(&policy), Some(100)); assert_eq!(state.next_delay(&policy), Some(100)); assert_eq!(state.next_delay(&policy), Some(100)); assert_eq!(state.next_delay(&policy), None); }
388
389 #[test]
390 fn exponential_delay_sequence() {
391 let policy = RetryPolicy {
392 strategy: RetryStrategy::ExponentialBackoff {
393 initial_ms: 100,
394 multiplier: 2.0,
395 max_ms: 10_000,
396 },
397 jitter: JitterMode::None,
398 max_attempts: 4,
399 retryable_errors: vec![],
400 };
401 let mut state = RetryState::new();
402
403 assert_eq!(state.next_delay(&policy), Some(0)); assert_eq!(state.next_delay(&policy), Some(100)); assert_eq!(state.next_delay(&policy), Some(200)); assert_eq!(state.next_delay(&policy), Some(400)); assert_eq!(state.next_delay(&policy), None);
408 }
409
410 #[test]
411 fn exponential_capped_at_max() {
412 let policy = RetryPolicy {
413 strategy: RetryStrategy::ExponentialBackoff {
414 initial_ms: 500,
415 multiplier: 10.0,
416 max_ms: 1_000,
417 },
418 jitter: JitterMode::None,
419 max_attempts: 5,
420 retryable_errors: vec![],
421 };
422 let mut state = RetryState::new();
423 state.next_delay(&policy); let d2 = state.next_delay(&policy).unwrap_or(0); let d3 = state.next_delay(&policy).unwrap_or(0); assert!(d2 <= 1_000);
428 assert!(d3 <= 1_000);
429 }
430
431 #[test]
432 fn fibonacci_sequence_correct() {
433 let policy = RetryPolicy {
434 strategy: RetryStrategy::Fibonacci {
435 initial_ms: 100,
436 max_ms: 10_000,
437 },
438 jitter: JitterMode::None,
439 max_attempts: 6,
440 retryable_errors: vec![],
441 };
442 let mut state = RetryState::new();
443
444 assert_eq!(state.next_delay(&policy), Some(0)); assert_eq!(state.next_delay(&policy), Some(100)); assert_eq!(state.next_delay(&policy), Some(100)); assert_eq!(state.next_delay(&policy), Some(200)); assert_eq!(state.next_delay(&policy), Some(300)); assert_eq!(state.next_delay(&policy), Some(500)); assert_eq!(state.next_delay(&policy), None);
451 }
452
453 #[test]
454 fn max_attempts_respected() {
455 let policy = fixed_policy(50, 3);
456 let mut state = RetryState::new();
457 let mut count = 0;
458 while state.next_delay(&policy).is_some() {
459 count += 1;
460 }
461 assert_eq!(count, 3);
462 }
463
464 #[tokio::test]
465 async fn successful_on_third_attempt() {
466 let policy = RetryPolicy {
467 strategy: RetryStrategy::FixedDelay { delay_ms: 0 },
468 jitter: JitterMode::None,
469 max_attempts: 5,
470 retryable_errors: vec![],
471 };
472
473 let call_count = Arc::new(Mutex::new(0u32));
474 let cc = call_count.clone();
475
476 let result: Result<u32, String> = retry_async(&policy, || {
477 let cc = cc.clone();
478 async move {
479 let mut n = cc.lock().unwrap_or_else(|e| e.into_inner());
480 *n += 1;
481 let attempt = *n;
482 drop(n);
483 if attempt < 3 {
484 Err("transient error".to_string())
485 } else {
486 Ok(attempt)
487 }
488 }
489 })
490 .await;
491
492 assert!(result.is_ok());
493 assert_eq!(result.unwrap(), 3);
494 assert_eq!(*call_count.lock().unwrap(), 3);
495 }
496
497 #[tokio::test]
498 async fn all_attempts_fail_returns_last_error() {
499 let policy = RetryPolicy {
500 strategy: RetryStrategy::FixedDelay { delay_ms: 0 },
501 jitter: JitterMode::None,
502 max_attempts: 3,
503 retryable_errors: vec![],
504 };
505
506 let result: Result<(), String> = retry_async(&policy, || async {
507 Err("permanent failure".to_string())
508 })
509 .await;
510
511 assert!(result.is_err());
512 assert_eq!(result.unwrap_err(), "permanent failure");
513 }
514
515 #[tokio::test]
516 async fn non_retryable_error_stops_immediately() {
517 let policy = RetryPolicy {
518 strategy: RetryStrategy::FixedDelay { delay_ms: 0 },
519 jitter: JitterMode::None,
520 max_attempts: 10,
521 retryable_errors: vec!["transient".to_string()],
522 };
523
524 let call_count = Arc::new(Mutex::new(0u32));
525 let cc = call_count.clone();
526
527 let result: Result<(), String> = retry_async(&policy, || {
528 let cc = cc.clone();
529 async move {
530 *cc.lock().unwrap() += 1;
531 Err("fatal: disk full".to_string())
532 }
533 })
534 .await;
535
536 assert!(result.is_err());
537 assert_eq!(*call_count.lock().unwrap(), 1);
539 }
540
541 #[test]
542 fn linear_backoff_sequence() {
543 let policy = RetryPolicy {
544 strategy: RetryStrategy::LinearBackoff {
545 initial_ms: 100,
546 increment_ms: 50,
547 max_ms: 10_000,
548 },
549 jitter: JitterMode::None,
550 max_attempts: 5,
551 retryable_errors: vec![],
552 };
553 let mut state = RetryState::new();
554
555 assert_eq!(state.next_delay(&policy), Some(0)); assert_eq!(state.next_delay(&policy), Some(100)); assert_eq!(state.next_delay(&policy), Some(150)); assert_eq!(state.next_delay(&policy), Some(200)); assert_eq!(state.next_delay(&policy), Some(250)); assert_eq!(state.next_delay(&policy), None);
561 }
562
563 #[test]
564 fn full_jitter_within_bounds() {
565 let policy = RetryPolicy {
566 strategy: RetryStrategy::FixedDelay { delay_ms: 1000 },
567 jitter: JitterMode::Full,
568 max_attempts: 10,
569 retryable_errors: vec![],
570 };
571 let mut state = RetryState::new();
572 state.next_delay(&policy); for _ in 0..9 {
575 if let Some(d) = state.next_delay(&policy) {
576 assert!(d <= 1000, "full jitter must be <= base delay");
577 }
578 }
579 }
580}