1use std::future::Future;
6use std::time::Duration;
7
8use crate::DbError;
9
10#[derive(Debug, Clone)]
12pub struct RetryPolicy {
13 pub max_retries: u32,
15 pub initial_delay: Duration,
17 pub max_delay: Duration,
19 pub backoff_factor: f64,
21 pub jitter: bool,
23}
24
25impl Default for RetryPolicy {
26 fn default() -> Self {
27 Self {
28 max_retries: 3,
29 initial_delay: Duration::from_millis(10),
30 max_delay: Duration::from_secs(1),
31 backoff_factor: 2.0,
32 jitter: true,
33 }
34 }
35}
36
37impl RetryPolicy {
38 pub fn delay(&self, retry: u32) -> Duration {
40 let base = self.initial_delay.as_millis() as f64;
41 let delay_ms = base * self.backoff_factor.powi(retry as i32);
42 let delay = Duration::from_millis(delay_ms as u64).min(self.max_delay);
43 if self.jitter {
44 let jitter_factor = 0.5 + rand_simple();
46 Duration::from_millis((delay.as_millis() as f64 * jitter_factor) as u64)
47 } else {
48 delay
49 }
50 }
51}
52
53fn rand_simple() -> f64 {
55 use std::sync::atomic::{AtomicU64, Ordering};
56 static SEED: AtomicU64 = AtomicU64::new(12345);
57 let s = SEED.fetch_add(2654435761, Ordering::Relaxed);
58 (s % 1000) as f64 / 1000.0
59}
60
61pub async fn retry_with_backoff<F, Fut, T>(
65 policy: &RetryPolicy,
66 mut operation: F,
67) -> Result<T, DbError>
68where
69 F: FnMut() -> Fut,
70 Fut: Future<Output = Result<T, DbError>>,
71{
72 let mut last_err = None;
73 for attempt in 0..=policy.max_retries {
74 match operation().await {
75 Ok(result) => return Ok(result),
76 Err(e) => {
77 if !e.is_retryable() || attempt == policy.max_retries {
78 return Err(e);
79 }
80 last_err = Some(e);
81 tokio::time::sleep(policy.delay(attempt)).await;
82 }
83 }
84 }
85 Err(last_err.unwrap_or(DbError::Internal("retry exhausted".into())))
86}
87
88#[cfg(test)]
89mod tests {
90 use super::*;
91 use std::sync::atomic::{AtomicU32, Ordering};
92 use std::sync::Arc;
93
94 #[test]
95 fn test_retry_policy_default() {
96 let p = RetryPolicy::default();
97 assert_eq!(p.max_retries, 3);
98 assert_eq!(p.initial_delay, Duration::from_millis(10));
99 assert_eq!(p.max_delay, Duration::from_secs(1));
100 assert!((p.backoff_factor - 2.0).abs() < f64::EPSILON);
101 assert!(p.jitter);
102 }
103
104 #[test]
105 fn test_retry_policy_delay_within_bounds() {
106 let p = RetryPolicy {
107 jitter: false,
108 ..Default::default()
109 };
110 let d0 = p.delay(0);
111 let d1 = p.delay(1);
112 assert_eq!(d0, Duration::from_millis(10));
114 assert_eq!(d1, Duration::from_millis(20));
115 }
116
117 #[test]
118 fn test_retry_policy_delay_capped_at_max() {
119 let p = RetryPolicy {
120 jitter: false,
121 max_delay: Duration::from_millis(50),
122 ..Default::default()
123 };
124 let d = p.delay(10);
126 assert_eq!(d, Duration::from_millis(50));
127 }
128
129 #[tokio::test]
130 async fn test_retry_succeeds_first_attempt() {
131 let policy = RetryPolicy::default();
132 let counter = Arc::new(AtomicU32::new(0));
133 let c = counter.clone();
134 let result: Result<u32, DbError> = retry_with_backoff(&policy, || {
135 let c = c.clone();
136 async move {
137 c.fetch_add(1, Ordering::SeqCst);
138 Ok(42u32)
139 }
140 })
141 .await;
142 assert_eq!(result.unwrap(), 42);
143 assert_eq!(counter.load(Ordering::SeqCst), 1);
144 }
145
146 #[tokio::test]
147 async fn test_retry_retries_on_retryable_error() {
148 let policy = RetryPolicy {
149 max_retries: 3,
150 initial_delay: Duration::from_millis(1),
151 max_delay: Duration::from_millis(5),
152 jitter: false,
153 ..Default::default()
154 };
155 let counter = Arc::new(AtomicU32::new(0));
156 let c = counter.clone();
157 let result: Result<u32, DbError> = retry_with_backoff(&policy, || {
158 let c = c.clone();
159 async move {
160 let n = c.fetch_add(1, Ordering::SeqCst);
161 if n < 2 {
162 Err(DbError::ConnectionError("timeout".to_string()))
163 } else {
164 Ok(42u32)
165 }
166 }
167 })
168 .await;
169 assert_eq!(result.unwrap(), 42);
170 assert_eq!(counter.load(Ordering::SeqCst), 3);
171 }
172
173 #[tokio::test]
174 async fn test_retry_does_not_retry_non_retryable_error() {
175 let policy = RetryPolicy::default();
176 let counter = Arc::new(AtomicU32::new(0));
177 let c = counter.clone();
178 let result: Result<u32, DbError> = retry_with_backoff(&policy, || {
179 let c = c.clone();
180 async move {
181 c.fetch_add(1, Ordering::SeqCst);
182 Err(DbError::QueryError("syntax error".to_string()))
183 }
184 })
185 .await;
186 assert!(result.is_err());
187 assert_eq!(counter.load(Ordering::SeqCst), 1);
189 }
190
191 #[tokio::test]
192 async fn test_retry_exhausted_after_max_retries() {
193 let policy = RetryPolicy {
194 max_retries: 2,
195 initial_delay: Duration::from_millis(1),
196 max_delay: Duration::from_millis(5),
197 jitter: false,
198 ..Default::default()
199 };
200 let counter = Arc::new(AtomicU32::new(0));
201 let c = counter.clone();
202 let result: Result<u32, DbError> = retry_with_backoff(&policy, || {
203 let c = c.clone();
204 async move {
205 c.fetch_add(1, Ordering::SeqCst);
206 Err(DbError::ConnectionError("timeout".to_string()))
207 }
208 })
209 .await;
210 assert!(result.is_err());
211 assert_eq!(counter.load(Ordering::SeqCst), 3);
213 }
214}