1use serde::{Deserialize, Serialize};
2
3#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
5#[non_exhaustive]
6pub struct RateLimitError {
7 pub retry_after: Option<u64>,
9 pub limit: Option<u64>,
11 pub remaining: Option<u64>,
13 pub reset: Option<u64>,
15}
16
17#[derive(Debug, Clone, PartialEq, Eq)]
19pub struct ClientError {
20 pub message: String,
21 pub code: u16,
22 pub rate_limit: Option<RateLimitError>,
23}
24
25impl std::fmt::Display for ClientError {
26 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
27 if let Some(ref rl) = self.rate_limit {
28 write!(
29 f,
30 "ClientError (code {}): {} [rate_limit: retry_after={:?}, limit={:?}, remaining={:?}, reset={:?}]",
31 self.code, self.message, rl.retry_after, rl.limit, rl.remaining, rl.reset
32 )
33 } else {
34 write!(f, "ClientError (code {}): {}", self.code, self.message)
35 }
36 }
37}
38
39impl std::error::Error for ClientError {}
40
41impl ClientError {
42 pub fn new(code: u16, message: impl Into<String>) -> Self {
44 Self {
45 code,
46 message: message.into(),
47 rate_limit: None,
48 }
49 }
50
51 pub fn with_rate_limit(
53 code: u16,
54 message: impl Into<String>,
55 rate_limit: RateLimitError,
56 ) -> Self {
57 Self {
58 code,
59 message: message.into(),
60 rate_limit: Some(rate_limit),
61 }
62 }
63
64 pub fn is_auth(&self) -> bool {
66 self.code == 401 || self.code == 403
67 }
68
69 pub fn is_rate_limited(&self) -> bool {
71 self.code == 429
72 }
73
74 pub fn is_not_found(&self) -> bool {
76 self.code == 404
77 }
78
79 pub fn is_server_error(&self) -> bool {
81 (500..600).contains(&self.code)
82 }
83
84 pub fn is_retryable(&self) -> bool {
86 self.is_rate_limited()
87 || self.is_server_error()
88 || (self.code == 0
89 && (self.message.to_ascii_lowercase().contains("timed out")
90 || self.message.to_ascii_lowercase().contains("timeout")
91 || self.message.to_ascii_lowercase().contains("connection reset")
92 || self.message.to_ascii_lowercase().contains("network stream error")))
93 }
94}
95
96pub fn parse_retry_after(s: &str) -> Option<u64> {
98 let trimmed = s.trim();
99 if trimmed.is_empty() {
100 return None;
101 }
102 if let Ok(secs) = trimmed.parse::<u64>() {
103 return Some(secs);
104 }
105 if let Ok(system_time) = httpdate::parse_http_date(trimmed) {
106 let now = std::time::SystemTime::now();
107 let secs = system_time
108 .duration_since(now)
109 .map(|d| d.as_secs())
110 .unwrap_or(0);
111 return Some(secs);
112 }
113 None
114}
115
116pub fn extract_rate_limit_headers(headers: &reqwest::header::HeaderMap) -> Option<RateLimitError> {
118 let parse_u64 = |keys: &[&str]| -> Option<u64> {
119 for &k in keys {
120 if let Some(val) = headers.get(k) {
121 if let Ok(s) = val.to_str() {
122 let trimmed = s.trim();
123 if !trimmed.is_empty() {
124 if let Ok(n) = trimmed.parse::<u64>() {
125 return Some(n);
126 }
127 }
128 }
129 }
130 }
131 None
132 };
133
134 let parse_retry_after_from_headers = |keys: &[&str]| -> Option<u64> {
135 for &k in keys {
136 if let Some(val) = headers.get(k) {
137 if let Ok(s) = val.to_str() {
138 if let Some(secs) = parse_retry_after(s) {
139 return Some(secs);
140 }
141 }
142 }
143 }
144 None
145 };
146
147 let retry_after = parse_retry_after_from_headers(&["retry-after", "x-retry-after"]);
148 let limit = parse_u64(&["ratelimit-limit", "x-ratelimit-limit", "x-rate-limit-limit"]);
149 let remaining = parse_u64(&["ratelimit-remaining", "x-ratelimit-remaining", "x-rate-limit-remaining"]);
150 let reset = parse_u64(&["ratelimit-reset", "x-ratelimit-reset", "x-rate-limit-reset"]);
151
152 if retry_after.is_some() || limit.is_some() || remaining.is_some() || reset.is_some() {
153 Some(RateLimitError {
154 retry_after,
155 limit,
156 remaining,
157 reset,
158 })
159 } else {
160 None
161 }
162}
163
164
165#[cfg(test)]
166mod tests {
167 use super::*;
168
169 #[test]
170 fn test_client_error_display() {
171 let err = ClientError::new(404, "Not Found");
172 assert_eq!(format!("{}", err), "ClientError (code 404): Not Found");
173 }
174
175 #[test]
176 fn test_client_error_debug() {
177 let err = ClientError::new(500, "Internal Error");
178 let debug_str = format!("{:?}", err);
179 assert!(debug_str.contains("500"));
180 assert!(debug_str.contains("Internal Error"));
181 }
182
183 #[test]
184 fn test_client_error_classification_methods() {
185 let auth_err = ClientError::new(401, "Unauthorized");
186 assert!(auth_err.is_auth());
187 assert!(!auth_err.is_server_error());
188 assert!(!auth_err.is_retryable());
189
190 let forbidden_err = ClientError::new(403, "Forbidden");
191 assert!(forbidden_err.is_auth());
192
193 let rate_err = ClientError::new(429, "Too Many Requests");
194 assert!(rate_err.is_rate_limited());
195 assert!(rate_err.is_retryable());
196
197 let server_err = ClientError::new(503, "Service Unavailable");
198 assert!(server_err.is_server_error());
199 assert!(server_err.is_retryable());
200
201 let timeout_err = ClientError::new(0, "operation timed out after 30s");
202 assert!(timeout_err.is_retryable());
203 }
204
205 #[test]
206 fn test_client_error_clone_and_eq() {
207 let err1 = ClientError::new(400, "Bad Request");
208 let err2 = err1.clone();
209 assert_eq!(err1, err2);
210 assert_eq!(err1.code, 400);
211 assert_eq!(err1.message, "Bad Request");
212
213 let err3 = ClientError::new(401, "Unauthorized");
214 assert_ne!(err1, err3);
215 }
216
217 #[test]
218 fn test_client_error_implements_std_error() {
219 let err: Box<dyn std::error::Error> = Box::new(ClientError::new(403, "Forbidden"));
220 assert_eq!(format!("{}", err), "ClientError (code 403): Forbidden");
221 assert!(err.source().is_none());
222 }
223
224 #[test]
225 fn test_extract_rate_limit_headers() {
226 let mut headers = reqwest::header::HeaderMap::new();
227 headers.insert("Retry-After", "60".parse().unwrap());
228 headers.insert("RateLimit-Limit", "1000".parse().unwrap());
229 headers.insert("RateLimit-Remaining", "5".parse().unwrap());
230 headers.insert("RateLimit-Reset", "1700000000".parse().unwrap());
231
232 let rl = extract_rate_limit_headers(&headers).expect("should extract rate limit headers");
233 assert_eq!(rl.retry_after, Some(60));
234 assert_eq!(rl.limit, Some(1000));
235 assert_eq!(rl.remaining, Some(5));
236 assert_eq!(rl.reset, Some(1700000000));
237
238 let err_with_rl = ClientError::with_rate_limit(429, "Rate limit exceeded", rl);
239 assert!(err_with_rl.is_rate_limited());
240 assert!(err_with_rl.rate_limit.is_some());
241 assert_eq!(err_with_rl.rate_limit.as_ref().unwrap().retry_after, Some(60));
242 assert!(format!("{}", err_with_rl).contains("[rate_limit: retry_after=Some(60)"));
243 }
244
245 #[test]
246 fn test_parse_retry_after_http_date() {
247 assert_eq!(parse_retry_after("120"), Some(120));
248 assert_eq!(parse_retry_after("Wed, 21 Oct 2015 07:28:00 GMT"), Some(0));
249
250 let mut headers = reqwest::header::HeaderMap::new();
251 headers.insert("Retry-After", "Wed, 21 Oct 2015 07:28:00 GMT".parse().unwrap());
252 let rl = extract_rate_limit_headers(&headers).expect("should parse HTTP-date Retry-After");
253 assert_eq!(rl.retry_after, Some(0));
254 }
255}
256
257