Skip to main content

xz_provider/
error.rs

1use std::time::Duration;
2
3use thiserror::Error;
4
5use crate::types::CapabilityRequest;
6
7/// Provider 错误类型
8///
9/// 分类:
10/// - Transient(可重试):Network, RateLimit, Overloaded, Timeout
11/// - Fatal(不可重试):Cancelled, Config, Auth, TokenLimit, Format, 等
12#[derive(Debug, Error)]
13pub enum ProviderError {
14    // ── 可重试错误(Transient)──
15    #[error("网络错误: {message}")]
16    Network { message: String, detail: Option<String> },
17
18    #[error("限流,{retry_after_ms}ms 后重试")]
19    RateLimit { retry_after_ms: u64 },
20
21    #[error("Provider 过载 (503),稍后重试")]
22    Overloaded,
23
24    #[error("超时 ({timeout_ms}ms)")]
25    Timeout { timeout_ms: u64 },
26
27    // ── 不可重试错误(Fatal)──
28    #[error("请求被取消")]
29    Cancelled,
30
31    #[error("配置错误: {0}")]
32    Config(String),
33
34    #[error("认证失败: {0}")]
35    Auth(String),
36
37    #[error("Token 超限: prompt={actual}, limit={limit}")]
38    TokenLimit { actual: usize, limit: usize },
39
40    #[error("输出格式错误: {0}")]
41    Format(String),
42
43    #[error("模型不存在: {0}")]
44    ModelNotFound(String),
45
46    #[error("不支持的能力: {model} 不支持 {capability}")]
47    UnsupportedCapability { model: String, capability: String },
48
49    #[error("没有模型满足能力要求: {required:?}")]
50    NoModelForCapability { required: CapabilityRequest },
51
52    #[error("Provider 错误 [{status}]: {message}")]
53    Internal { status: u16, message: String },
54
55    #[error("没有可用的 Provider 路由: {0}")]
56    NoRoute(String),
57
58    #[error("无法获取 API Key: {0}")]
59    KeySource(String),
60}
61
62impl ProviderError {
63    /// 判断是否为可重试错误
64    pub fn is_retryable(&self) -> bool {
65        matches!(
66            self,
67            ProviderError::Network { .. }
68                | ProviderError::RateLimit { .. }
69                | ProviderError::Overloaded
70                | ProviderError::Timeout { .. }
71        )
72    }
73
74    /// 判断是否为用户主动取消
75    pub fn is_cancelled(&self) -> bool {
76        matches!(self, ProviderError::Cancelled)
77    }
78
79    /// 建议的重试等待时间
80    pub fn retry_after(&self) -> Option<Duration> {
81        match self {
82            ProviderError::RateLimit { retry_after_ms } => {
83                Some(Duration::from_millis(*retry_after_ms))
84            }
85            ProviderError::Overloaded => Some(Duration::from_secs(5)),
86            _ => None,
87        }
88    }
89}
90
91/// 重试策略
92#[derive(Debug, Clone)]
93pub struct RetryStrategy {
94    pub max_retries: u32,
95    pub base_delay: Duration,
96    pub max_delay: Duration,
97    pub jitter: bool,
98}
99
100impl Default for RetryStrategy {
101    fn default() -> Self {
102        Self {
103            max_retries: 3,
104            base_delay: Duration::from_secs(1),
105            max_delay: Duration::from_secs(30),
106            jitter: true,
107        }
108    }
109}
110
111impl RetryStrategy {
112    pub fn new(max_retries: u32, base_delay: Duration, max_delay: Duration, jitter: bool) -> Self {
113        Self { max_retries, base_delay, max_delay, jitter }
114    }
115
116    /// 计算第 n 次重试的等待时间(指数退避),不包含抖动时基数为 2^attempt
117    pub fn delay_for_attempt(&self, attempt: u32) -> Duration {
118        let base_ms = self.base_delay.as_millis() as f64;
119        let delay = base_ms * 2u64.pow(attempt.min(30)) as f64;
120        let delay = delay.min(self.max_delay.as_millis() as f64);
121        if self.jitter {
122            let jitter = rand_dummy();
123            Duration::from_millis((delay * (0.5 + jitter * 0.5)) as u64)
124        } else {
125            Duration::from_millis(delay as u64)
126        }
127    }
128}
129
130fn rand_dummy() -> f64 {
131    (std::time::SystemTime::now()
132        .duration_since(std::time::UNIX_EPOCH)
133        .unwrap_or_default()
134        .subsec_nanos() as f64
135        / 1_000_000_000.0)
136        .fract()
137}
138
139impl From<reqwest::Error> for ProviderError {
140    fn from(e: reqwest::Error) -> Self {
141        if e.is_timeout() {
142            ProviderError::Timeout { timeout_ms: 0 }
143        } else if e.is_connect() || e.is_request() {
144            ProviderError::Network { message: e.to_string(), detail: None }
145        } else if let Some(status) = e.status() {
146            if status.as_u16() == 503 {
147                ProviderError::Overloaded
148            } else if status.as_u16() == 429 {
149                ProviderError::RateLimit { retry_after_ms: 0 }
150            } else if status.as_u16() == 401 || status.as_u16() == 403 {
151                ProviderError::Auth(e.to_string())
152            } else {
153                ProviderError::Internal { status: status.as_u16(), message: e.to_string() }
154            }
155        } else {
156            ProviderError::Network { message: e.to_string(), detail: None }
157        }
158    }
159}
160
161impl From<serde_json::Error> for ProviderError {
162    fn from(e: serde_json::Error) -> Self {
163        ProviderError::Format(e.to_string())
164    }
165}
166
167#[cfg(test)]
168mod tests {
169    use super::*;
170
171    #[test]
172    fn test_is_retryable_network() {
173        let err = ProviderError::Network { message: "refused".into(), detail: None };
174        assert!(err.is_retryable());
175        assert!(!err.is_cancelled());
176    }
177
178    #[test]
179    fn test_is_retryable_rate_limit() {
180        let err = ProviderError::RateLimit { retry_after_ms: 5000 };
181        assert!(err.is_retryable());
182    }
183
184    #[test]
185    fn test_is_retryable_overloaded() {
186        let err = ProviderError::Overloaded;
187        assert!(err.is_retryable());
188    }
189
190    #[test]
191    fn test_is_retryable_timeout() {
192        let err = ProviderError::Timeout { timeout_ms: 30000 };
193        assert!(err.is_retryable());
194    }
195
196    #[test]
197    fn test_is_not_retryable() {
198        assert!(!ProviderError::Cancelled.is_retryable());
199        assert!(!ProviderError::Config("bad".into()).is_retryable());
200        assert!(!ProviderError::Auth("denied".into()).is_retryable());
201        assert!(!ProviderError::TokenLimit { actual: 100, limit: 50 }.is_retryable());
202        assert!(!ProviderError::Format("bad json".into()).is_retryable());
203        assert!(!ProviderError::ModelNotFound("gpt-3".into()).is_retryable());
204        assert!(
205            !ProviderError::UnsupportedCapability {
206                model: "gpt".into(),
207                capability: "vision".into()
208            }
209            .is_retryable()
210        );
211        assert!(
212            !ProviderError::NoModelForCapability { required: CapabilityRequest::default() }
213                .is_retryable()
214        );
215        assert!(!ProviderError::Internal { status: 500, message: "error".into() }.is_retryable());
216        assert!(!ProviderError::NoRoute("nowhere".into()).is_retryable());
217    }
218
219    #[test]
220    fn test_is_cancelled() {
221        assert!(ProviderError::Cancelled.is_cancelled());
222        assert!(!ProviderError::Config("test".into()).is_cancelled());
223        assert!(!ProviderError::Network { message: "x".into(), detail: None }.is_cancelled());
224    }
225
226    #[test]
227    fn test_retry_after_rate_limit() {
228        let err = ProviderError::RateLimit { retry_after_ms: 2000 };
229        assert_eq!(err.retry_after(), Some(Duration::from_millis(2000)));
230    }
231
232    #[test]
233    fn test_retry_after_overloaded() {
234        let err = ProviderError::Overloaded;
235        assert_eq!(err.retry_after(), Some(Duration::from_secs(5)));
236    }
237
238    #[test]
239    fn test_retry_after_fatal_errors() {
240        let fatal_errors = vec![
241            ProviderError::Network { message: "x".into(), detail: None },
242            ProviderError::Timeout { timeout_ms: 1000 },
243            ProviderError::Cancelled,
244            ProviderError::Config("x".into()),
245            ProviderError::Auth("x".into()),
246        ];
247        for err in fatal_errors {
248            assert!(err.retry_after().is_none(), "{:?} should not have retry_after", err);
249        }
250    }
251
252    #[test]
253    fn test_error_display_network() {
254        let err = ProviderError::Network { message: "refused".into(), detail: None };
255        let msg = format!("{}", err);
256        assert!(msg.contains("refused"));
257    }
258
259    #[test]
260    fn test_error_display_rate_limit() {
261        let err = ProviderError::RateLimit { retry_after_ms: 5000 };
262        let msg = format!("{}", err);
263        assert!(msg.contains("5000"));
264    }
265
266    #[test]
267    fn test_retry_strategy_default() {
268        let s = RetryStrategy::default();
269        assert_eq!(s.max_retries, 3);
270        assert_eq!(s.base_delay, Duration::from_secs(1));
271        assert_eq!(s.max_delay, Duration::from_secs(30));
272        assert!(s.jitter);
273    }
274
275    #[test]
276    fn test_retry_strategy_exponential_no_jitter() {
277        let s = RetryStrategy::new(5, Duration::from_millis(100), Duration::from_secs(10), false);
278        assert_eq!(s.delay_for_attempt(0), Duration::from_millis(100));
279        assert_eq!(s.delay_for_attempt(1), Duration::from_millis(200));
280        assert_eq!(s.delay_for_attempt(2), Duration::from_millis(400));
281        assert_eq!(s.delay_for_attempt(3), Duration::from_millis(800));
282    }
283
284    #[test]
285    fn test_retry_strategy_max_delay_cap() {
286        let s = RetryStrategy::new(10, Duration::from_secs(1), Duration::from_millis(500), false);
287        // 2^10 * 1000ms exceeds 500ms cap
288        assert_eq!(s.delay_for_attempt(10), Duration::from_millis(500));
289    }
290
291    #[test]
292    fn test_retry_strategy_jitter_within_bounds() {
293        let s = RetryStrategy::new(3, Duration::from_millis(100), Duration::from_secs(10), true);
294        for attempt in 0..5 {
295            let delay = s.delay_for_attempt(attempt);
296            let max_delay = (100.0 * 2u64.pow(attempt.min(30)) as f64).min(10_000.0);
297            let ms = delay.as_millis() as f64;
298            assert!(ms >= max_delay * 0.4, "attempt={} ms={} min={}", attempt, ms, max_delay * 0.4);
299            assert!(ms <= max_delay * 1.1, "attempt={} ms={} max={}", attempt, ms, max_delay * 1.1);
300        }
301    }
302
303    #[test]
304    fn test_retry_strategy_display() {
305        let s = RetryStrategy::default();
306        assert_eq!(s.max_retries, 3);
307    }
308}