use super::Provider;
use crate::util::error::HttpError;
use crate::{ChatRequest, ChatResponse};
use async_trait::async_trait;
use std::time::Duration;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum ErrorClass {
Retryable,
NonRetryable,
}
impl ErrorClass {
const fn reason_label(self) -> &'static str {
match self {
Self::Retryable => "retryable",
Self::NonRetryable => "non_retryable",
}
}
}
const NON_RETRYABLE_HINTS: &[&str] = &[
"insufficient balance",
"insufficient_quota",
"quota exhausted",
"quota exceeded",
"error code 1113",
];
fn classify_err(err: &anyhow::Error) -> ErrorClass {
let msg = err.to_string();
let lower = msg.to_lowercase();
let status = err.downcast_ref::<HttpError>().map(|e| e.status);
if NON_RETRYABLE_HINTS.iter().any(|h| lower.contains(h)) {
return ErrorClass::NonRetryable;
}
if status.is_some_and(|c| (400..500).contains(&c) && c != 408 && c != 429) {
return ErrorClass::NonRetryable;
}
if lower.contains("model")
&& (lower.contains("not found")
|| lower.contains("unknown")
|| lower.contains("unsupported")
|| lower.contains("does not exist"))
{
return ErrorClass::NonRetryable;
}
ErrorClass::Retryable
}
fn parse_retry_after_ms(err: &anyhow::Error) -> Option<u64> {
if let Some(http_err) = err.downcast_ref::<HttpError>() {
return http_err.retry_after_ms;
}
None
}
pub struct ReliableProvider {
name: String,
provider: Box<dyn Provider>,
max_retries: u32,
base_backoff_ms: u64,
}
impl ReliableProvider {
#[must_use]
pub fn new(
name: String,
provider: Box<dyn Provider>,
max_retries: u32,
base_backoff_ms: u64,
) -> Self {
Self {
name,
provider,
max_retries,
base_backoff_ms: base_backoff_ms.max(50),
}
}
fn compute_backoff(base: u64, err: &anyhow::Error) -> u64 {
if let Some(retry_after) = parse_retry_after_ms(err) {
retry_after.min(30_000).max(base)
} else {
let half_range = base / 2;
base - base / 4 + (rand::random::<u64>() % half_range)
}
}
}
#[async_trait]
impl Provider for ReliableProvider {
async fn warmup(&self) -> anyhow::Result<()> {
self.provider.warmup().await
}
async fn chat(&self, request: ChatRequest) -> anyhow::Result<ChatResponse> {
let mut failures = Vec::new();
let mut backoff_ms = self.base_backoff_ms;
for attempt in 0..=self.max_retries {
match self.provider.chat(request.clone()).await {
Ok(resp) => {
if attempt > 0 {
tracing::info!(
provider = self.name,
attempt,
"Provider recovered after retry"
);
}
return Ok(resp);
}
Err(e) => {
let class = classify_err(&e);
let error_detail = e.to_string();
let reason = class.reason_label();
failures.push(format!(
"provider={} attempt {}/{}: {}; error={}",
self.name,
attempt + 1,
self.max_retries + 1,
reason,
error_detail,
));
let can_retry = class == ErrorClass::Retryable;
if can_retry && attempt < self.max_retries {
let wait = Self::compute_backoff(backoff_ms, &e);
if !crate::shutdown::sleep_or_shutdown(Duration::from_millis(wait)).await {
tracing::info!(
provider = self.name,
attempt = attempt + 1,
"Provider shutting down — aborting retry loop"
);
break;
}
tracing::warn!(
provider = self.name,
attempt = attempt + 1,
reason,
error = %error_detail,
"Provider call failed, retrying"
);
backoff_ms = backoff_ms.saturating_mul(2);
} else {
let log_msg = match class {
ErrorClass::NonRetryable => "Non-retryable error, aborting",
ErrorClass::Retryable => "Exhausted retries",
};
tracing::warn!(
provider = self.name,
attempt = attempt + 1,
reason,
error = %error_detail,
"{log_msg}"
);
break;
}
}
}
}
anyhow::bail!("All attempts failed.\n{}", failures.join("\n"))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ChatMessage;
use crate::providers::test_request;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
fn test_err(status: u16, body: &str) -> anyhow::Error {
anyhow::Error::from(HttpError::new(status, "test", body, None))
}
struct TestProvider {
calls: Arc<AtomicUsize>,
fail_until_attempt: usize,
response_text: &'static str,
error: &'static str,
context_overflow: bool,
tool_schema_error: bool,
tool_calls: Vec<crate::ToolCall>,
warmup_fails: bool,
}
impl TestProvider {
fn new(response_text: &'static str) -> Self {
Self {
calls: Arc::new(AtomicUsize::new(0)),
fail_until_attempt: 0,
response_text,
error: "mock error",
context_overflow: false,
tool_schema_error: false,
tool_calls: Vec::new(),
warmup_fails: false,
}
}
fn with_fail(mut self, until_attempt: usize, error: &'static str) -> Self {
self.fail_until_attempt = until_attempt;
self.error = error;
self
}
fn with_context_overflow(mut self, fail_until: usize) -> Self {
self.context_overflow = true;
self.fail_until_attempt = fail_until;
self
}
fn with_tool_schema_error(mut self, fail_until: usize) -> Self {
self.tool_schema_error = true;
self.fail_until_attempt = fail_until;
self
}
fn with_calls(mut self, calls: Arc<AtomicUsize>) -> Self {
self.calls = calls;
self
}
fn with_warmup_fail(mut self) -> Self {
self.warmup_fails = true;
self
}
fn make_error(&self) -> String {
if self.context_overflow {
"request (8968 tokens) exceeds the available context size (8448 tokens), try increasing it".to_string()
} else if self.tool_schema_error {
"tool call validation failed: attempted to call tool 'recall' which was not in request".to_string()
} else {
self.error.to_string()
}
}
fn check_fail(&self, attempt: usize) -> bool {
attempt <= self.fail_until_attempt
}
}
#[async_trait]
impl Provider for TestProvider {
async fn chat(&self, _request: ChatRequest) -> anyhow::Result<ChatResponse> {
let call = self.calls.fetch_add(1, Ordering::SeqCst);
if self.check_fail(call + 1) {
if self.context_overflow {
return Err(test_err(400, &self.make_error()));
}
if self.tool_schema_error {
return Err(test_err(400, &self.make_error()));
}
anyhow::bail!("{}", self.make_error());
}
Ok(ChatResponse {
text: Some(self.response_text.to_string()),
tool_calls: self.tool_calls.clone(),
..Default::default()
})
}
async fn warmup(&self) -> anyhow::Result<()> {
if self.warmup_fails {
anyhow::bail!("warmup failed");
}
Ok(())
}
}
#[test]
fn retryable_error_classification() {
let is_non_retryable =
|e: &anyhow::Error| matches!(classify_err(e), ErrorClass::NonRetryable);
assert!(is_non_retryable(&test_err(401, "Unauthorized")));
assert!(is_non_retryable(&test_err(403, "Forbidden")));
assert!(is_non_retryable(&test_err(400, "invalid api key")));
assert!(is_non_retryable(&anyhow::anyhow!("model not found")));
assert!(is_non_retryable(&anyhow::anyhow!("model 'xyz' is unknown")));
assert!(is_non_retryable(&anyhow::anyhow!("insufficient balance")));
assert!(is_non_retryable(&anyhow::anyhow!("insufficient_quota")));
assert!(is_non_retryable(&anyhow::anyhow!("quota exhausted")));
assert!(is_non_retryable(&anyhow::anyhow!("error code 1113")));
assert!(!is_non_retryable(&anyhow::anyhow!("500 Server Error")));
assert!(!is_non_retryable(&anyhow::anyhow!("502 Bad Gateway")));
assert!(!is_non_retryable(&anyhow::anyhow!(
"503 Service Unavailable"
)));
assert!(!is_non_retryable(&anyhow::anyhow!("connection reset")));
assert!(!is_non_retryable(&anyhow::anyhow!(
"model overloaded, try again later"
)));
}
#[tokio::test]
async fn chat_retries_then_recovers() {
let calls = Arc::new(AtomicUsize::new(0));
let provider = ReliableProvider::new(
"primary".into(),
Box::new(
TestProvider::new("history ok")
.with_fail(1, "temporary")
.with_calls(calls.clone()),
) as Box<dyn Provider>,
2,
50,
);
let messages = vec![ChatMessage::system("system"), ChatMessage::user("hello")];
let result = provider
.chat(test_request(messages.clone(), None))
.await
.unwrap();
assert_eq!(result.text.as_deref(), Some("history ok"));
assert_eq!(calls.load(Ordering::SeqCst), 2);
}
#[test]
fn backoff_and_retry_after() {
let with_retry = HttpError::new(429, "test", "rate limited", Some(5000));
assert_eq!(
parse_retry_after_ms(&anyhow::Error::from(with_retry)),
Some(5000)
);
let no_retry = test_err(429, "rate limit");
assert_eq!(parse_retry_after_ms(&no_retry), None);
let structured =
anyhow::Error::from(HttpError::new(429, "test", "rate limited", Some(3_000)));
assert_eq!(ReliableProvider::compute_backoff(500, &structured), 3_000);
let with_long_retry =
anyhow::Error::from(HttpError::new(429, "test", "rate limit", Some(120_000)));
assert_eq!(
ReliableProvider::compute_backoff(500, &with_long_retry),
30_000
);
let no_header = test_err(500, "error");
let backoff = ReliableProvider::compute_backoff(500, &no_header);
assert!(
(375..625).contains(&backoff),
"expected backoff in [375, 625), got {backoff}"
);
}
#[test]
fn classify_err_typed_path() {
assert!(matches!(
classify_err(&test_err(429, "Too Many Requests")),
ErrorClass::Retryable
));
assert!(matches!(
classify_err(&test_err(429, "rate limit exceeded")),
ErrorClass::Retryable
));
assert_eq!(
classify_err(&test_err(429, "insufficient balance")),
ErrorClass::NonRetryable
);
assert_eq!(
classify_err(&test_err(429, "quota exhausted")),
ErrorClass::NonRetryable
);
assert!(matches!(
classify_err(&test_err(400, "Bad Request")),
ErrorClass::NonRetryable
));
assert!(matches!(
classify_err(&test_err(403, "Forbidden")),
ErrorClass::NonRetryable
));
assert!(matches!(
classify_err(&test_err(408, "Request Timeout")),
ErrorClass::Retryable
));
assert!(matches!(
classify_err(&test_err(500, "Internal Server Error")),
ErrorClass::Retryable
));
assert!(matches!(
classify_err(&test_err(400, "exceeds the context window of this model")),
ErrorClass::NonRetryable
));
assert!(matches!(
classify_err(&test_err(400, "tool call validation failed")),
ErrorClass::NonRetryable
));
assert!(matches!(
classify_err(&test_err(403, "unauthorized")),
ErrorClass::NonRetryable
));
assert!(matches!(
classify_err(&test_err(404, "model not found")),
ErrorClass::NonRetryable
));
assert_eq!(
classify_err(&test_err(429, "error code 1113")),
ErrorClass::NonRetryable
);
assert_eq!(
classify_err(&test_err(
502,
"Your chosen model is down or we received an invalid response from it"
)),
ErrorClass::Retryable
);
}
#[tokio::test]
async fn chat_returns_aggregated_error_when_all_retries_exhausted() {
let provider = ReliableProvider::new(
"p1".into(),
Box::new(TestProvider::new("never").with_fail(usize::MAX, "p1 chat error"))
as Box<dyn Provider>,
0,
1,
);
let messages = vec![ChatMessage::user("hello")];
let request = test_request(messages.clone(), None);
let err = provider
.chat(request)
.await
.expect_err("all attempts should fail");
let msg = err.to_string();
assert!(msg.contains("All attempts failed"));
assert!(msg.contains("provider=p1"));
assert!(msg.contains("error=p1 chat error"));
assert!(msg.contains("retryable"));
}
#[tokio::test]
async fn warmup_propagates_inner_error() {
let inner = TestProvider::new("unused").with_warmup_fail();
let provider =
ReliableProvider::new("test".into(), Box::new(inner) as Box<dyn Provider>, 0, 1);
let err = provider
.warmup()
.await
.expect_err("warmup should propagate error");
assert!(
err.to_string().contains("warmup failed"),
"expected 'warmup failed', got: {err}"
);
}
#[tokio::test]
async fn warmup_ok_when_inner_succeeds() {
let inner = TestProvider::new("ok");
let provider =
ReliableProvider::new("test".into(), Box::new(inner) as Box<dyn Provider>, 0, 1);
provider.warmup().await.expect("warmup should succeed");
}
#[test]
fn context_window_error_classification() {
let is_non_retryable =
|e: &anyhow::Error| matches!(classify_err(e), ErrorClass::NonRetryable);
assert!(is_non_retryable(&test_err(
400,
"request (8968 tokens) exceeds the available context size (8448 tokens)",
)));
assert!(is_non_retryable(&test_err(
400,
"This model's maximum context length is 8192 tokens",
)));
assert!(is_non_retryable(&test_err(
400,
"maximum context length of this model is 128K tokens",
)));
assert!(is_non_retryable(&test_err(401, "Unauthorized")));
}
#[tokio::test]
async fn chat_context_window_exceeded_is_not_retried() {
let calls = Arc::new(AtomicUsize::new(0));
let provider = ReliableProvider::new(
"primary".into(),
Box::new(
TestProvider::new("ok after overflow")
.with_context_overflow(2)
.with_calls(calls.clone()),
) as Box<dyn Provider>,
3,
1,
);
let messages = vec![ChatMessage::user("test")];
let result = provider.chat(test_request(messages.clone(), None)).await;
assert!(
result.is_err(),
"context window errors are non-retryable, should fail immediately"
);
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"should not retry context overflow"
);
}
#[test]
fn tool_schema_error_detection() {
use ErrorClass::NonRetryable;
for msg in [
r#"Groq API error (400 Bad Request): {"error":{"message":"tool call validation failed: attempted to call tool 'recall' which was not in request"}}"#,
"tool 'search' which was not in request",
"function 'foo' not found in tool list",
"invalid_tool_call: no matching function",
] {
assert!(
matches!(classify_err(&test_err(400, msg)), NonRetryable),
"should detect: {msg}"
);
}
assert!(
matches!(
classify_err(&test_err(400, "invalid api key provided")),
NonRetryable
),
"pure 400 should be NonRetryable"
);
}
#[test]
fn non_retryable_hints_are_classified_non_retryable() {
for hint in NON_RETRYABLE_HINTS {
let err = anyhow::anyhow!("some error: {hint}");
assert!(
matches!(classify_err(&err), ErrorClass::NonRetryable),
"hint '{hint}' should be classified as NonRetryable"
);
}
}
#[tokio::test]
async fn chat_tool_schema_error_is_not_retried() {
let calls = Arc::new(AtomicUsize::new(0));
let provider = ReliableProvider::new(
"primary".into(),
Box::new(
TestProvider::new("unused")
.with_tool_schema_error(10)
.with_calls(calls.clone()),
) as Box<dyn Provider>,
3,
1,
);
let messages = vec![ChatMessage::user("test")];
let result = provider.chat(test_request(messages.clone(), None)).await;
assert!(
result.is_err(),
"tool schema errors are non-retryable, should fail immediately"
);
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"should not retry tool schema errors"
);
}
}