tokio_prompt_orchestrator/
retry_budget.rs1use dashmap::DashMap;
7use std::sync::{
8 atomic::{AtomicU32, AtomicU64, Ordering},
9 Arc, Mutex,
10};
11use std::time::{Duration, Instant};
12
13#[derive(Debug, Clone)]
19pub enum BackoffStrategy {
20 Fixed(Duration),
22 Linear {
24 base: Duration,
26 increment: Duration,
28 },
29 Exponential {
31 base: Duration,
33 max: Duration,
35 multiplier: f64,
37 },
38 DecorrelatedJitter {
41 base: Duration,
43 max: Duration,
45 seed: u64,
47 },
48}
49
50#[derive(Debug, thiserror::Error)]
56pub enum RetryError {
57 #[error("retry budget exhausted after {attempts} attempts")]
59 BudgetExhausted {
60 attempts: u32,
62 },
63 #[error("non-retryable error: {0}")]
65 NonRetryableError(String),
66 #[error("next delay would exceed remaining time budget")]
68 MaxDelayExceeded,
69}
70
71#[derive(Debug, Clone)]
77pub struct AttemptRecord {
78 pub attempt_num: u32,
80 pub started_at: Instant,
82 pub duration_ms: u64,
84 pub success: bool,
86 pub error_message: Option<String>,
88}
89
90#[derive(Debug, Clone)]
96pub struct BudgetReport {
97 pub attempts_used: u32,
99 pub attempts_remaining: u32,
101 pub time_used_ms: u64,
103 pub success_rate: f64,
105 pub avg_attempt_ms: f64,
107}
108
109pub 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 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 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 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 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 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 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 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 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 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 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
265pub struct RetryPolicy {
272 budgets: DashMap<String, Arc<RetryBudget>>,
273}
274
275impl RetryPolicy {
276 pub fn new() -> Self {
278 Self {
279 budgets: DashMap::new(),
280 }
281 }
282
283 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 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 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}