use std::sync::atomic::{AtomicU8, Ordering};
use crate::error::{Error, Result};
use crate::provider::{CompletionRequest, CompletionResponse, Provider};
#[derive(Debug)]
pub struct Fallback<A, B> {
primary: A,
secondary: B,
name: String,
served: AtomicU8,
}
impl<A: Provider, B: Provider> Fallback<A, B> {
pub fn new(primary: A, secondary: B) -> Self {
let name = format!("{} -> {}", primary.name(), secondary.name());
Self {
primary,
secondary,
name,
served: AtomicU8::new(0),
}
}
fn note(&self, who: u8) {
self.served.store(who, Ordering::Relaxed);
}
}
fn worth_another_provider(e: &Error) -> bool {
matches!(e, Error::Provider { kind, .. } if kind.is_retryable())
}
impl<A: Provider + Sync, B: Provider + Sync> Provider for Fallback<A, B> {
#[cfg(feature = "media")]
fn accepts_images(&self) -> bool {
self.primary.accepts_images() && self.secondary.accepts_images()
}
async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse> {
match self.primary.complete(request.clone()).await {
Ok(response) => {
self.note(1);
Ok(response)
}
Err(e) if worth_another_provider(&e) => {
tracing::warn!(
primary = self.primary.name(),
secondary = self.secondary.name(),
error = %e,
"provider failed; falling over"
);
let out = self.secondary.complete(request).await;
self.note(if out.is_ok() { 2 } else { 0 });
out
}
Err(e) => {
self.note(0);
Err(e)
}
}
}
fn name(&self) -> &str {
&self.name
}
fn endpoint(&self) -> Option<&str> {
self.primary.endpoint()
}
fn endpoints(&self) -> Vec<&str> {
let mut out = self.primary.endpoints();
out.extend(self.secondary.endpoints());
out
}
fn last_served(&self) -> Option<String> {
match self.served.load(Ordering::Relaxed) {
1 => Some(
self.primary
.last_served()
.unwrap_or_else(|| self.primary.name().to_string()),
),
2 => Some(
self.secondary
.last_served()
.unwrap_or_else(|| self.secondary.name().to_string()),
),
_ => None,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::ProviderErrorKind;
struct Fixed {
label: &'static str,
result: fn() -> Result<CompletionResponse>,
endpoint: Option<&'static str>,
}
impl Provider for Fixed {
async fn complete(&self, _request: CompletionRequest) -> Result<CompletionResponse> {
(self.result)()
}
fn name(&self) -> &str {
self.label
}
fn endpoint(&self) -> Option<&str> {
self.endpoint
}
}
fn ok() -> Result<CompletionResponse> {
Ok(CompletionResponse {
text: Some("hello".into()),
..Default::default()
})
}
fn down() -> Result<CompletionResponse> {
Err(Error::provider_status(503, None, "down"))
}
fn bad_key() -> Result<CompletionResponse> {
Err(Error::provider_status(401, None, "bad key"))
}
fn boom() -> Result<CompletionResponse> {
panic!("the secondary must not be called");
}
#[allow(clippy::needless_update)] fn req() -> CompletionRequest {
CompletionRequest {
system: String::new(),
user: "hi".into(),
tools: Vec::new(),
..Default::default()
}
}
fn p(label: &'static str, result: fn() -> Result<CompletionResponse>) -> Fixed {
Fixed {
label,
result,
endpoint: None,
}
}
#[tokio::test]
async fn a_down_primary_is_answered_by_the_secondary() {
let f = Fallback::new(p("first", down), p("second", ok));
let out = f.complete(req()).await.unwrap();
assert_eq!(out.text.as_deref(), Some("hello"));
assert_eq!(f.last_served().as_deref(), Some("second"));
assert_eq!(f.name(), "first -> second");
}
#[tokio::test]
async fn a_working_primary_is_never_backed_up() {
let f = Fallback::new(p("first", ok), p("second", boom));
assert!(f.complete(req()).await.is_ok());
assert_eq!(f.last_served().as_deref(), Some("first"));
}
#[tokio::test]
async fn a_failure_the_secondary_would_share_does_not_fall_over() {
let f = Fallback::new(p("first", bad_key), p("second", boom));
let err = f.complete(req()).await.unwrap_err();
let Error::Provider { kind, status, .. } = err else {
panic!("expected a provider error, got {err:?}");
};
assert_eq!(kind, ProviderErrorKind::Auth);
assert_eq!(status, Some(401));
assert_eq!(f.last_served(), None);
}
#[tokio::test]
async fn both_failing_reports_the_secondarys_error() {
let f = Fallback::new(p("first", down), p("second", bad_key));
let err = f.complete(req()).await.unwrap_err();
let Error::Provider { kind, .. } = err else {
panic!("expected a provider error");
};
assert_eq!(kind, ProviderErrorKind::Auth);
}
#[tokio::test]
async fn three_providers_nest_and_the_leaf_is_what_gets_recorded() {
let f = Fallback::new(p("a", down), Fallback::new(p("b", down), p("c", ok)));
assert!(f.complete(req()).await.is_ok());
assert_eq!(f.name(), "a -> b -> c");
assert_eq!(f.last_served().as_deref(), Some("c"));
}
#[test]
fn authorization_sees_every_endpoint_in_the_chain() {
let a = Fixed {
label: "a",
result: ok,
endpoint: Some("https://a.example/v1"),
};
let b = Fixed {
label: "b",
result: ok,
endpoint: Some("https://b.example/v1"),
};
let f = Fallback::new(a, b);
assert_eq!(
f.endpoints(),
vec!["https://a.example/v1", "https://b.example/v1"]
);
let g = Fallback::new(p("x", ok), p("y", ok));
assert!(g.endpoints().is_empty());
}
}