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
135pub 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}