Skip to main content

sz_orm_core/
retry.rs

1//! 通用错误重试器
2//!
3//! 支持指数退避 + 抖动的可配置重试策略
4
5use std::future::Future;
6use std::time::Duration;
7
8use crate::DbError;
9
10/// 重试策略
11#[derive(Debug, Clone)]
12pub struct RetryPolicy {
13    /// 最大重试次数
14    pub max_retries: u32,
15    /// 初始延迟
16    pub initial_delay: Duration,
17    /// 最大延迟
18    pub max_delay: Duration,
19    /// 退避因子
20    pub backoff_factor: f64,
21    /// 是否添加抖动
22    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    /// 计算第 n 次重试的延迟
39    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            // 简单抖动:在 50%-150% 范围内随机
45            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
53/// 简单伪随机(不依赖 rand crate)
54fn 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
61/// 带重试执行异步操作
62///
63/// 根据 `RetryPolicy` 重试可重试的 `DbError`,使用指数退避 + 抖动策略。
64pub 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        // 无抖动时:delay(0) = 10ms, delay(1) = 20ms
113        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        // retry=10 → base * 2^10 = 10240ms,应被限制到 50ms
125        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        // 非可重试错误应立即返回,不重试
188        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        // max_retries=2 → 尝试 3 次(attempt 0,1,2)
212        assert_eq!(counter.load(Ordering::SeqCst), 3);
213    }
214}