1use std::time::Duration;
2
3use thiserror::Error;
4
5use crate::types::CapabilityRequest;
6
7#[derive(Debug, Error)]
13pub enum ProviderError {
14 #[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 #[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 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 pub fn is_cancelled(&self) -> bool {
76 matches!(self, ProviderError::Cancelled)
77 }
78
79 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#[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 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 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}