Skip to main content

gateway_core/
error.rs

1use serde::{Deserialize, Serialize};
2
3#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
4pub struct DependencyFailure {
5    pub provider: String,
6    pub status: Option<u16>,
7    pub message: String,
8}
9
10#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
11pub enum ProviderError {
12    #[error("invalid provider request: {0}")]
13    InvalidRequest(String),
14    #[error("context window exceeded: {0}")]
15    ContextWindowExceeded(String),
16    #[error("surface unsupported by provider: {0}")]
17    Unsupported(String),
18    #[error("upstream model unavailable")]
19    ModelUnavailable(Vec<DependencyFailure>),
20    #[error("provider dependency failed")]
21    Dependency(Vec<DependencyFailure>),
22    #[error("provider stream was invalid: {0}")]
23    InvalidStream(String),
24    #[error("provider stream was rate limited: {0}")]
25    RateLimitedStream(String),
26    #[error("all provider circuits are open")]
27    AllCircuitsOpen(Vec<String>),
28}
29
30impl ProviderError {
31    pub fn code(&self) -> &'static str {
32        match self {
33            Self::InvalidRequest(_) => "invalid_request",
34            Self::ContextWindowExceeded(_) => "context_window_exceeded",
35            Self::Unsupported(_) => "unsupported",
36            Self::ModelUnavailable(_) => "model_unavailable",
37            Self::Dependency(_) => "provider_dependency_failed",
38            Self::InvalidStream(_) => "invalid_stream",
39            Self::RateLimitedStream(_) => "provider_rate_limited",
40            Self::AllCircuitsOpen(_) => "all_provider_circuits_open",
41        }
42    }
43
44    pub fn from_upstream(provider: impl Into<String>, status: u16, body: &str) -> Self {
45        let provider = provider.into();
46        let message = extract_message(body);
47        if is_context_length_error(body) || is_context_length_error(&message) {
48            return Self::ContextWindowExceeded(message);
49        }
50        let failure = DependencyFailure {
51            provider,
52            status: Some(status),
53            message,
54        };
55        if status == 404 {
56            Self::ModelUnavailable(vec![failure])
57        } else if (400..500).contains(&status) && status != 429 {
58            Self::InvalidRequest(failure.message)
59        } else {
60            Self::Dependency(vec![failure])
61        }
62    }
63
64    pub fn transport(provider: impl Into<String>, message: impl Into<String>) -> Self {
65        Self::Dependency(vec![DependencyFailure {
66            provider: provider.into(),
67            status: None,
68            message: message.into(),
69        }])
70    }
71
72    pub fn is_retryable(&self) -> bool {
73        match self {
74            Self::Dependency(failures) => {
75                !failures.is_empty()
76                    && failures.iter().all(|failure| {
77                        failure
78                            .status
79                            .is_none_or(|status| status == 429 || status >= 500)
80                    })
81            }
82            Self::ModelUnavailable(_) => true,
83            _ => false,
84        }
85    }
86
87    pub fn affects_provider_health(&self) -> bool {
88        matches!(self, Self::Dependency(_)) && self.is_retryable()
89    }
90
91    pub fn is_stream_rate_limited(&self) -> bool {
92        matches!(self, Self::RateLimitedStream(_))
93    }
94
95    pub fn is_credential_rate_limited(&self) -> bool {
96        match self {
97            Self::RateLimitedStream(_) => true,
98            Self::Dependency(failures) => {
99                failures.iter().any(|failure| failure.status == Some(429))
100            }
101            _ => false,
102        }
103    }
104}
105
106fn extract_message(body: &str) -> String {
107    serde_json::from_str::<serde_json::Value>(body)
108        .ok()
109        .and_then(|value| {
110            value
111                .pointer("/error/message")
112                .or_else(|| value.get("message"))
113                .and_then(serde_json::Value::as_str)
114                .map(str::to_owned)
115        })
116        .unwrap_or_else(|| body.to_owned())
117}
118
119fn is_context_length_error(text: &str) -> bool {
120    let text = text.to_ascii_lowercase();
121    [
122        "context_length_exceeded",
123        "context length",
124        "context window",
125        "prompt is too long",
126        "prompt too long",
127        "maximum number of tokens",
128        "too many tokens",
129        "maximum prompt length",
130    ]
131    .iter()
132    .any(|signal| text.contains(signal))
133}
134
135/// Recognize only explicit provider rate-limit markers in an SSE JSON payload.
136pub fn is_rate_limit_payload(value: &serde_json::Value) -> bool {
137    let error = value.get("error");
138    let error_shaped =
139        error.is_some() || value.get("type").and_then(serde_json::Value::as_str) == Some("error");
140    if !error_shaped {
141        return false;
142    }
143    let status_is_429 = [value.get("status"), value.pointer("/error/status")]
144        .into_iter()
145        .flatten()
146        .any(|status| {
147            status.as_u64() == Some(429) || status.as_str().is_some_and(|status| status == "429")
148        });
149    if status_is_429 {
150        return true;
151    }
152    [
153        value
154            .pointer("/error/type")
155            .and_then(serde_json::Value::as_str),
156        value.pointer("/error/code").and_then(|code| {
157            code.as_str()
158                .or_else(|| (code.as_u64() == Some(429)).then_some("429"))
159        }),
160        value.pointer("/type").and_then(serde_json::Value::as_str),
161        value.pointer("/code").and_then(serde_json::Value::as_str),
162    ]
163    .into_iter()
164    .flatten()
165    .any(|signal| signal.contains("rate_limit") || signal == "429")
166}
167
168#[cfg(test)]
169mod tests {
170    use super::*;
171
172    #[test]
173    fn normalizes_context_limit_signals_without_retrying_or_degrading_health() {
174        for body in [
175            r#"{"error":{"code":"context_length_exceeded","message":"too long"}}"#,
176            r#"{"error":{"message":"prompt is too long: 250000 tokens"}}"#,
177            r#"{"message":"input exceeds the maximum number of tokens"}"#,
178        ] {
179            let error = ProviderError::from_upstream("provider", 400, body);
180            assert!(matches!(error, ProviderError::ContextWindowExceeded(_)));
181            assert!(!error.is_retryable());
182            assert!(!error.affects_provider_health());
183        }
184    }
185
186    #[test]
187    fn rate_limits_and_server_failures_retry_and_affect_health() {
188        for status in [429, 500, 502, 503, 599] {
189            let error = ProviderError::from_upstream("provider", status, "upstream unavailable");
190            assert!(matches!(error, ProviderError::Dependency(_)));
191            assert!(error.is_retryable(), "status {status}");
192            assert!(error.affects_provider_health(), "status {status}");
193        }
194        let transport = ProviderError::transport("provider", "timeout");
195        assert!(transport.is_retryable());
196        assert!(transport.affects_provider_health());
197    }
198
199    #[test]
200    fn authentication_and_other_client_failures_are_permanent_but_not_unhealthy() {
201        for status in [400, 401, 403, 422] {
202            let error = ProviderError::from_upstream("provider", status, "invalid request");
203            assert!(matches!(error, ProviderError::InvalidRequest(_)));
204            assert!(!error.is_retryable(), "status {status}");
205            assert!(!error.affects_provider_health(), "status {status}");
206        }
207    }
208
209    #[test]
210    fn missing_model_fails_over_without_marking_provider_unhealthy() {
211        let error = ProviderError::from_upstream("foundry", 404, "missing deployment");
212        assert!(matches!(error, ProviderError::ModelUnavailable(_)));
213        assert!(error.is_retryable());
214        assert!(!error.affects_provider_health());
215    }
216
217    #[test]
218    fn recognizes_only_explicit_stream_rate_limit_shapes() {
219        for body in [
220            r#"{"type":"error","error":{"type":"rate_limit_error"}}"#,
221            r#"{"error":{"code":"rate_limit_exceeded"}}"#,
222            r#"{"error":{"code":429}}"#,
223            r#"{"error":{"status":429}}"#,
224        ] {
225            let value: serde_json::Value = serde_json::from_str(body).unwrap();
226            assert!(is_rate_limit_payload(&value), "{body}");
227        }
228        for body in [
229            r#"{"error":{"type":"overloaded_error"}}"#,
230            r#"{"error":{"message":"try again later"}}"#,
231            r#"{"status":500}"#,
232            r#"{"type":"rate_limits.updated","rate_limits":{"requests":10}}"#,
233        ] {
234            let value: serde_json::Value = serde_json::from_str(body).unwrap();
235            assert!(!is_rate_limit_payload(&value), "{body}");
236        }
237    }
238}