1use serde::{Deserialize, Deserializer, Serialize};
4
5#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
7pub struct ApiErrorPayload {
8 #[serde(default, deserialize_with = "deserialize_null_as_default")]
10 pub message: String,
11 #[serde(rename = "type", default)]
13 pub error_type: Option<String>,
14 #[serde(default)]
16 pub param: Option<String>,
17 #[serde(default, deserialize_with = "deserialize_string_or_int")]
19 pub code: Option<String>,
20}
21
22fn deserialize_null_as_default<'de, D>(deserializer: D) -> Result<String, D::Error>
23where
24 D: Deserializer<'de>,
25{
26 let opt = Option::<String>::deserialize(deserializer)?;
27 Ok(opt.unwrap_or_default())
28}
29
30fn deserialize_string_or_int<'de, D>(deserializer: D) -> Result<Option<String>, D::Error>
31where
32 D: Deserializer<'de>,
33{
34 let opt = Option::<serde_json::Value>::deserialize(deserializer)?;
35 Ok(opt.and_then(|val| match val {
36 serde_json::Value::String(s) => Some(s),
37 serde_json::Value::Number(n) => Some(n.to_string()),
38 serde_json::Value::Bool(b) => Some(b.to_string()),
39 serde_json::Value::Null => None,
40 other => Some(other.to_string()),
41 }))
42}
43
44fn deserialize_api_error_payload<'de, D>(deserializer: D) -> Result<ApiErrorPayload, D::Error>
45where
46 D: Deserializer<'de>,
47{
48 #[derive(Deserialize)]
49 #[serde(untagged)]
50 enum ErrorPayloadOrString {
51 Payload(ApiErrorPayload),
52 String(String),
53 }
54
55 let val = ErrorPayloadOrString::deserialize(deserializer)?;
56 Ok(match val {
57 ErrorPayloadOrString::Payload(p) => p,
58 ErrorPayloadOrString::String(s) => ApiErrorPayload {
59 message: s,
60 error_type: None,
61 param: None,
62 code: None,
63 },
64 })
65}
66
67#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
69pub struct ApiErrorEnvelope {
70 #[serde(deserialize_with = "deserialize_api_error_payload")]
72 pub error: ApiErrorPayload,
73}
74
75#[derive(Debug, thiserror::Error)]
77pub enum ClientError {
78 #[error("HTTP transport error: {0}")]
80 Http(#[from] reqwest::Error),
81
82 #[error("API error{}: {message}", .status.map(|s| format!(" (status {s})")).unwrap_or_default())]
84 Api {
85 status: Option<reqwest::StatusCode>,
87 message: String,
89 error_type: Option<String>,
91 code: Option<String>,
93 param: Option<String>,
95 },
96
97 #[error(
99 "JSON serialization error: {source}{}",
100 .raw_payload.as_deref().map(|p| format!(" (raw payload: {p})")).unwrap_or_default()
101 )]
102 Serialization {
103 #[source]
105 source: serde_json::Error,
106 raw_payload: Option<String>,
108 },
109
110 #[error("Stream error: {0}")]
112 Stream(String),
113
114 #[error("Missing API key: {0}")]
116 MissingApiKey(String),
117
118 #[error("Invalid URL: {0}")]
120 InvalidUrl(String),
121
122 #[error("Invalid HTTP header: {0}")]
124 InvalidHeader(String),
125}
126
127impl From<serde_json::Error> for ClientError {
128 fn from(source: serde_json::Error) -> Self {
129 Self::Serialization {
130 source,
131 raw_payload: None,
132 }
133 }
134}
135
136impl ClientError {
137 pub fn status(&self) -> Option<reqwest::StatusCode> {
139 match self {
140 Self::Api { status, .. } => *status,
141 Self::Http(err) => err.status(),
142 _ => None,
143 }
144 }
145
146 pub fn is_rate_limited(&self) -> bool {
148 if self.status() == Some(reqwest::StatusCode::TOO_MANY_REQUESTS) {
149 return true;
150 }
151 if let Self::Api { code, .. } = self
152 && let Some(c) = code.as_deref()
153 {
154 return c == "429" || c == "rate_limit_exceeded" || c == "rate_limit";
155 }
156 false
157 }
158
159 pub fn is_transient(&self) -> bool {
161 if let Some(status) = self.status()
162 && (status.is_server_error() || status == reqwest::StatusCode::REQUEST_TIMEOUT)
163 {
164 return true;
165 }
166 if let Self::Api {
167 code, error_type, ..
168 } = self
169 {
170 if let Some(c) = code.as_deref()
171 && matches!(
172 c,
173 "500"
174 | "502"
175 | "503"
176 | "504"
177 | "server_error"
178 | "overloaded"
179 | "timeout"
180 | "service_unavailable"
181 | "gateway_timeout"
182 | "engine_overloaded"
183 )
184 {
185 return true;
186 }
187 if let Some(t) = error_type.as_deref()
188 && (t == "server_error" || t == "timeout" || t == "service_unavailable")
189 {
190 return true;
191 }
192 }
193 if let Self::Http(err) = self {
194 #[cfg(not(target_arch = "wasm32"))]
195 return err.is_timeout() || err.is_connect();
196 #[cfg(target_arch = "wasm32")]
197 return err.is_timeout();
198 }
199 false
200 }
201}
202
203#[cfg(test)]
204mod tests {
205 use super::*;
206
207 #[test]
208 fn test_parse_string_error_code() {
209 let json = r#"{"error": {"message": "Rate limit exceeded", "type": "tokens", "code": "rate_limit_exceeded"}}"#;
210 let env: ApiErrorEnvelope = serde_json::from_str(json).unwrap();
211 assert_eq!(env.error.code.as_deref(), Some("rate_limit_exceeded"));
212 assert_eq!(env.error.message, "Rate limit exceeded");
213 }
214
215 #[test]
216 fn test_parse_numeric_error_code() {
217 let json = r#"{"error": {"message": "Payment required", "code": 402}}"#;
218 let env: ApiErrorEnvelope = serde_json::from_str(json).unwrap();
219 assert_eq!(env.error.code.as_deref(), Some("402"));
220 assert_eq!(env.error.message, "Payment required");
221 }
222
223 #[test]
224 fn test_error_status_and_helpers() {
225 let err429 = ClientError::Api {
226 status: Some(reqwest::StatusCode::TOO_MANY_REQUESTS),
227 message: "Rate limit reached".to_string(),
228 error_type: None,
229 code: None,
230 param: None,
231 };
232 assert_eq!(
233 err429.status(),
234 Some(reqwest::StatusCode::TOO_MANY_REQUESTS)
235 );
236 assert!(err429.is_rate_limited());
237 assert!(!err429.is_transient());
238
239 let err_code_rate_limit = ClientError::Api {
240 status: None,
241 message: "Too fast".to_string(),
242 error_type: None,
243 code: Some("rate_limit_exceeded".to_string()),
244 param: None,
245 };
246 assert_eq!(err_code_rate_limit.status(), None);
247 assert!(err_code_rate_limit.is_rate_limited());
248 assert!(!err_code_rate_limit.is_transient());
249
250 let err500 = ClientError::Api {
251 status: Some(reqwest::StatusCode::INTERNAL_SERVER_ERROR),
252 message: "Internal server error".to_string(),
253 error_type: None,
254 code: None,
255 param: None,
256 };
257 assert_eq!(
258 err500.status(),
259 Some(reqwest::StatusCode::INTERNAL_SERVER_ERROR)
260 );
261 assert!(!err500.is_rate_limited());
262 assert!(err500.is_transient());
263
264 let err_in_band_503 = ClientError::Api {
265 status: None,
266 message: "Model overloaded".to_string(),
267 error_type: Some("server_error".to_string()),
268 code: Some("503".to_string()),
269 param: None,
270 };
271 assert!(err_in_band_503.is_transient());
272
273 let err408 = ClientError::Api {
274 status: Some(reqwest::StatusCode::REQUEST_TIMEOUT),
275 message: "Timeout".to_string(),
276 error_type: None,
277 code: None,
278 param: None,
279 };
280 assert!(err408.is_transient());
281
282 let err_service_unavailable = ClientError::Api {
283 status: None,
284 message: "Service unavailable".to_string(),
285 error_type: Some("service_unavailable".to_string()),
286 code: Some("service_unavailable".to_string()),
287 param: None,
288 };
289 assert!(err_service_unavailable.is_transient());
290
291 let err_gateway_timeout = ClientError::Api {
292 status: None,
293 message: "Gateway timeout".to_string(),
294 error_type: None,
295 code: Some("gateway_timeout".to_string()),
296 param: None,
297 };
298 assert!(err_gateway_timeout.is_transient());
299 }
300
301 #[test]
302 fn test_parse_null_or_missing_error_message() {
303 let json_null = r#"{"error": {"message": null, "code": 500}}"#;
304 let env_null: ApiErrorEnvelope = serde_json::from_str(json_null).unwrap();
305 assert_eq!(env_null.error.message, "");
306 assert_eq!(env_null.error.code.as_deref(), Some("500"));
307
308 let json_missing = r#"{"error": {"code": 503}}"#;
309 let env_missing: ApiErrorEnvelope = serde_json::from_str(json_missing).unwrap();
310 assert_eq!(env_missing.error.message, "");
311 assert_eq!(env_missing.error.code.as_deref(), Some("503"));
312 }
313
314 #[test]
315 fn test_parse_bare_string_error_message() {
316 let json_str = r#"{"error": "model is currently overloaded, please try again"}"#;
317 let env: ApiErrorEnvelope = serde_json::from_str(json_str).unwrap();
318 assert_eq!(
319 env.error.message,
320 "model is currently overloaded, please try again"
321 );
322 assert_eq!(env.error.code, None);
323 assert_eq!(env.error.error_type, None);
324 }
325
326 #[test]
327 fn test_parse_various_error_codes() {
328 let json_float = r#"{"error": {"message": "Rate limit", "code": 429.5}}"#;
329 let env_float: ApiErrorEnvelope = serde_json::from_str(json_float).unwrap();
330 assert_eq!(env_float.error.code.as_deref(), Some("429.5"));
331
332 let json_bool = r#"{"error": {"message": "Invalid", "code": true}}"#;
333 let env_bool: ApiErrorEnvelope = serde_json::from_str(json_bool).unwrap();
334 assert_eq!(env_bool.error.code.as_deref(), Some("true"));
335
336 let json_null_code = r#"{"error": {"message": "Error", "code": null}}"#;
337 let env_null_code: ApiErrorEnvelope = serde_json::from_str(json_null_code).unwrap();
338 assert_eq!(env_null_code.error.code, None);
339 }
340
341 #[test]
342 fn test_serialization_error_display_with_and_without_raw_payload() {
343 let err_no_payload: ClientError =
344 serde_json::from_str::<serde_json::Value>("invalid json{")
345 .unwrap_err()
346 .into();
347 let display_no_payload = err_no_payload.to_string();
348 assert!(display_no_payload.starts_with("JSON serialization error:"));
349 assert!(!display_no_payload.contains("raw payload"));
350
351 let err_with_payload = ClientError::Serialization {
352 source: serde_json::from_str::<serde_json::Value>("invalid json{").unwrap_err(),
353 raw_payload: Some("data: {malformed}".to_string()),
354 };
355 let display_with_payload = err_with_payload.to_string();
356 assert!(display_with_payload.contains("raw payload: data: {malformed}"));
357 }
358}