Skip to main content

heartbit_core/llm/
error_class.rs

1//! Error classification for LLM API errors — distinguishes retryable from fatal conditions.
2
3use crate::error::Error;
4
5/// Actionable classification of LLM provider errors.
6#[derive(Debug, Clone, Copy, PartialEq, Eq)]
7pub enum ErrorClass {
8    /// The conversation context exceeds the model's context window.
9    ContextOverflow,
10    /// Rate limited (HTTP 429). Already retried by `RetryingProvider`.
11    RateLimited,
12    /// Authentication failure (HTTP 401/403).
13    AuthError,
14    /// Server-side failure (HTTP 500/502/503/529).
15    ServerError,
16    /// Client error that is not overflow (other HTTP 400).
17    InvalidRequest,
18    /// Transport-level failure (`Error::Http`): TCP/DNS/TLS/timeout.
19    /// Treated as transient — the same signal used by `RetryingProvider`.
20    Network,
21    /// Unrecognized error — no actionable recovery.
22    Unknown,
23}
24
25/// Classify an [`Error`] into an actionable [`ErrorClass`].
26///
27/// Primarily useful for `Error::Api` errors where the HTTP status code and
28/// message body determine recovery strategy.
29pub fn classify(error: &Error) -> ErrorClass {
30    // Unwrap WithPartialUsage to classify the inner error.
31    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
59/// Check if an error message indicates context overflow.
60///
61/// Uses case-insensitive substring matching (no regex dependency).
62fn 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    // Mistral phrasing: "Prompt contains 325070 tokens …, too large for model".
80    // Compound match (both fragments required) so a plain "prompt contains"
81    // in an unrelated validation error can't false-positive.
82    lower.contains("prompt contains") && lower.contains("token")
83}
84
85#[cfg(test)]
86mod tests {
87    use super::*;
88
89    // --- Context overflow (regression pins) ---
90
91    #[test]
92    fn classify_mistral_openrouter_overflow_as_context_overflow() {
93        // EXACT error body from TUI session 6a24e4bb-4159588 (2026-06-07):
94        // OpenRouter wraps Mistral's raw error; the nested `raw` payload carries
95        // "maximum context length", which the substring patterns must match.
96        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        // Same Mistral message WITHOUT the OpenRouter wrapper (direct API) —
106        // must classify on the "prompt contains … tokens" phrasing alone.
107        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    // --- Auth errors ---
115
116    #[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    // --- Rate limited ---
135
136    #[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    // --- Server errors ---
146
147    #[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    // --- Context overflow (400 with overflow patterns) ---
184
185    #[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    /// `max_tokens` in a 400 message can mean parameter validation (e.g.,
231    /// "max_tokens: 4096 must be less than ..."), not context overflow.
232    /// We should NOT classify it as ContextOverflow.
233    #[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    // --- Invalid request (400 without overflow pattern) ---
288
289    #[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    // --- HTTP / network errors ---
299
300    #[test]
301    fn classify_http_error_as_network() {
302        // Build a reqwest error by making a request to an invalid URL.
303        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    // --- Other error variants ---
315
316    #[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    // --- WithPartialUsage unwrapping ---
347
348    #[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    // --- Unknown status codes ---
377
378    #[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}