1pub mod anthropic;
8pub mod openai;
9pub mod preflight;
10pub mod retry;
11pub(crate) mod sse;
12
13use crate::message::{CompletionRequest, CompletionResponse, Usage};
14use anyhow::Result;
15use async_trait::async_trait;
16use tokio::sync::mpsc::UnboundedSender;
17
18#[derive(Debug, Clone)]
20pub enum StreamEvent {
21 TextDelta(String),
22 ThinkingDelta(String),
23 ToolUseStart {
25 name: String,
26 },
27 Usage(Usage),
36}
37
38pub type StreamSink = UnboundedSender<StreamEvent>;
39
40#[async_trait]
41pub trait Provider: Send + Sync {
42 fn id(&self) -> &str;
44
45 fn default_model(&self) -> &str;
47
48 fn vision(&self) -> bool {
61 false
62 }
63
64 async fn complete(
67 &self,
68 req: &CompletionRequest,
69 sink: Option<&StreamSink>,
70 ) -> Result<CompletionResponse>;
71}
72
73pub fn build(cfg: &crate::config::ProviderConfig) -> Result<Box<dyn Provider>> {
75 match cfg.kind.as_str() {
76 "anthropic" => Ok(Box::new(anthropic::Anthropic::from_config(cfg)?)),
77 "openai" | "openai-compatible" | "local" => {
78 Ok(Box::new(openai::OpenAiCompatible::from_config(cfg)?))
79 }
80 other => {
81 anyhow::bail!("unknown provider kind {other:?} (expected: anthropic, openai, local)")
82 }
83 }
84}
85
86pub struct Failover {
101 primary: Box<dyn Provider>,
102 fallbacks: Vec<(String, Box<dyn Provider>)>,
103}
104
105impl Failover {
106 pub fn new(primary: Box<dyn Provider>, fallbacks: Vec<(String, Box<dyn Provider>)>) -> Self {
107 Failover { primary, fallbacks }
108 }
109}
110
111fn failover_worthy(e: &anyhow::Error) -> bool {
112 e.downcast_ref::<retry::ProviderError>()
113 .is_some_and(retry::ProviderError::transient)
114}
115
116#[async_trait]
117impl Provider for Failover {
118 fn id(&self) -> &str {
119 self.primary.id()
120 }
121
122 fn default_model(&self) -> &str {
123 self.primary.default_model()
124 }
125
126 async fn complete(
127 &self,
128 req: &CompletionRequest,
129 sink: Option<&StreamSink>,
130 ) -> Result<CompletionResponse> {
131 let mut last = match self.primary.complete(req, sink).await {
132 Ok(response) => return Ok(response),
133 Err(e) if failover_worthy(&e) => e,
134 Err(e) => return Err(e),
135 };
136
137 for (name, provider) in &self.fallbacks {
138 tracing::warn!(
139 error = %last,
140 fallback = %name,
141 "provider failed transiently after retries; falling back"
142 );
143 let fb_req = CompletionRequest {
144 model: provider.default_model().to_string(),
145 ..req.clone()
146 };
147 match provider.complete(&fb_req, sink).await {
148 Ok(response) => return Ok(response),
149 Err(e) if failover_worthy(&e) => last = e,
150 Err(e) => return Err(e),
151 }
152 }
153 Err(last.context(format!(
154 "the primary and {} fallback(s) all failed transiently",
155 self.fallbacks.len()
156 )))
157 }
158}
159
160#[cfg(test)]
161mod failover_tests {
162 use super::*;
163 use crate::message::{Block, Message, StopReason};
164 use retry::ProviderError;
165 use std::sync::atomic::{AtomicUsize, Ordering};
166 use std::sync::{Arc, Mutex};
167
168 struct Failing {
170 error: fn() -> anyhow::Error,
171 calls: Arc<AtomicUsize>,
172 }
173
174 #[async_trait]
175 impl Provider for Failing {
176 fn id(&self) -> &str {
177 "failing"
178 }
179 fn default_model(&self) -> &str {
180 "primary-model"
181 }
182 async fn complete(
183 &self,
184 _req: &CompletionRequest,
185 _sink: Option<&StreamSink>,
186 ) -> Result<CompletionResponse> {
187 self.calls.fetch_add(1, Ordering::SeqCst);
188 Err((self.error)())
189 }
190 }
191
192 struct Recording {
194 model_seen: Arc<Mutex<Option<String>>>,
195 calls: Arc<AtomicUsize>,
196 }
197
198 #[async_trait]
199 impl Provider for Recording {
200 fn id(&self) -> &str {
201 "recording"
202 }
203 fn default_model(&self) -> &str {
204 "fallback-model"
205 }
206 async fn complete(
207 &self,
208 req: &CompletionRequest,
209 _sink: Option<&StreamSink>,
210 ) -> Result<CompletionResponse> {
211 self.calls.fetch_add(1, Ordering::SeqCst);
212 *self.model_seen.lock().unwrap() = Some(req.model.clone());
213 Ok(CompletionResponse {
214 message: Message::assistant(vec![Block::text("from the fallback")]),
215 stop_reason: StopReason::EndTurn,
216 usage: Usage::default(),
217 refusal: None,
218 model: "fallback-model".into(),
219 malformed_tool_args: 0,
220 })
221 }
222 }
223
224 fn req() -> CompletionRequest {
225 CompletionRequest {
226 model: "primary-model".into(),
227 system: None,
228 messages: vec![Message::user("hi")],
229 tools: Vec::new(),
230 max_tokens: 64,
231 effort: None,
232 thinking: false,
233 cache_prompt: false,
234 }
235 }
236
237 type Rig = (
238 Failover,
239 Arc<AtomicUsize>,
240 Arc<Mutex<Option<String>>>,
241 Arc<AtomicUsize>,
242 );
243
244 fn rig(error: fn() -> anyhow::Error) -> Rig {
245 let primary_calls = Arc::new(AtomicUsize::new(0));
246 let fallback_calls = Arc::new(AtomicUsize::new(0));
247 let model_seen = Arc::new(Mutex::new(None));
248 let failover = Failover::new(
249 Box::new(Failing {
250 error,
251 calls: Arc::clone(&primary_calls),
252 }),
253 vec![(
254 "small".into(),
255 Box::new(Recording {
256 model_seen: Arc::clone(&model_seen),
257 calls: Arc::clone(&fallback_calls),
258 }) as Box<dyn Provider>,
259 )],
260 );
261 (failover, primary_calls, model_seen, fallback_calls)
262 }
263
264 #[tokio::test]
265 async fn a_transient_exhaustion_falls_back_and_the_fallback_answers_as_itself() {
266 let (failover, _, model_seen, _) =
267 rig(|| anyhow::Error::new(ProviderError::Overloaded).context("anthropic 529: busy"));
268
269 let response = failover.complete(&req(), None).await.unwrap();
270
271 assert_eq!(response.message.text(), "from the fallback");
272 assert_eq!(
275 model_seen.lock().unwrap().as_deref(),
276 Some("fallback-model")
277 );
278 }
279
280 #[tokio::test]
281 async fn terminal_classes_never_fall_back() {
282 for error in [
286 (|| anyhow::Error::new(ProviderError::Invalid("bad".into()))) as fn() -> anyhow::Error,
287 || anyhow::Error::new(ProviderError::Auth),
288 || anyhow::Error::new(ProviderError::ContextOverflow),
289 ] {
290 let (failover, primary_calls, _, fallback_calls) = rig(error);
291 let err = failover.complete(&req(), None).await.unwrap_err();
292 assert!(err.downcast_ref::<ProviderError>().is_some());
293 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
294 assert_eq!(
295 fallback_calls.load(Ordering::SeqCst),
296 0,
297 "the fallback was consulted"
298 );
299 }
300 }
301
302 #[tokio::test]
303 async fn an_unclassified_error_never_falls_back_because_it_may_be_mid_stream() {
304 let (failover, _, _, fallback_calls) = rig(|| anyhow::anyhow!("stream aborted mid-body"));
308
309 let err = failover.complete(&req(), None).await.unwrap_err();
310 assert!(err.to_string().contains("stream aborted"));
311 assert_eq!(fallback_calls.load(Ordering::SeqCst), 0);
312 }
313}