1use std::fmt;
2
3use thiserror::Error;
4
5#[doc = include_str!("docs/llm_error.md")]
6#[derive(Debug, Error, Clone)]
7pub enum LlmError {
8 #[error(transparent)]
9 ReasoningValidation(#[from] crate::catalog::ReasoningEffortError),
10 #[error("Disabling reasoning is not implemented for model '{model}' on this transport")]
11 UnsupportedDisableTransport { model: String },
12 #[error("{0} environment variable not set")]
14 MissingApiKey(String),
15 #[error("{0}")]
18 Provider(#[from] ProviderError),
19 #[error("IO error reading stream: {0}")]
21 IoError(String),
22 #[error("JSON parsing error: {0}")]
24 JsonParsing(String),
25 #[error("Failed to parse tool parameters for {tool_name}: {error}")]
27 ToolParameterParsing { tool_name: String, error: String },
28 #[error("OAuth error: {0}")]
30 OAuthError(String),
31 #[error("Unsupported content: {0}")]
33 UnsupportedContent(String),
34 #[error("Provider '{provider}' requires a URL configured via providers.{provider}.url")]
36 MissingProviderUrl { provider: String },
37 #[error("Unknown provider: {provider}")]
39 UnknownProvider { provider: String },
40 #[error("No models provided")]
42 EmptyModelSpec,
43 #[error("providers.{provider}.{field} cannot be used with multiple {provider} models in one alloy spec")]
47 DuplicateProvider { provider: String, field: String },
48 #[error("Invalid model spec: {0}")]
50 InvalidModelSpec(String),
51 #[error("Failed to build provider request: {0}")]
54 ProviderRequest(String),
55 #[error("Invalid argument: {0}")]
57 InvalidArgument(String),
58}
59
60#[derive(Debug, Clone, Copy, PartialEq, Eq)]
61pub enum ProviderErrorKind {
62 Authentication,
63 Api,
64 RateLimit,
65 Server,
66 Timeout,
67 Network,
68 StreamInterrupted,
69 Unknown,
70}
71
72impl ProviderErrorKind {
73 pub fn is_retryable(&self) -> bool {
74 matches!(
75 self,
76 Self::RateLimit | Self::Server | Self::Timeout | Self::Network | Self::StreamInterrupted | Self::Unknown
77 )
78 }
79}
80
81#[derive(Debug, Clone, PartialEq, Eq)]
82pub struct ProviderError {
83 pub kind: ProviderErrorKind,
84 pub message: String,
85 pub http_status: Option<u16>,
86 pub request_id: Option<String>,
87 pub code: Option<String>,
88}
89
90impl ProviderError {
91 pub fn new(kind: ProviderErrorKind, message: impl Into<String>) -> Self {
92 Self { kind, message: message.into(), http_status: None, request_id: None, code: None }
93 }
94
95 pub fn authentication(message: impl Into<String>) -> Self {
96 Self::new(ProviderErrorKind::Authentication, message)
97 }
98
99 pub fn api(message: impl Into<String>) -> Self {
100 Self::new(ProviderErrorKind::Api, message)
101 }
102
103 pub fn rate_limit(message: impl Into<String>) -> Self {
104 Self::new(ProviderErrorKind::RateLimit, message)
105 }
106
107 pub fn server(message: impl Into<String>) -> Self {
108 Self::new(ProviderErrorKind::Server, message)
109 }
110
111 pub fn timeout(message: impl Into<String>) -> Self {
112 Self::new(ProviderErrorKind::Timeout, message)
113 }
114
115 pub fn network(message: impl Into<String>) -> Self {
116 Self::new(ProviderErrorKind::Network, message)
117 }
118
119 pub fn stream_interrupted(message: impl Into<String>) -> Self {
120 Self::new(ProviderErrorKind::StreamInterrupted, message)
121 }
122
123 pub fn from_http_status(status: u16, message: impl Into<String>) -> Self {
124 match status {
125 401 | 403 => Self::authentication(message),
126 408 | 504 => Self::timeout(message),
127 429 => Self::rate_limit(message),
128 s if (500..600).contains(&s) => Self::server(message),
129 _ => Self::api(message),
130 }
131 .with_http_status(status)
132 }
133
134 pub fn with_http_status(mut self, status: u16) -> Self {
135 self.http_status = Some(status);
136 self
137 }
138
139 pub fn with_http_metadata(mut self, status: Option<u16>, request_id: Option<String>) -> Self {
140 self.http_status = status;
141 self.request_id = request_id;
142 self
143 }
144
145 pub fn with_code(mut self, code: Option<String>) -> Self {
146 self.code = code;
147 self
148 }
149
150 pub fn with_request_id(mut self, request_id: Option<String>) -> Self {
151 self.request_id = request_id;
152 self
153 }
154
155 pub fn is_retryable(&self) -> bool {
156 self.kind.is_retryable()
157 }
158}
159
160impl fmt::Display for ProviderError {
161 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
162 let prefix = match self.kind {
163 ProviderErrorKind::Authentication => "Authentication error",
164 ProviderErrorKind::Api => "API error",
165 ProviderErrorKind::RateLimit => "Rate limited",
166 ProviderErrorKind::Server => "Server error",
167 ProviderErrorKind::Timeout => "Request timed out",
168 ProviderErrorKind::Network => "Network error",
169 ProviderErrorKind::StreamInterrupted => "Stream interrupted",
170 ProviderErrorKind::Unknown => "Provider error",
171 };
172 write!(f, "{prefix}: {}", self.message)?;
173 let mut diagnostics = Vec::new();
174 if let Some(status) = self.http_status {
175 diagnostics.push(format!("status {status}"));
176 }
177 if let Some(code) = &self.code {
178 diagnostics.push(format!("code {code}"));
179 }
180 if let Some(request_id) = &self.request_id {
181 diagnostics.push(format!("request_id {request_id}"));
182 }
183 if !diagnostics.is_empty() {
184 write!(f, " ({})", diagnostics.join(", "))?;
185 }
186 Ok(())
187 }
188}
189
190impl std::error::Error for ProviderError {}
191
192impl LlmError {
193 pub fn is_retryable(&self) -> bool {
194 self.provider().is_some_and(ProviderError::is_retryable)
195 }
196
197 pub fn provider(&self) -> Option<&ProviderError> {
198 match self {
199 Self::Provider(error) => Some(error),
200 _ => None,
201 }
202 }
203}
204
205impl From<reqwest::Error> for LlmError {
206 fn from(error: reqwest::Error) -> Self {
207 if error.is_timeout() {
208 return ProviderError::timeout(error.to_string()).into();
209 }
210 if error.is_connect() || error.is_request() {
211 return ProviderError::network(error.to_string()).into();
212 }
213 match error.status().map(|s| s.as_u16()) {
214 Some(status) => ProviderError::from_http_status(status, error.to_string()).into(),
215 None => ProviderError::network(error.to_string()).into(),
216 }
217 }
218}
219
220impl From<serde_json::Error> for LlmError {
221 fn from(error: serde_json::Error) -> Self {
222 LlmError::JsonParsing(error.to_string())
223 }
224}
225
226impl From<std::io::Error> for LlmError {
227 fn from(error: std::io::Error) -> Self {
228 LlmError::IoError(error.to_string())
229 }
230}
231
232impl From<reqwest::header::InvalidHeaderValue> for LlmError {
233 fn from(error: reqwest::header::InvalidHeaderValue) -> Self {
234 LlmError::ProviderRequest(error.to_string())
235 }
236}
237
238impl From<async_openai::error::OpenAIError> for LlmError {
239 fn from(error: async_openai::error::OpenAIError) -> Self {
240 use async_openai::error::OpenAIError;
241 match error {
242 OpenAIError::Reqwest(e) => LlmError::from(e),
243 OpenAIError::StreamError(e) => ProviderError::stream_interrupted(e.to_string()).into(),
244 OpenAIError::ApiError(api_err) => {
245 let status = api_err.status_code.as_u16();
246 let code = api_err.api_error.code.clone();
247 let message = format!("{status} {}", api_err.api_error);
248 ProviderError::from_http_status(status, message).with_code(code).into()
249 }
250 OpenAIError::JSONDeserialize(e, _) => LlmError::JsonParsing(e.to_string()),
251 OpenAIError::Boxed(e) => e
252 .downcast::<ProviderError>()
253 .map_or_else(|error| ProviderError::api(error.to_string()).into(), |error| (*error).into()),
254 OpenAIError::FileSaveError(s) | OpenAIError::FileReadError(s) => LlmError::IoError(s),
255 OpenAIError::InvalidArgument(s) => LlmError::InvalidArgument(s),
256 }
257 }
258}
259
260#[cfg(feature = "codex")]
261impl From<aether_auth::OAuthError> for LlmError {
262 fn from(error: aether_auth::OAuthError) -> Self {
263 LlmError::OAuthError(error.to_string())
264 }
265}
266
267pub type Result<T> = std::result::Result<T, LlmError>;
268
269#[cfg(test)]
270mod tests {
271 use super::*;
272
273 #[test]
274 fn retryable_kinds_cover_transient_failures() {
275 assert!(!ProviderErrorKind::Authentication.is_retryable());
276 assert!(!ProviderErrorKind::Api.is_retryable());
277 assert!(ProviderErrorKind::RateLimit.is_retryable());
278 assert!(ProviderErrorKind::Server.is_retryable());
279 assert!(ProviderErrorKind::Timeout.is_retryable());
280 assert!(ProviderErrorKind::Network.is_retryable());
281 assert!(ProviderErrorKind::StreamInterrupted.is_retryable());
282 assert!(ProviderErrorKind::Unknown.is_retryable());
283 }
284
285 #[test]
286 fn http_status_classification_covers_known_statuses() {
287 assert_eq!(ProviderError::from_http_status(401, "d").kind, ProviderErrorKind::Authentication);
288 assert_eq!(ProviderError::from_http_status(403, "d").kind, ProviderErrorKind::Authentication);
289 assert_eq!(ProviderError::from_http_status(408, "d").kind, ProviderErrorKind::Timeout);
290 assert!(ProviderError::from_http_status(408, "d").is_retryable());
291 assert_eq!(ProviderError::from_http_status(429, "d").kind, ProviderErrorKind::RateLimit);
292 assert_eq!(ProviderError::from_http_status(500, "d").kind, ProviderErrorKind::Server);
293 assert_eq!(ProviderError::from_http_status(503, "d").kind, ProviderErrorKind::Server);
294 assert_eq!(ProviderError::from_http_status(400, "d").kind, ProviderErrorKind::Api);
295 assert!(!ProviderError::from_http_status(400, "d").is_retryable());
296 assert_eq!(ProviderError::from_http_status(404, "d").kind, ProviderErrorKind::Api);
297 }
298
299 #[test]
300 fn display_includes_diagnostics_when_present() {
301 let error = ProviderError::server("boom").with_http_status(200).with_code(Some("server_error".into()));
302 let text = error.to_string();
303 assert!(text.contains("boom"));
304 assert!(text.contains("status 200"));
305 assert!(text.contains("code server_error"));
306 let error = error.with_request_id(Some("req-1".into()));
307 assert!(error.to_string().contains("request_id req-1"));
308 }
309
310 #[test]
311 fn display_omits_suffix_when_no_metadata() {
312 let error = ProviderError::api("bad");
313 assert_eq!(error.to_string(), "API error: bad");
314 }
315
316 #[test]
317 fn is_retryable() {
318 assert!(LlmError::from(ProviderError::rate_limit("rl")).is_retryable());
319 assert!(LlmError::from(ProviderError::server("x").with_http_status(503)).is_retryable());
320 assert!(LlmError::from(ProviderError::server("stream-level")).is_retryable());
321 assert!(LlmError::from(ProviderError::timeout("t")).is_retryable());
322 assert!(LlmError::from(ProviderError::network("n")).is_retryable());
323 assert!(LlmError::from(ProviderError::stream_interrupted("s")).is_retryable());
324
325 assert!(!LlmError::from(ProviderError::api("x")).is_retryable());
326 assert!(!LlmError::from(ProviderError::authentication("x")).is_retryable());
327 assert!(!LlmError::MissingApiKey("x".into()).is_retryable());
328 assert!(!LlmError::IoError("x".into()).is_retryable());
329 assert!(!LlmError::JsonParsing("x".into()).is_retryable());
330 assert!(!LlmError::ToolParameterParsing { tool_name: "t".into(), error: "e".into() }.is_retryable());
331 assert!(!LlmError::OAuthError("x".into()).is_retryable());
332 assert!(!LlmError::UnsupportedContent("x".into()).is_retryable());
333 assert!(!LlmError::MissingProviderUrl { provider: "azure-foundry".into() }.is_retryable());
334 assert!(!LlmError::UnknownProvider { provider: "foo".into() }.is_retryable());
335 assert!(!LlmError::EmptyModelSpec.is_retryable());
336 assert!(
337 !LlmError::DuplicateProvider { provider: "bedrock".into(), field: "inferenceProfileArn".into() }
338 .is_retryable()
339 );
340 assert!(!LlmError::InvalidModelSpec("x".into()).is_retryable());
341 assert!(!LlmError::ProviderRequest("x".into()).is_retryable());
342 assert!(!LlmError::InvalidArgument("x".into()).is_retryable());
343 }
344
345 #[test]
346 fn async_openai_api_error_preserves_status_and_code() {
347 use async_openai::error::{ApiError, ApiErrorResponse};
348 let response = ApiErrorResponse {
349 status_code: reqwest::StatusCode::SERVICE_UNAVAILABLE,
350 api_error: ApiError {
351 message: "overloaded".to_string(),
352 r#type: None,
353 param: None,
354 code: Some("server_error".to_string()),
355 misalignment: None,
356 },
357 };
358 let error = LlmError::from(async_openai::error::OpenAIError::ApiError(response));
359 let provider = error.provider().expect("expected provider error");
360 assert_eq!(provider.kind, ProviderErrorKind::Server);
361 assert_eq!(provider.http_status, Some(503));
362 assert_eq!(provider.code.as_deref(), Some("server_error"));
363 assert!(error.is_retryable());
364 }
365
366 #[test]
367 fn async_openai_boxed_provider_error_preserves_classification_and_diagnostics() {
368 let provider = ProviderError::rate_limit("slow down")
369 .with_http_status(429)
370 .with_code(Some("429".into()))
371 .with_request_id(Some("request-123".into()));
372 let error = LlmError::from(async_openai::error::OpenAIError::Boxed(Box::new(provider.clone())));
373
374 assert_eq!(error.provider(), Some(&provider));
375 assert!(error.is_retryable());
376 }
377
378 #[test]
379 fn async_openai_stream_error_is_interruption() {
380 let io = std::io::Error::other("eof");
381 let error = LlmError::from(async_openai::error::OpenAIError::StreamError(Box::new(
382 async_openai::error::StreamError::EventStream(io.to_string()),
383 )));
384 let provider = error.provider().expect("expected provider error");
385 assert_eq!(provider.kind, ProviderErrorKind::StreamInterrupted);
386 assert!(error.is_retryable());
387 }
388
389 #[test]
390 fn async_openai_non_error_body_stays_a_json_parsing_error() {
391 let body = r#"{"id":"chatcmpl-1","object":"chat.completion","choices":[]}"#;
392 let parse_error = serde_json::from_str::<String>(body).unwrap_err();
393
394 let error = LlmError::from(async_openai::error::OpenAIError::JSONDeserialize(parse_error, body.to_string()));
395
396 assert!(matches!(error, LlmError::JsonParsing(_)), "got {error:?}");
397 assert!(!error.is_retryable());
398 }
399}