pub mod anthropic;
pub mod openai;
pub mod preflight;
pub mod retry;
pub(crate) mod sse;
use crate::message::{CompletionRequest, CompletionResponse, Usage};
use anyhow::Result;
use async_trait::async_trait;
use tokio::sync::mpsc::UnboundedSender;
#[derive(Debug, Clone)]
pub enum StreamEvent {
TextDelta(String),
ThinkingDelta(String),
ToolUseStart {
name: String,
},
Usage(Usage),
}
pub type StreamSink = UnboundedSender<StreamEvent>;
#[async_trait]
pub trait Provider: Send + Sync {
fn id(&self) -> &str;
fn default_model(&self) -> &str;
fn vision(&self) -> bool {
false
}
async fn complete(
&self,
req: &CompletionRequest,
sink: Option<&StreamSink>,
) -> Result<CompletionResponse>;
}
pub fn build(cfg: &crate::config::ProviderConfig) -> Result<Box<dyn Provider>> {
match cfg.kind.as_str() {
"anthropic" => Ok(Box::new(anthropic::Anthropic::from_config(cfg)?)),
"openai" | "openai-compatible" | "local" => {
Ok(Box::new(openai::OpenAiCompatible::from_config(cfg)?))
}
other => {
anyhow::bail!("unknown provider kind {other:?} (expected: anthropic, openai, local)")
}
}
}
pub struct Failover {
primary: Box<dyn Provider>,
fallbacks: Vec<(String, Box<dyn Provider>)>,
}
impl Failover {
pub fn new(primary: Box<dyn Provider>, fallbacks: Vec<(String, Box<dyn Provider>)>) -> Self {
Failover { primary, fallbacks }
}
}
fn failover_worthy(e: &anyhow::Error) -> bool {
e.downcast_ref::<retry::ProviderError>()
.is_some_and(retry::ProviderError::transient)
}
#[async_trait]
impl Provider for Failover {
fn id(&self) -> &str {
self.primary.id()
}
fn default_model(&self) -> &str {
self.primary.default_model()
}
async fn complete(
&self,
req: &CompletionRequest,
sink: Option<&StreamSink>,
) -> Result<CompletionResponse> {
let mut last = match self.primary.complete(req, sink).await {
Ok(response) => return Ok(response),
Err(e) if failover_worthy(&e) => e,
Err(e) => return Err(e),
};
for (name, provider) in &self.fallbacks {
tracing::warn!(
error = %last,
fallback = %name,
"provider failed transiently after retries; falling back"
);
let fb_req = CompletionRequest {
model: provider.default_model().to_string(),
..req.clone()
};
match provider.complete(&fb_req, sink).await {
Ok(response) => return Ok(response),
Err(e) if failover_worthy(&e) => last = e,
Err(e) => return Err(e),
}
}
Err(last.context(format!(
"the primary and {} fallback(s) all failed transiently",
self.fallbacks.len()
)))
}
}
#[cfg(test)]
mod failover_tests {
use super::*;
use crate::message::{Block, Message, StopReason};
use retry::ProviderError;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
struct Failing {
error: fn() -> anyhow::Error,
calls: Arc<AtomicUsize>,
}
#[async_trait]
impl Provider for Failing {
fn id(&self) -> &str {
"failing"
}
fn default_model(&self) -> &str {
"primary-model"
}
async fn complete(
&self,
_req: &CompletionRequest,
_sink: Option<&StreamSink>,
) -> Result<CompletionResponse> {
self.calls.fetch_add(1, Ordering::SeqCst);
Err((self.error)())
}
}
struct Recording {
model_seen: Arc<Mutex<Option<String>>>,
calls: Arc<AtomicUsize>,
}
#[async_trait]
impl Provider for Recording {
fn id(&self) -> &str {
"recording"
}
fn default_model(&self) -> &str {
"fallback-model"
}
async fn complete(
&self,
req: &CompletionRequest,
_sink: Option<&StreamSink>,
) -> Result<CompletionResponse> {
self.calls.fetch_add(1, Ordering::SeqCst);
*self.model_seen.lock().unwrap() = Some(req.model.clone());
Ok(CompletionResponse {
message: Message::assistant(vec![Block::text("from the fallback")]),
stop_reason: StopReason::EndTurn,
usage: Usage::default(),
refusal: None,
model: "fallback-model".into(),
malformed_tool_args: 0,
})
}
}
fn req() -> CompletionRequest {
CompletionRequest {
model: "primary-model".into(),
system: None,
messages: vec![Message::user("hi")],
tools: Vec::new(),
max_tokens: 64,
effort: None,
thinking: false,
cache_prompt: false,
}
}
type Rig = (
Failover,
Arc<AtomicUsize>,
Arc<Mutex<Option<String>>>,
Arc<AtomicUsize>,
);
fn rig(error: fn() -> anyhow::Error) -> Rig {
let primary_calls = Arc::new(AtomicUsize::new(0));
let fallback_calls = Arc::new(AtomicUsize::new(0));
let model_seen = Arc::new(Mutex::new(None));
let failover = Failover::new(
Box::new(Failing {
error,
calls: Arc::clone(&primary_calls),
}),
vec![(
"small".into(),
Box::new(Recording {
model_seen: Arc::clone(&model_seen),
calls: Arc::clone(&fallback_calls),
}) as Box<dyn Provider>,
)],
);
(failover, primary_calls, model_seen, fallback_calls)
}
#[tokio::test]
async fn a_transient_exhaustion_falls_back_and_the_fallback_answers_as_itself() {
let (failover, _, model_seen, _) =
rig(|| anyhow::Error::new(ProviderError::Overloaded).context("anthropic 529: busy"));
let response = failover.complete(&req(), None).await.unwrap();
assert_eq!(response.message.text(), "from the fallback");
assert_eq!(
model_seen.lock().unwrap().as_deref(),
Some("fallback-model")
);
}
#[tokio::test]
async fn terminal_classes_never_fall_back() {
for error in [
(|| anyhow::Error::new(ProviderError::Invalid("bad".into()))) as fn() -> anyhow::Error,
|| anyhow::Error::new(ProviderError::Auth),
|| anyhow::Error::new(ProviderError::ContextOverflow),
] {
let (failover, primary_calls, _, fallback_calls) = rig(error);
let err = failover.complete(&req(), None).await.unwrap_err();
assert!(err.downcast_ref::<ProviderError>().is_some());
assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
assert_eq!(
fallback_calls.load(Ordering::SeqCst),
0,
"the fallback was consulted"
);
}
}
#[tokio::test]
async fn an_unclassified_error_never_falls_back_because_it_may_be_mid_stream() {
let (failover, _, _, fallback_calls) = rig(|| anyhow::anyhow!("stream aborted mid-body"));
let err = failover.complete(&req(), None).await.unwrap_err();
assert!(err.to_string().contains("stream aborted"));
assert_eq!(fallback_calls.load(Ordering::SeqCst), 0);
}
}