use crate::error::{Result, ShimError};
use crate::log::{LogEntry, Logger, RequestTimer};
use crate::policy::DispatchPolicyContext;
use crate::router::Router;
use serde_json::Value;
use std::time::Duration;
#[derive(Debug, Clone)]
pub struct FallbackConfig {
pub models: Vec<String>,
pub max_retries: u32,
pub initial_backoff: Duration,
pub retryable_statuses: Vec<u16>,
}
impl Default for FallbackConfig {
fn default() -> Self {
Self {
models: Vec::new(),
max_retries: 2,
initial_backoff: Duration::from_millis(500),
retryable_statuses: vec![429, 500, 502, 503, 529],
}
}
}
impl FallbackConfig {
pub fn new(models: Vec<String>) -> Self {
Self {
models,
..Default::default()
}
}
pub fn max_retries(mut self, n: u32) -> Self {
self.max_retries = n;
self
}
pub fn initial_backoff(mut self, d: Duration) -> Self {
self.initial_backoff = d;
self
}
}
fn is_retryable(err: &ShimError, retryable_statuses: &[u16]) -> bool {
match err {
ShimError::ProviderError { status, .. } => retryable_statuses.contains(status),
ShimError::Http(_) => true, _ => false,
}
}
pub async fn completion_with_fallback(
router: &Router,
request: &Value,
config: &FallbackConfig,
logger: Option<&Logger>,
) -> Result<Value> {
completion_with_fallback_inner(router, request, config, logger, None).await
}
pub async fn completion_with_fallback_and_policy(
router: &Router,
request: &Value,
config: &FallbackConfig,
logger: Option<&Logger>,
policy_context: &DispatchPolicyContext,
) -> Result<Value> {
completion_with_fallback_inner(router, request, config, logger, Some(policy_context)).await
}
async fn completion_with_fallback_inner(
router: &Router,
request: &Value,
config: &FallbackConfig,
logger: Option<&Logger>,
policy_context: Option<&DispatchPolicyContext>,
) -> Result<Value> {
let models = if config.models.is_empty() {
vec![request
.get("model")
.and_then(|m| m.as_str())
.ok_or(ShimError::MissingModel)?
.to_string()]
} else {
config.models.clone()
};
let mut errors: Vec<String> = Vec::new();
let mut policy_limited_providers = std::collections::HashSet::new();
let mut provider_limit_error: Option<ShimError> = None;
let mut attempted_distinct_provider_after_limit = false;
let client = crate::bound_client(router);
for model_str in &models {
let mut req = request.clone();
req["model"] = Value::String(model_str.clone());
let req = match router.expand_route(&req) {
Ok(expanded) => expanded.into_owned(),
Err(e) => {
errors.push(format!("{}: {}", model_str, e));
continue;
}
};
let (provider, model) = match router.resolve(model_str) {
Ok(r) => r,
Err(e) => {
errors.push(format!("{}: {}", model_str, e));
continue;
}
};
if policy_limited_providers.contains(provider.name()) {
errors.push(format!(
"{}: skipped provider {} after attempt policy refusal",
model_str,
provider.name()
));
continue;
}
if !policy_limited_providers.is_empty() {
attempted_distinct_provider_after_limit = true;
}
let mut backoff = config.initial_backoff;
for attempt in 0..=config.max_retries {
if !router.breaker().admit(provider.name()).await {
errors.push(format!(
"{}: circuit open for provider {}",
model_str,
provider.name()
));
break; }
let timer = RequestTimer::start();
let dispatch_outcome = client
.completion_dispatch(provider, &model, &req, policy_context)
.await;
if let Some(outcome) = crate::client::ShimClient::breaker_outcome(&dispatch_outcome) {
client.observe(provider, outcome).await;
}
match dispatch_outcome {
Ok(result) => {
if let Some(logger) = logger {
logger.log(&LogEntry::from_response(
provider.name(),
model_str,
&result,
timer.elapsed(),
));
}
return Ok(result);
}
Err(crate::client::DispatchFailure::Upstream(error)) => {
if is_retryable(&error, &config.retryable_statuses)
&& attempt < config.max_retries
{
errors.push(format!(
"{} (attempt {}): {}",
model_str,
attempt + 1,
error
));
tokio::time::sleep(backoff).await;
backoff *= 2;
continue;
}
if let Some(logger) = logger {
logger.log(&LogEntry::from_error(
provider.name(),
model_str,
&error.to_string(),
timer.elapsed(),
));
}
errors.push(format!("{}: {}", model_str, error));
break; }
Err(crate::client::DispatchFailure::Local(error)) => {
if let Some(logger) = logger {
logger.log(&LogEntry::from_error(
provider.name(),
model_str,
&error.to_string(),
timer.elapsed(),
));
}
errors.push(format!("{}: {}", model_str, error));
break;
}
Err(crate::client::DispatchFailure::PolicyRefusal(refusal))
if refusal.kind() == crate::policy::AttemptPolicyRefusalKind::ProviderLimit =>
{
let error = refusal.into_shim_error();
errors.push(format!("{}: {}", model_str, error));
policy_limited_providers.insert(provider.name().to_owned());
provider_limit_error = Some(error);
attempted_distinct_provider_after_limit = false;
break;
}
Err(crate::client::DispatchFailure::PolicyRefusal(refusal)) => {
let error = refusal.into_shim_error();
if let Some(logger) = logger {
logger.log(&LogEntry::from_error(
provider.name(),
model_str,
&error.to_string(),
timer.elapsed(),
));
}
return Err(error);
}
Err(crate::client::DispatchFailure::PolicyObservation(policy_error)) => {
let error = policy_error.into_shim_error();
if let Some(logger) = logger {
logger.log(&LogEntry::from_error(
provider.name(),
model_str,
&error.to_string(),
timer.elapsed(),
));
}
return Err(error);
}
}
}
}
if !attempted_distinct_provider_after_limit {
if let Some(error) = provider_limit_error {
return Err(error);
}
}
Err(ShimError::AllFailed(errors))
}