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("Unknown provider: {provider}")]
98 UnknownProvider { provider: String },
99 #[error("No models provided")]
101 EmptyModelSpec,
102 #[error("providers.{provider}.{field} cannot be used with multiple {provider} models in one alloy spec")]
106 DuplicateProvider { provider: String, field: String },
107 #[error("Invalid model spec: {0}")]
109 InvalidModelSpec(String),
110 #[error("Failed to build provider request: {0}")]
113 ProviderRequest(String),
114 #[error("Invalid argument: {0}")]
116 InvalidArgument(String),
117}
118
119impl LlmError {
120 pub fn is_retryable(&self) -> bool {
124 matches!(
125 self,
126 LlmError::RateLimited(_)
127 | LlmError::ServerError { .. }
128 | LlmError::Timeout(_)
129 | LlmError::Network(_)
130 | LlmError::StreamInterrupted(_)
131 )
132 }
133}
134
135impl From<reqwest::Error> for LlmError {
136 fn from(error: reqwest::Error) -> Self {
137 if error.is_timeout() {
138 return LlmError::Timeout(error.to_string());
139 }
140 if error.is_connect() || error.is_request() {
141 return LlmError::Network(error.to_string());
142 }
143 match error.status().map(|s| s.as_u16()) {
144 Some(429) => LlmError::RateLimited(error.to_string()),
145 Some(s) if (500..600).contains(&s) => LlmError::ServerError { status: Some(s), message: error.to_string() },
146 _ => LlmError::ApiRequest(error.to_string()),
147 }
148 }
149}
150
151impl From<serde_json::Error> for LlmError {
152 fn from(error: serde_json::Error) -> Self {
153 LlmError::JsonParsing(error.to_string())
154 }
155}
156
157impl From<std::io::Error> for LlmError {
158 fn from(error: std::io::Error) -> Self {
159 LlmError::IoError(error.to_string())
160 }
161}
162
163impl From<reqwest::header::InvalidHeaderValue> for LlmError {
164 fn from(error: reqwest::header::InvalidHeaderValue) -> Self {
165 LlmError::InvalidApiKey(error.to_string())
166 }
167}
168
169impl From<async_openai::error::OpenAIError> for LlmError {
170 fn from(error: async_openai::error::OpenAIError) -> Self {
171 use async_openai::error::OpenAIError;
172 match error {
173 OpenAIError::Reqwest(e) => LlmError::from(e),
174 OpenAIError::StreamError(e) => LlmError::StreamInterrupted(e.to_string()),
175 OpenAIError::ApiError(api_err) => LlmError::ApiError(api_err.to_string()),
176 OpenAIError::JSONDeserialize(e, _) => LlmError::JsonParsing(e.to_string()),
177 OpenAIError::FileSaveError(s) | OpenAIError::FileReadError(s) => LlmError::IoError(s),
178 OpenAIError::InvalidArgument(s) => LlmError::InvalidArgument(s),
179 }
180 }
181}
182
183#[cfg(feature = "codex")]
184impl From<aether_auth::OAuthError> for LlmError {
185 fn from(error: aether_auth::OAuthError) -> Self {
186 LlmError::OAuthError(error.to_string())
187 }
188}
189
190pub type Result<T> = std::result::Result<T, LlmError>;
191
192#[cfg(test)]
193mod tests {
194 use super::*;
195
196 #[test]
197 fn is_retryable() {
198 assert!(LlmError::RateLimited("rl".into()).is_retryable());
199 assert!(LlmError::ServerError { status: Some(503), message: "x".into() }.is_retryable());
200 assert!(LlmError::ServerError { status: None, message: "stream-level".into() }.is_retryable());
201 assert!(LlmError::Timeout("t".into()).is_retryable());
202 assert!(LlmError::Network("n".into()).is_retryable());
203 assert!(LlmError::StreamInterrupted("s".into()).is_retryable());
204
205 assert!(!LlmError::ApiError("x".into()).is_retryable());
206 assert!(!LlmError::ApiRequest("x".into()).is_retryable());
207 assert!(!LlmError::MissingApiKey("x".into()).is_retryable());
208 assert!(!LlmError::InvalidApiKey("x".into()).is_retryable());
209 assert!(!LlmError::HttpClientCreation("x".into()).is_retryable());
210 assert!(!LlmError::IoError("x".into()).is_retryable());
211 assert!(!LlmError::JsonParsing("x".into()).is_retryable());
212 assert!(!LlmError::ToolParameterParsing { tool_name: "t".into(), error: "e".into() }.is_retryable());
213 assert!(!LlmError::OAuthError("x".into()).is_retryable());
214 assert!(!LlmError::UnsupportedContent("x".into()).is_retryable());
215 assert!(!LlmError::MissingProviderUrl { provider: "azure-foundry".into() }.is_retryable());
216 assert!(!LlmError::UnknownProvider { provider: "foo".into() }.is_retryable());
217 assert!(!LlmError::EmptyModelSpec.is_retryable());
218 assert!(
219 !LlmError::DuplicateProvider { provider: "bedrock".into(), field: "inferenceProfileArn".into() }
220 .is_retryable()
221 );
222 assert!(!LlmError::InvalidModelSpec("x".into()).is_retryable());
223 assert!(!LlmError::ProviderRequest("x".into()).is_retryable());
224 assert!(!LlmError::InvalidArgument("x".into()).is_retryable());
225 assert!(!LlmError::ContextOverflow(ContextOverflowError::new("p", None, None, None, "m")).is_retryable());
226 }
227}