1use std::fmt;
2
3use thiserror::Error;
4
5#[derive(Debug, Clone, PartialEq, Eq)]
6pub struct ContextOverflowError {
7 pub provider: String,
8 pub model: Option<String>,
9 pub requested_tokens: Option<u32>,
10 pub max_tokens: Option<u32>,
11 pub message: String,
12}
13
14impl ContextOverflowError {
15 pub fn new(
16 provider: impl Into<String>,
17 model: Option<String>,
18 requested_tokens: Option<u32>,
19 max_tokens: Option<u32>,
20 message: impl Into<String>,
21 ) -> Self {
22 Self { provider: provider.into(), model, requested_tokens, max_tokens, message: message.into() }
23 }
24}
25
26impl fmt::Display for ContextOverflowError {
27 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
28 let model = self.model.as_deref().unwrap_or("unknown-model");
29 match (self.requested_tokens, self.max_tokens) {
30 (Some(requested), Some(max)) => write!(
31 f,
32 "{} (provider={}, model={}, requested={}, max={})",
33 self.message, self.provider, model, requested, max
34 ),
35 _ => write!(f, "{} (provider={}, model={})", self.message, self.provider, model),
36 }
37 }
38}
39
40#[doc = include_str!("docs/llm_error.md")]
41#[derive(Debug, Error, Clone)]
42pub enum LlmError {
43 #[error("{0} environment variable not set")]
45 MissingApiKey(String),
46 #[error("Invalid API key: {0}")]
48 InvalidApiKey(String),
49 #[error("Failed to create HTTP client: {0}")]
51 HttpClientCreation(String),
52 #[error("API request failed: {0}")]
54 ApiRequest(String),
55 #[error("API error: {0}")]
57 ApiError(String),
58 #[error("Rate limited: {0}")]
60 RateLimited(String),
61 #[error("Server error (status {status:?}): {message}")]
65 ServerError { status: Option<u16>, message: String },
66 #[error("Request timed out: {0}")]
68 Timeout(String),
69 #[error("Network error: {0}")]
71 Network(String),
72 #[error("Stream interrupted: {0}")]
74 StreamInterrupted(String),
75 #[error("Context overflow: {0}")]
77 ContextOverflow(ContextOverflowError),
78 #[error("IO error reading stream: {0}")]
80 IoError(String),
81 #[error("JSON parsing error: {0}")]
83 JsonParsing(String),
84 #[error("Failed to parse tool parameters for {tool_name}: {error}")]
86 ToolParameterParsing { tool_name: String, error: String },
87 #[error("OAuth error: {0}")]
89 OAuthError(String),
90 #[error("Unsupported content: {0}")]
92 UnsupportedContent(String),
93 #[error("Provider '{provider}' requires a URL configured via providers.{provider}.url")]
95 MissingProviderUrl { provider: String },
96 #[error("{0}")]
98 Other(String),
99}
100
101impl LlmError {
102 pub fn is_retryable(&self) -> bool {
106 matches!(
107 self,
108 LlmError::RateLimited(_)
109 | LlmError::ServerError { .. }
110 | LlmError::Timeout(_)
111 | LlmError::Network(_)
112 | LlmError::StreamInterrupted(_)
113 )
114 }
115}
116
117impl From<reqwest::Error> for LlmError {
118 fn from(error: reqwest::Error) -> Self {
119 if error.is_timeout() {
120 return LlmError::Timeout(error.to_string());
121 }
122 if error.is_connect() || error.is_request() {
123 return LlmError::Network(error.to_string());
124 }
125 match error.status().map(|s| s.as_u16()) {
126 Some(429) => LlmError::RateLimited(error.to_string()),
127 Some(s) if (500..600).contains(&s) => LlmError::ServerError { status: Some(s), message: error.to_string() },
128 _ => LlmError::ApiRequest(error.to_string()),
129 }
130 }
131}
132
133impl From<serde_json::Error> for LlmError {
134 fn from(error: serde_json::Error) -> Self {
135 LlmError::JsonParsing(error.to_string())
136 }
137}
138
139impl From<std::io::Error> for LlmError {
140 fn from(error: std::io::Error) -> Self {
141 LlmError::IoError(error.to_string())
142 }
143}
144
145impl From<reqwest::header::InvalidHeaderValue> for LlmError {
146 fn from(error: reqwest::header::InvalidHeaderValue) -> Self {
147 LlmError::InvalidApiKey(error.to_string())
148 }
149}
150
151impl From<async_openai::error::OpenAIError> for LlmError {
152 fn from(error: async_openai::error::OpenAIError) -> Self {
153 use async_openai::error::OpenAIError;
154 match error {
155 OpenAIError::Reqwest(e) => LlmError::from(e),
156 OpenAIError::StreamError(e) => LlmError::StreamInterrupted(e.to_string()),
157 OpenAIError::ApiError(api_err) => LlmError::ApiError(api_err.to_string()),
158 OpenAIError::JSONDeserialize(e, _) => LlmError::JsonParsing(e.to_string()),
159 OpenAIError::FileSaveError(s) | OpenAIError::FileReadError(s) => LlmError::IoError(s),
160 OpenAIError::InvalidArgument(s) => LlmError::Other(s),
161 }
162 }
163}
164
165#[cfg(feature = "codex")]
166impl From<aether_auth::OAuthError> for LlmError {
167 fn from(error: aether_auth::OAuthError) -> Self {
168 LlmError::OAuthError(error.to_string())
169 }
170}
171
172pub type Result<T> = std::result::Result<T, LlmError>;
173
174#[cfg(test)]
175mod tests {
176 use super::*;
177
178 #[test]
179 fn is_retryable() {
180 assert!(LlmError::RateLimited("rl".into()).is_retryable());
181 assert!(LlmError::ServerError { status: Some(503), message: "x".into() }.is_retryable());
182 assert!(LlmError::ServerError { status: None, message: "stream-level".into() }.is_retryable());
183 assert!(LlmError::Timeout("t".into()).is_retryable());
184 assert!(LlmError::Network("n".into()).is_retryable());
185 assert!(LlmError::StreamInterrupted("s".into()).is_retryable());
186
187 assert!(!LlmError::ApiError("x".into()).is_retryable());
188 assert!(!LlmError::ApiRequest("x".into()).is_retryable());
189 assert!(!LlmError::MissingApiKey("x".into()).is_retryable());
190 assert!(!LlmError::InvalidApiKey("x".into()).is_retryable());
191 assert!(!LlmError::HttpClientCreation("x".into()).is_retryable());
192 assert!(!LlmError::IoError("x".into()).is_retryable());
193 assert!(!LlmError::JsonParsing("x".into()).is_retryable());
194 assert!(!LlmError::ToolParameterParsing { tool_name: "t".into(), error: "e".into() }.is_retryable());
195 assert!(!LlmError::OAuthError("x".into()).is_retryable());
196 assert!(!LlmError::UnsupportedContent("x".into()).is_retryable());
197 assert!(!LlmError::MissingProviderUrl { provider: "azure-foundry".into() }.is_retryable());
198 assert!(!LlmError::Other("x".into()).is_retryable());
199 assert!(!LlmError::ContextOverflow(ContextOverflowError::new("p", None, None, None, "m")).is_retryable());
200 }
201}