Skip to main content

tokio_prompt_orchestrator/
retry_budget.rs

1//! Retry budgeting with exponential backoff tracking.
2//!
3//! Provides per-key retry budgets with configurable backoff strategies,
4//! time-based budget enforcement, and per-attempt history recording.
5
6use dashmap::DashMap;
7use std::sync::{
8    atomic::{AtomicU32, AtomicU64, Ordering},
9    Arc, Mutex,
10};
11use std::time::{Duration, Instant};
12
13// ---------------------------------------------------------------------------
14// BackoffStrategy
15// ---------------------------------------------------------------------------
16
17/// Backoff strategy used to compute the delay before each retry attempt.
18#[derive(Debug, Clone)]
19pub enum BackoffStrategy {
20    /// Wait the same fixed duration before every attempt.
21    Fixed(Duration),
22    /// Increase the delay linearly: `base + attempt * increment`.
23    Linear {
24        /// Starting delay.
25        base: Duration,
26        /// Amount added per attempt.
27        increment: Duration,
28    },
29    /// Exponential backoff: `base * multiplier^attempt`, capped at `max`.
30    Exponential {
31        /// Initial delay.
32        base: Duration,
33        /// Maximum delay.
34        max: Duration,
35        /// Growth multiplier (e.g. `2.0` for doubling).
36        multiplier: f64,
37    },
38    /// Decorrelated jitter: `min(max, random_between(base, prev * 3))`.
39    /// Uses a simple LCG seeded with `seed`.
40    DecorrelatedJitter {
41        /// Minimum delay.
42        base: Duration,
43        /// Maximum delay.
44        max: Duration,
45        /// Initial LCG seed.
46        seed: u64,
47    },
48}
49
50// ---------------------------------------------------------------------------
51// RetryError
52// ---------------------------------------------------------------------------
53
54/// Errors returned when the retry machinery refuses to issue another attempt.
55#[derive(Debug, thiserror::Error)]
56pub enum RetryError {
57    /// No more attempts remain (either the count or time budget was exceeded).
58    #[error("retry budget exhausted after {attempts} attempts")]
59    BudgetExhausted {
60        /// Number of attempts that were made.
61        attempts: u32,
62    },
63    /// The error is not retryable (e.g. a 4xx client error).
64    #[error("non-retryable error: {0}")]
65    NonRetryableError(String),
66    /// The computed delay would exceed the remaining time budget.
67    #[error("next delay would exceed remaining time budget")]
68    MaxDelayExceeded,
69}
70
71// ---------------------------------------------------------------------------
72// AttemptRecord
73// ---------------------------------------------------------------------------
74
75/// Record of a single attempt.
76#[derive(Debug, Clone)]
77pub struct AttemptRecord {
78    /// Attempt number (1-based).
79    pub attempt_num: u32,
80    /// Instant at which the attempt started.
81    pub started_at: Instant,
82    /// How long the attempt took in milliseconds.
83    pub duration_ms: u64,
84    /// Whether the attempt succeeded.
85    pub success: bool,
86    /// Optional error message if the attempt failed.
87    pub error_message: Option<String>,
88}
89
90// ---------------------------------------------------------------------------
91// BudgetReport
92// ---------------------------------------------------------------------------
93
94/// Summary report of a [`RetryBudget`]'s current state.
95#[derive(Debug, Clone)]
96pub struct BudgetReport {
97    /// Number of attempts consumed so far.
98    pub attempts_used: u32,
99    /// Number of attempts still available.
100    pub attempts_remaining: u32,
101    /// Total milliseconds consumed across all attempts.
102    pub time_used_ms: u64,
103    /// Fraction of attempts that succeeded.
104    pub success_rate: f64,
105    /// Average duration per attempt in milliseconds.
106    pub avg_attempt_ms: f64,
107}
108
109// ---------------------------------------------------------------------------
110// RetryBudget
111// ---------------------------------------------------------------------------
112
113/// A budget that limits how many times an operation may be retried, and for
114/// how long in total.
115pub struct RetryBudget {
116    max_attempts: u32,
117    total_time_budget: Duration,
118    strategy: BackoffStrategy,
119    current_attempts: AtomicU32,
120    total_elapsed_ms: AtomicU64,
121    history: Mutex<Vec<AttemptRecord>>,
122}
123
124impl RetryBudget {
125    /// Create a new `RetryBudget`.
126    pub fn new(max_attempts: u32, time_budget: Duration, strategy: BackoffStrategy) -> Self {
127        Self {
128            max_attempts,
129            total_time_budget: time_budget,
130            strategy,
131            current_attempts: AtomicU32::new(0),
132            total_elapsed_ms: AtomicU64::new(0),
133            history: Mutex::new(Vec::new()),
134        }
135    }
136
137    /// Compute the delay before attempt number `attempt` (0-indexed).
138    ///
139    /// Returns `None` if the budget is exhausted.
140    pub fn next_delay(&self, attempt: u32) -> Option<Duration> {
141        if !self.can_retry(Duration::from_millis(self.total_elapsed_ms.load(Ordering::Relaxed))) {
142            return None;
143        }
144
145        let delay = match &self.strategy {
146            BackoffStrategy::Fixed(d) => *d,
147
148            BackoffStrategy::Linear { base, increment } => {
149                *base + *increment * attempt
150            }
151
152            BackoffStrategy::Exponential { base, max, multiplier } => {
153                let factor = multiplier.powi(attempt as i32);
154                let ms = (base.as_millis() as f64 * factor) as u64;
155                let computed = Duration::from_millis(ms);
156                computed.min(*max)
157            }
158
159            BackoffStrategy::DecorrelatedJitter { base, max, seed } => {
160                // LCG: next = (a * prev + c) % m
161                let mut state = seed.wrapping_add(attempt as u64 * 6364136223846793005);
162                state = state
163                    .wrapping_mul(6364136223846793005)
164                    .wrapping_add(1442695040888963407);
165                let prev_ms = if attempt == 0 {
166                    base.as_millis() as u64
167                } else {
168                    // approximate prev as base * 3^(attempt-1), capped
169                    let prev_factor = 3u64.saturating_pow(attempt - 1);
170                    (base.as_millis() as u64).saturating_mul(prev_factor)
171                };
172                let range = (prev_ms * 3).max(base.as_millis() as u64) - base.as_millis() as u64;
173                let jitter_ms = if range == 0 {
174                    0
175                } else {
176                    base.as_millis() as u64 + state % range
177                };
178                Duration::from_millis(jitter_ms).min(*max)
179            }
180        };
181
182        Some(delay)
183    }
184
185    /// Returns `true` if another attempt is permitted given that `elapsed`
186    /// time has already been spent.
187    pub fn can_retry(&self, elapsed: Duration) -> bool {
188        let attempts = self.current_attempts.load(Ordering::Relaxed);
189        attempts < self.max_attempts && elapsed < self.total_time_budget
190    }
191
192    /// Record the outcome of an attempt.
193    pub fn record_attempt(&self, success: bool, duration_ms: u64, error: Option<String>) {
194        let attempt_num = self.current_attempts.fetch_add(1, Ordering::Relaxed) + 1;
195        self.total_elapsed_ms
196            .fetch_add(duration_ms, Ordering::Relaxed);
197
198        let record = AttemptRecord {
199            attempt_num,
200            started_at: Instant::now(),
201            duration_ms,
202            success,
203            error_message: error,
204        };
205
206        if let Ok(mut guard) = self.history.lock() {
207            guard.push(record);
208        }
209    }
210
211    /// Reset all counters and clear attempt history.
212    pub fn reset(&self) {
213        self.current_attempts.store(0, Ordering::Relaxed);
214        self.total_elapsed_ms.store(0, Ordering::Relaxed);
215        if let Ok(mut guard) = self.history.lock() {
216            guard.clear();
217        }
218    }
219
220    /// Return the number of attempts still available.
221    pub fn remaining_attempts(&self) -> u32 {
222        let used = self.current_attempts.load(Ordering::Relaxed);
223        self.max_attempts.saturating_sub(used)
224    }
225
226    /// Return the remaining time budget given that `elapsed` has already been
227    /// consumed.  Returns `None` if the budget is already exhausted.
228    pub fn remaining_time(&self, elapsed: Duration) -> Option<Duration> {
229        if elapsed >= self.total_time_budget {
230            None
231        } else {
232            Some(self.total_time_budget - elapsed)
233        }
234    }
235
236    /// Produce a summary report of the current budget state.
237    pub fn budget_report(&self) -> BudgetReport {
238        let attempts_used = self.current_attempts.load(Ordering::Relaxed);
239        let attempts_remaining = self.max_attempts.saturating_sub(attempts_used);
240        let time_used_ms = self.total_elapsed_ms.load(Ordering::Relaxed);
241
242        let (success_rate, avg_attempt_ms) = if let Ok(guard) = self.history.lock() {
243            if guard.is_empty() {
244                (0.0, 0.0)
245            } else {
246                let successes = guard.iter().filter(|r| r.success).count() as f64;
247                let total = guard.len() as f64;
248                let avg = guard.iter().map(|r| r.duration_ms as f64).sum::<f64>() / total;
249                (successes / total, avg)
250            }
251        } else {
252            (0.0, 0.0)
253        };
254
255        BudgetReport {
256            attempts_used,
257            attempts_remaining,
258            time_used_ms,
259            success_rate,
260            avg_attempt_ms,
261        }
262    }
263}
264
265// ---------------------------------------------------------------------------
266// RetryPolicy
267// ---------------------------------------------------------------------------
268
269/// A collection of named [`RetryBudget`]s, keyed by an arbitrary string
270/// (e.g. provider name, endpoint, or request type).
271pub struct RetryPolicy {
272    budgets: DashMap<String, Arc<RetryBudget>>,
273}
274
275impl RetryPolicy {
276    /// Create a new, empty `RetryPolicy`.
277    pub fn new() -> Self {
278        Self {
279            budgets: DashMap::new(),
280        }
281    }
282
283    /// Retrieve the budget for `key`, creating it with `config` if it does
284    /// not yet exist.
285    ///
286    /// `config` is `(max_attempts, time_budget, strategy)`.
287    pub fn get_or_create(
288        &self,
289        key: &str,
290        config: (u32, Duration, BackoffStrategy),
291    ) -> Arc<RetryBudget> {
292        if let Some(budget) = self.budgets.get(key) {
293            return Arc::clone(&*budget);
294        }
295
296        let (max_attempts, time_budget, strategy) = config;
297        let budget = Arc::new(RetryBudget::new(max_attempts, time_budget, strategy));
298        self.budgets.insert(key.to_string(), Arc::clone(&budget));
299        budget
300    }
301
302    /// Return a summary report for every tracked budget.
303    pub fn policy_summary(&self) -> Vec<(String, BudgetReport)> {
304        self.budgets
305            .iter()
306            .map(|entry| (entry.key().clone(), entry.value().budget_report()))
307            .collect()
308    }
309}
310
311impl Default for RetryPolicy {
312    fn default() -> Self {
313        Self::new()
314    }
315}
316
317#[cfg(test)]
318mod tests {
319    use super::*;
320
321    #[test]
322    fn test_fixed_backoff() {
323        let budget = RetryBudget::new(
324            3,
325            Duration::from_secs(60),
326            BackoffStrategy::Fixed(Duration::from_millis(100)),
327        );
328        assert_eq!(budget.next_delay(0), Some(Duration::from_millis(100)));
329        assert_eq!(budget.next_delay(2), Some(Duration::from_millis(100)));
330    }
331
332    #[test]
333    fn test_exponential_backoff() {
334        let budget = RetryBudget::new(
335            5,
336            Duration::from_secs(60),
337            BackoffStrategy::Exponential {
338                base: Duration::from_millis(100),
339                max: Duration::from_secs(10),
340                multiplier: 2.0,
341            },
342        );
343        let d0 = budget.next_delay(0).expect("should have delay");
344        let d1 = budget.next_delay(1).expect("should have delay");
345        assert!(d1 > d0);
346    }
347
348    #[test]
349    fn test_budget_exhausted() {
350        let budget = RetryBudget::new(
351            2,
352            Duration::from_secs(60),
353            BackoffStrategy::Fixed(Duration::from_millis(10)),
354        );
355        budget.record_attempt(false, 10, Some("err".to_string()));
356        budget.record_attempt(false, 10, Some("err".to_string()));
357        assert!(!budget.can_retry(Duration::from_millis(20)));
358        assert_eq!(budget.next_delay(2), None);
359    }
360
361    #[test]
362    fn test_reset() {
363        let budget = RetryBudget::new(
364            3,
365            Duration::from_secs(60),
366            BackoffStrategy::Fixed(Duration::from_millis(50)),
367        );
368        budget.record_attempt(true, 50, None);
369        assert_eq!(budget.remaining_attempts(), 2);
370        budget.reset();
371        assert_eq!(budget.remaining_attempts(), 3);
372    }
373
374    #[test]
375    fn test_policy_get_or_create() {
376        let policy = RetryPolicy::new();
377        let b1 = policy.get_or_create(
378            "api",
379            (
380                3,
381                Duration::from_secs(30),
382                BackoffStrategy::Fixed(Duration::from_millis(100)),
383            ),
384        );
385        let b2 = policy.get_or_create(
386            "api",
387            (
388                5,
389                Duration::from_secs(60),
390                BackoffStrategy::Fixed(Duration::from_millis(200)),
391            ),
392        );
393        // Should return the same budget
394        assert_eq!(Arc::ptr_eq(&b1, &b2), true);
395    }
396
397    #[test]
398    fn test_budget_report() {
399        let budget = RetryBudget::new(
400            5,
401            Duration::from_secs(60),
402            BackoffStrategy::Fixed(Duration::from_millis(100)),
403        );
404        budget.record_attempt(true, 50, None);
405        budget.record_attempt(false, 80, Some("timeout".to_string()));
406        let report = budget.budget_report();
407        assert_eq!(report.attempts_used, 2);
408        assert_eq!(report.attempts_remaining, 3);
409        assert!((report.success_rate - 0.5).abs() < 1e-9);
410    }
411}