1use crate::error::Error;
4
5#[derive(Debug, Clone, Copy, PartialEq, Eq)]
7pub enum ErrorClass {
8 ContextOverflow,
10 RateLimited,
12 AuthError,
14 ServerError,
16 InvalidRequest,
18 Network,
21 Unknown,
23}
24
25pub fn classify(error: &Error) -> ErrorClass {
30 let inner = match error {
32 Error::WithPartialUsage { source, .. } => source.as_ref(),
33 other => other,
34 };
35
36 match inner {
37 Error::Api { status, message } => classify_api(*status, message),
38 Error::Http(_) => ErrorClass::Network,
39 _ => ErrorClass::Unknown,
40 }
41}
42
43fn classify_api(status: u16, message: &str) -> ErrorClass {
44 match status {
45 401 | 403 => ErrorClass::AuthError,
46 429 => ErrorClass::RateLimited,
47 500 | 502 | 503 | 529 => ErrorClass::ServerError,
48 400 => {
49 if is_context_overflow(message) {
50 ErrorClass::ContextOverflow
51 } else {
52 ErrorClass::InvalidRequest
53 }
54 }
55 _ => ErrorClass::Unknown,
56 }
57}
58
59fn is_context_overflow(message: &str) -> bool {
63 const PATTERNS: &[&str] = &[
64 "prompt is too long",
65 "maximum context length",
66 "context_length_exceeded",
67 "context window",
68 "too many tokens",
69 "input is too long",
70 "exceeds the model's maximum context",
71 "request too large",
72 "content too large",
73 ];
74
75 let lower = message.to_lowercase();
76 if PATTERNS.iter().any(|p| lower.contains(p)) {
77 return true;
78 }
79 lower.contains("prompt contains") && lower.contains("token")
83}
84
85#[cfg(test)]
86mod tests {
87 use super::*;
88
89 #[test]
92 fn classify_mistral_openrouter_overflow_as_context_overflow() {
93 let err = Error::Api {
97 status: 400,
98 message: r#"{"error":{"message":"Provider returned error","code":400,"metadata":{"raw":"{\"object\":\"error\",\"message\":\"Prompt contains 325070 tokens and 0 draft tokens, too large for model with 262144 maximum context length\",\"type\":\"invalid_request_invalid_args\",\"param\":null,\"code\":\"3051\",\"raw_status_code\":400}","provider_name":"Mistral","is_byok":false}},"user_id":"user_x"}"#.into(),
99 };
100 assert_eq!(classify(&err), ErrorClass::ContextOverflow);
101 }
102
103 #[test]
104 fn classify_bare_prompt_contains_tokens_as_context_overflow() {
105 let err = Error::Api {
108 status: 400,
109 message: "Prompt contains 325070 tokens and 0 draft tokens, too large for model".into(),
110 };
111 assert_eq!(classify(&err), ErrorClass::ContextOverflow);
112 }
113
114 #[test]
117 fn classify_401_as_auth_error() {
118 let err = Error::Api {
119 status: 401,
120 message: "Unauthorized".into(),
121 };
122 assert_eq!(classify(&err), ErrorClass::AuthError);
123 }
124
125 #[test]
126 fn classify_403_as_auth_error() {
127 let err = Error::Api {
128 status: 403,
129 message: "Forbidden".into(),
130 };
131 assert_eq!(classify(&err), ErrorClass::AuthError);
132 }
133
134 #[test]
137 fn classify_429_as_rate_limited() {
138 let err = Error::Api {
139 status: 429,
140 message: "Too Many Requests".into(),
141 };
142 assert_eq!(classify(&err), ErrorClass::RateLimited);
143 }
144
145 #[test]
148 fn classify_500_as_server_error() {
149 let err = Error::Api {
150 status: 500,
151 message: "Internal Server Error".into(),
152 };
153 assert_eq!(classify(&err), ErrorClass::ServerError);
154 }
155
156 #[test]
157 fn classify_502_as_server_error() {
158 let err = Error::Api {
159 status: 502,
160 message: "Bad Gateway".into(),
161 };
162 assert_eq!(classify(&err), ErrorClass::ServerError);
163 }
164
165 #[test]
166 fn classify_503_as_server_error() {
167 let err = Error::Api {
168 status: 503,
169 message: "Service Unavailable".into(),
170 };
171 assert_eq!(classify(&err), ErrorClass::ServerError);
172 }
173
174 #[test]
175 fn classify_529_as_server_error() {
176 let err = Error::Api {
177 status: 529,
178 message: "Overloaded".into(),
179 };
180 assert_eq!(classify(&err), ErrorClass::ServerError);
181 }
182
183 #[test]
186 fn classify_400_prompt_too_long() {
187 let err = Error::Api {
188 status: 400,
189 message: "prompt is too long".into(),
190 };
191 assert_eq!(classify(&err), ErrorClass::ContextOverflow);
192 }
193
194 #[test]
195 fn classify_400_maximum_context_length() {
196 let err = Error::Api {
197 status: 400,
198 message: "This request exceeds the maximum context length".into(),
199 };
200 assert_eq!(classify(&err), ErrorClass::ContextOverflow);
201 }
202
203 #[test]
204 fn classify_400_context_length_exceeded() {
205 let err = Error::Api {
206 status: 400,
207 message: "context_length_exceeded".into(),
208 };
209 assert_eq!(classify(&err), ErrorClass::ContextOverflow);
210 }
211
212 #[test]
213 fn classify_400_request_too_large() {
214 let err = Error::Api {
215 status: 400,
216 message: "request too large for this model".into(),
217 };
218 assert_eq!(classify(&err), ErrorClass::ContextOverflow);
219 }
220
221 #[test]
222 fn classify_400_content_too_large() {
223 let err = Error::Api {
224 status: 400,
225 message: "content too large".into(),
226 };
227 assert_eq!(classify(&err), ErrorClass::ContextOverflow);
228 }
229
230 #[test]
234 fn classify_400_max_tokens_parameter_is_not_overflow() {
235 let err = Error::Api {
236 status: 400,
237 message: "max_tokens: 4096 must be less than 2048".into(),
238 };
239 assert_eq!(classify(&err), ErrorClass::InvalidRequest);
240 }
241
242 #[test]
243 fn classify_400_context_window() {
244 let err = Error::Api {
245 status: 400,
246 message: "exceeds the context window".into(),
247 };
248 assert_eq!(classify(&err), ErrorClass::ContextOverflow);
249 }
250
251 #[test]
252 fn classify_400_too_many_tokens() {
253 let err = Error::Api {
254 status: 400,
255 message: "too many tokens in the request".into(),
256 };
257 assert_eq!(classify(&err), ErrorClass::ContextOverflow);
258 }
259
260 #[test]
261 fn classify_400_input_too_long() {
262 let err = Error::Api {
263 status: 400,
264 message: "input is too long for model".into(),
265 };
266 assert_eq!(classify(&err), ErrorClass::ContextOverflow);
267 }
268
269 #[test]
270 fn classify_400_exceeds_model_maximum_context() {
271 let err = Error::Api {
272 status: 400,
273 message: "exceeds the model's maximum context length".into(),
274 };
275 assert_eq!(classify(&err), ErrorClass::ContextOverflow);
276 }
277
278 #[test]
279 fn classify_400_case_insensitive() {
280 let err = Error::Api {
281 status: 400,
282 message: "PROMPT IS TOO LONG".into(),
283 };
284 assert_eq!(classify(&err), ErrorClass::ContextOverflow);
285 }
286
287 #[test]
290 fn classify_400_generic_as_invalid_request() {
291 let err = Error::Api {
292 status: 400,
293 message: "invalid parameter: temperature must be between 0 and 1".into(),
294 };
295 assert_eq!(classify(&err), ErrorClass::InvalidRequest);
296 }
297
298 #[test]
301 fn classify_http_error_as_network() {
302 let rt = tokio::runtime::Builder::new_current_thread()
304 .enable_all()
305 .build()
306 .expect("test runtime");
307 let http_err = rt
308 .block_on(reqwest::get("http://[::0]:1"))
309 .expect_err("should fail");
310 let err = Error::Http(http_err);
311 assert_eq!(classify(&err), ErrorClass::Network);
312 }
313
314 #[test]
317 fn classify_agent_error_as_unknown() {
318 let err = Error::Agent("something went wrong".into());
319 assert_eq!(classify(&err), ErrorClass::Unknown);
320 }
321
322 #[test]
323 fn classify_max_turns_exceeded_as_unknown() {
324 let err = Error::MaxTurnsExceeded(10);
325 assert_eq!(classify(&err), ErrorClass::Unknown);
326 }
327
328 #[test]
329 fn classify_truncated_as_unknown() {
330 let err = Error::Truncated;
331 assert_eq!(classify(&err), ErrorClass::Unknown);
332 }
333
334 #[test]
335 fn classify_config_error_as_unknown() {
336 let err = Error::Config("bad config".into());
337 assert_eq!(classify(&err), ErrorClass::Unknown);
338 }
339
340 #[test]
341 fn classify_mcp_error_as_unknown() {
342 let err = Error::Mcp("connection refused".into());
343 assert_eq!(classify(&err), ErrorClass::Unknown);
344 }
345
346 #[test]
349 fn classify_unwraps_with_partial_usage() {
350 use crate::llm::types::TokenUsage;
351
352 let inner = Error::Api {
353 status: 429,
354 message: "rate limited".into(),
355 };
356 let wrapped = inner.with_partial_usage(TokenUsage {
357 input_tokens: 100,
358 output_tokens: 50,
359 ..Default::default()
360 });
361 assert_eq!(classify(&wrapped), ErrorClass::RateLimited);
362 }
363
364 #[test]
365 fn classify_unwraps_partial_usage_context_overflow() {
366 use crate::llm::types::TokenUsage;
367
368 let inner = Error::Api {
369 status: 400,
370 message: "prompt is too long".into(),
371 };
372 let wrapped = inner.with_partial_usage(TokenUsage::default());
373 assert_eq!(classify(&wrapped), ErrorClass::ContextOverflow);
374 }
375
376 #[test]
379 fn classify_unknown_status_as_unknown() {
380 let err = Error::Api {
381 status: 418,
382 message: "I'm a teapot".into(),
383 };
384 assert_eq!(classify(&err), ErrorClass::Unknown);
385 }
386}