Skip to main content

cera_client/
error.rs

1//! Error types and API error payload definitions for `cera-client`.
2
3use serde::{Deserialize, Deserializer, Serialize};
4
5/// Detailed error payload returned by OpenAI-compatible endpoints.
6#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
7pub struct ApiErrorPayload {
8    /// Human-readable description of the error.
9    #[serde(default, deserialize_with = "deserialize_null_as_default")]
10    pub message: String,
11    /// Error category or classification (e.g. `invalid_request_error`).
12    #[serde(rename = "type", default)]
13    pub error_type: Option<String>,
14    /// Specific parameter that caused the error, if applicable.
15    #[serde(default)]
16    pub param: Option<String>,
17    /// Machine-readable error code (e.g. `rate_limit_exceeded` or OpenRouter numeric codes).
18    #[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/// JSON envelope containing the API error payload.
68#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
69pub struct ApiErrorEnvelope {
70    /// Inner error payload.
71    #[serde(deserialize_with = "deserialize_api_error_payload")]
72    pub error: ApiErrorPayload,
73}
74
75/// Errors that can occur when calling OpenAI and OpenRouter endpoints.
76#[derive(Debug, thiserror::Error)]
77pub enum ClientError {
78    /// Network or HTTP transport failure.
79    #[error("HTTP transport error: {0}")]
80    Http(#[from] reqwest::Error),
81
82    /// The remote endpoint returned an error status code and payload.
83    #[error("API error{}: {message}", .status.map(|s| format!(" (status {s})")).unwrap_or_default())]
84    Api {
85        /// HTTP response status code, if available.
86        status: Option<reqwest::StatusCode>,
87        /// Human-readable error message.
88        message: String,
89        /// Error classification string.
90        error_type: Option<String>,
91        /// Machine-readable error code.
92        code: Option<String>,
93        /// Parameter that caused the error.
94        param: Option<String>,
95    },
96
97    /// JSON serialization or deserialization failure.
98    #[error(
99        "JSON serialization error: {source}{}",
100        .raw_payload.as_deref().map(|p| format!(" (raw payload: {p})")).unwrap_or_default()
101    )]
102    Serialization {
103        /// Underlying serde_json error.
104        #[source]
105        source: serde_json::Error,
106        /// Raw unparsed payload if available for diagnostics.
107        raw_payload: Option<String>,
108    },
109
110    /// Server-Sent Events streaming error or protocol violation.
111    #[error("Stream error: {0}")]
112    Stream(String),
113
114    /// An API key was required but not provided or found in the environment.
115    #[error("Missing API key: {0}")]
116    MissingApiKey(String),
117
118    /// The provided base URL or endpoint could not be parsed.
119    #[error("Invalid URL: {0}")]
120    InvalidUrl(String),
121
122    /// An HTTP header value contains invalid characters.
123    #[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    /// Returns the HTTP status code if this error originated from an API response.
138    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    /// Returns true if the error represents an HTTP 429 Too Many Requests response or rate limit code.
147    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    /// Returns true if the error is a transient server failure (HTTP 5xx, 408, or connection error).
160    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}