Skip to main content

tokio_prompt_orchestrator/
retry_policy.rs

1//! # Retry Policy
2//!
3//! Configurable retry logic with multiple back-off strategies, jitter modes,
4//! and async execution support.
5//!
6//! ## Example
7//!
8//! ```rust
9//! use tokio_prompt_orchestrator::retry_policy::{
10//!     JitterMode, RetryPolicy, RetryStrategy, retry_async,
11//! };
12//!
13//! # tokio_test::block_on(async {
14//! let policy = RetryPolicy {
15//!     strategy: RetryStrategy::ExponentialBackoff {
16//!         initial_ms: 10,
17//!         multiplier: 2.0,
18//!         max_ms: 1000,
19//!     },
20//!     jitter: JitterMode::Full,
21//!     max_attempts: 3,
22//!     retryable_errors: vec![],
23//! };
24//!
25//! let call_count = std::sync::atomic::AtomicU32::new(0);
26//! let result: Result<(), String> = retry_async(&policy, || async {
27//!     let n = call_count.fetch_add(1, std::sync::atomic::Ordering::SeqCst) + 1;
28//!     if n < 3 { Err("transient".to_string()) } else { Ok(()) }
29//! }).await;
30//! assert!(result.is_ok());
31//! # });
32//! ```
33
34use std::future::Future;
35
36// ── RetryStrategy ─────────────────────────────────────────────────────────────
37
38/// The back-off strategy that determines how long to wait between retry
39/// attempts.
40#[derive(Debug, Clone, PartialEq)]
41pub enum RetryStrategy {
42    /// Do not retry — fail immediately after the first attempt.
43    NoRetry,
44    /// Wait a constant `delay_ms` milliseconds between every attempt.
45    FixedDelay {
46        /// Constant delay in milliseconds.
47        delay_ms: u64,
48    },
49    /// Multiply the previous delay by `multiplier` each attempt, capped at
50    /// `max_ms`.
51    ExponentialBackoff {
52        /// Delay before the first retry in milliseconds.
53        initial_ms: u64,
54        /// Multiplicative factor applied after each attempt.
55        multiplier: f64,
56        /// Upper bound on the delay in milliseconds.
57        max_ms: u64,
58    },
59    /// Increase the delay by `increment_ms` each attempt, capped at `max_ms`.
60    LinearBackoff {
61        /// Delay before the first retry in milliseconds.
62        initial_ms: u64,
63        /// Amount added to the previous delay each attempt.
64        increment_ms: u64,
65        /// Upper bound on the delay in milliseconds.
66        max_ms: u64,
67    },
68    /// Use the Fibonacci sequence as delays (F(1)=initial, F(2)=initial, …),
69    /// capped at `max_ms`.
70    Fibonacci {
71        /// Seed value for the first two terms of the sequence.
72        initial_ms: u64,
73        /// Upper bound on the delay in milliseconds.
74        max_ms: u64,
75    },
76}
77
78// ── JitterMode ────────────────────────────────────────────────────────────────
79
80/// Jitter strategy to reduce correlated retry storms.
81#[derive(Debug, Clone, PartialEq)]
82pub enum JitterMode {
83    /// No jitter — use the computed delay as-is.
84    None,
85    /// Full jitter: `random in [0, delay]`.
86    Full,
87    /// Equal jitter: `delay/2 + random in [0, delay/2]`.
88    Equal,
89    /// Decorrelated jitter: `random in [initial, 3 * prev_delay]`.
90    Decorrelated,
91}
92
93// ── RetryPolicy ───────────────────────────────────────────────────────────────
94
95/// A complete retry configuration.
96#[derive(Debug, Clone)]
97pub struct RetryPolicy {
98    /// Delay back-off strategy.
99    pub strategy: RetryStrategy,
100    /// Jitter mode applied to the computed delay.
101    pub jitter: JitterMode,
102    /// Maximum total number of attempts (including the first).
103    pub max_attempts: u32,
104    /// Error message substrings that are considered retryable.  An empty list
105    /// means all errors are retryable.
106    pub retryable_errors: Vec<String>,
107}
108
109impl RetryPolicy {
110    /// Returns `true` if the given error message is eligible for a retry.
111    ///
112    /// If `retryable_errors` is empty every error is considered retryable.
113    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
121// ── LCG RNG ───────────────────────────────────────────────────────────────────
122
123/// Minimal linear-congruential generator for deterministic, allocation-free
124/// jitter without pulling in `rand`.
125struct Lcg {
126    state: u64,
127}
128
129impl Lcg {
130    fn new(seed: u64) -> Self {
131        // Knuth's MMIX constants.
132        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    /// Returns a pseudo-random value in `[0, max]`.
139    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// ── RetryState ────────────────────────────────────────────────────────────────
151
152/// Mutable state tracked across retry attempts.
153#[derive(Debug, Clone)]
154pub struct RetryState {
155    /// The attempt number that was most recently executed (1-based).
156    pub attempt: u32,
157    /// The delay that was used before the most recent attempt (0 for the first
158    /// attempt).
159    pub last_delay_ms: u64,
160    /// Cumulative delay accumulated across all waits.
161    pub total_delay_ms: u64,
162    /// Previous-previous delay used by the Fibonacci strategy.
163    fib_prev: u64,
164    /// Previous delay used by the Fibonacci strategy.
165    fib_curr: u64,
166}
167
168impl RetryState {
169    /// Creates a fresh [`RetryState`] ready for the first attempt.
170    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    /// Advances the state and returns the delay in milliseconds that should
181    /// be waited before the next attempt, or `None` if `max_attempts` has been
182    /// reached.
183    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        // First attempt: no delay.
191        if self.attempt == 1 {
192            return Some(0);
193        }
194
195        // Compute base delay from strategy.
196        let retry_num = self.attempt - 1; // 1 on first retry
197        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                    // Bootstrap: F(1) = initial, F(2) = initial
224                    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        // Apply jitter.
237        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// ── RetryMetrics ──────────────────────────────────────────────────────────────
280
281/// Summary of a completed retry sequence.
282#[derive(Debug, Clone, Default)]
283pub struct RetryMetrics {
284    /// Total number of calls made (including successes and failures).
285    pub total_attempts: u32,
286    /// The attempt number on which the call succeeded, or `None` if all failed.
287    pub successful_on_attempt: Option<u32>,
288    /// Total milliseconds slept across all waits.
289    pub total_delay_ms: u64,
290}
291
292// ── retry_async ───────────────────────────────────────────────────────────────
293
294/// Executes `f` up to `policy.max_attempts` times, sleeping between attempts
295/// according to the policy, and returns the first `Ok` result or the last
296/// error.
297///
298/// # Type Parameters
299///
300/// - `F`: An async closure/function that produces a `Future`.
301/// - `Fut`: The `Future` type returned by `f`.
302/// - `T`: The success type.
303/// - `E`: The error type; must implement [`std::fmt::Display`] for retryability
304///   checks.
305pub 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    // Stops when max attempts are exhausted.
315    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    // With max_attempts >= 1 the loop always records an error before ending.
334    // With max_attempts == 0 no attempt was made; make exactly one rather than panic.
335    match last_err {
336        Some(e) => Err(e),
337        None => f().await,
338    }
339}
340
341// ── Tests ─────────────────────────────────────────────────────────────────────
342
343#[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        // First attempt: delay = 0 (the initial call).
372        assert_eq!(state.next_delay(&policy), Some(0));
373        // Second attempt: max_attempts=1 already exhausted.
374        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));   // attempt 1: no pre-wait
383        assert_eq!(state.next_delay(&policy), Some(100)); // attempt 2
384        assert_eq!(state.next_delay(&policy), Some(100)); // attempt 3
385        assert_eq!(state.next_delay(&policy), Some(100)); // attempt 4
386        assert_eq!(state.next_delay(&policy), None);       // exhausted
387    }
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));    // attempt 1
404        assert_eq!(state.next_delay(&policy), Some(100));  // attempt 2: 100 * 2^0
405        assert_eq!(state.next_delay(&policy), Some(200));  // attempt 3: 100 * 2^1
406        assert_eq!(state.next_delay(&policy), Some(400));  // attempt 4: 100 * 2^2
407        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); // attempt 1
424
425        let d2 = state.next_delay(&policy).unwrap_or(0); // 500
426        let d3 = state.next_delay(&policy).unwrap_or(0); // 5000 -> capped 1000
427        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));   // attempt 1
445        assert_eq!(state.next_delay(&policy), Some(100)); // F(1) = 100
446        assert_eq!(state.next_delay(&policy), Some(100)); // F(2) = 100
447        assert_eq!(state.next_delay(&policy), Some(200)); // F(3) = 200
448        assert_eq!(state.next_delay(&policy), Some(300)); // F(4) = 300
449        assert_eq!(state.next_delay(&policy), Some(500)); // F(5) = 500
450        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        // Should stop after first attempt because error doesn't match "transient".
538        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));   // attempt 1
556        assert_eq!(state.next_delay(&policy), Some(100)); // 100 + 50*0
557        assert_eq!(state.next_delay(&policy), Some(150)); // 100 + 50*1
558        assert_eq!(state.next_delay(&policy), Some(200)); // 100 + 50*2
559        assert_eq!(state.next_delay(&policy), Some(250)); // 100 + 50*3
560        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); // skip attempt 1
573
574        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}