use super::Provider;
use crate::providers::error::ProviderError;
use crate::util::http::extract_http_status;
use crate::{ChatRequest, ChatResponse, StreamEvent, StreamResult};
use async_trait::async_trait;
use futures_util::stream;
use std::time::Duration;
use reqwest;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum ErrorClass {
Retryable,
NonRetryable,
ToolSchemaError,
}
impl ErrorClass {
const fn reason_label(self) -> &'static str {
match self {
Self::Retryable => "retryable",
Self::NonRetryable => "non_retryable",
Self::ToolSchemaError => "tool_schema_error",
}
}
}
const CTX_HINTS: &[&str] = &[
"exceeds the context window",
"exceeds the available context size",
"context window of this model",
"maximum context length",
"context length exceeded",
"too many tokens",
"token limit exceeded",
"prompt is too long",
"input is too long",
"prompt exceeds max length",
];
const TOOL_SCHEMA_HINTS: &[&str] = &[
"tool call validation failed",
"which was not in request",
"not found in tool list",
"invalid_tool_call",
];
const AUTH_HINTS: &[&str] = &[
"invalid api key",
"incorrect api key",
"missing api key",
"api key not set",
"authentication failed",
"auth failed",
"unauthorized",
"forbidden",
"permission denied",
"access denied",
"invalid token",
];
const BILLING_HINTS: &[&str] = &[
"insufficient balance",
"insufficient_quota",
"quota exhausted",
"quota exceeded",
"error code 1113",
];
fn classify_fallback(lower: &str) -> ErrorClass {
if AUTH_HINTS.iter().any(|h| lower.contains(h)) {
return ErrorClass::NonRetryable;
}
if lower.contains("model")
&& (lower.contains("not found")
|| lower.contains("unknown")
|| lower.contains("unsupported")
|| lower.contains("does not exist")
|| lower.contains("invalid"))
{
return ErrorClass::NonRetryable;
}
if BILLING_HINTS.iter().any(|h| lower.contains(h)) {
return ErrorClass::NonRetryable;
}
ErrorClass::Retryable
}
#[inline]
fn classify_by_status_code(status: Option<u16>, lower: &str) -> ErrorClass {
if status.is_some_and(is_non_retryable_4xx) {
ErrorClass::NonRetryable
} else {
classify_fallback(lower)
}
}
fn is_non_retryable_4xx(code: u16) -> bool {
(400..500).contains(&code) && code != 408 && code != 429
}
fn classify_err(err: &anyhow::Error) -> ErrorClass {
let msg = err.to_string();
let lower = msg.to_lowercase();
if CTX_HINTS.iter().any(|h| lower.contains(h)) {
return ErrorClass::NonRetryable;
}
if TOOL_SCHEMA_HINTS.iter().any(|h| lower.contains(h)) {
return ErrorClass::ToolSchemaError;
}
if let Some(provider_err) = err.downcast_ref::<ProviderError>() {
return classify_by_status_code(Some(provider_err.status), &lower);
}
if let Some(transport_err) = err.downcast_ref::<reqwest::Error>() {
return classify_transport_err(transport_err, &lower);
}
classify_by_status_code(extract_http_status(&msg), &lower)
}
fn classify_transport_err(transport_err: &reqwest::Error, lower: &str) -> ErrorClass {
if transport_err.is_timeout() || transport_err.is_connect() {
return ErrorClass::Retryable;
}
if transport_err.is_builder() || transport_err.is_redirect() {
return ErrorClass::NonRetryable;
}
if transport_err.is_status()
&& let Some(status) = transport_err.status()
{
let code = status.as_u16();
return classify_by_status_code(Some(code), lower);
}
ErrorClass::Retryable
}
fn parse_retry_after_ms(err: &anyhow::Error) -> Option<u64> {
if let Some(provider_err) = err.downcast_ref::<ProviderError>() {
return provider_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).max(1); let jittered_backoff = base - base / 4 + (rand::random::<u64>() % half_range);
jittered_backoff.max(50) }
}
}
#[async_trait]
impl Provider for ReliableProvider {
async fn warmup(&self) -> anyhow::Result<()> {
if let Err(e) = self.provider.warmup().await {
tracing::warn!(
provider = self.name,
"Connection warmup failed (non-fatal): {e}"
);
}
Ok(())
}
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 {
if crate::shutdown::shutdown_token().is_cancelled() {
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"
);
let wait = Self::compute_backoff(backoff_ms, &e);
if !crate::shutdown::sleep_or_shutdown(Duration::from_millis(wait)).await {
break;
}
backoff_ms = backoff_ms.saturating_mul(2);
} else {
let log_msg = match class {
ErrorClass::NonRetryable | ErrorClass::ToolSchemaError => {
"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"))
}
fn stream_chat(
&self,
request: ChatRequest,
) -> stream::BoxStream<'static, StreamResult<StreamEvent>> {
self.provider.stream_chat(request)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{ChatMessage, ToolSpec};
use futures_util::StreamExt;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
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>,
}
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(),
}
}
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_tool_calls(mut self, tool_calls: Vec<crate::ToolCall>) -> Self {
self.tool_calls = tool_calls;
self
}
fn with_calls(mut self, calls: Arc<AtomicUsize>) -> Self {
self.calls = calls;
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) {
anyhow::bail!("{}", self.make_error());
}
Ok(ChatResponse {
text: Some(self.response_text.to_string()),
tool_calls: self.tool_calls.clone(),
..Default::default()
})
}
fn stream_chat(
&self,
_request: ChatRequest,
) -> stream::BoxStream<'static, StreamResult<StreamEvent>> {
stream::iter(vec![
Ok(StreamEvent::ToolCall(crate::ToolCall {
id: "call_1".to_string(),
name: "shell".to_string(),
arguments: serde_json::json!({"command": "date"}),
})),
Ok(StreamEvent::Final),
])
.boxed()
}
}
#[test]
fn retryable_error_classification() {
let is_non_retryable =
|e: &anyhow::Error| matches!(classify_err(e), ErrorClass::NonRetryable);
assert!(is_non_retryable(&anyhow::anyhow!("401 Unauthorized")));
assert!(is_non_retryable(&anyhow::anyhow!("403 Forbidden")));
assert!(is_non_retryable(&anyhow::anyhow!("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!("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,
1,
);
let messages = vec![ChatMessage::system("system"), ChatMessage::user("hello")];
let result = provider
.chat(ChatRequest {
messages: messages.clone(),
tools: None,
model: "test".to_string(),
allow_image_parts: false,
temperature: 0.1,
reasoning_effort: None,
provider_order: None,
provider_allow_fallbacks: None,
})
.await
.unwrap();
assert_eq!(result.text.as_deref(), Some("history ok"));
assert_eq!(calls.load(Ordering::SeqCst), 2);
}
#[test]
fn rate_limit_utilities() {
let with_retry_after =
anyhow::Error::from(ProviderError::new(429, "test", "rate limited", Some(3_000)));
assert_eq!(
ReliableProvider::compute_backoff(500, &with_retry_after),
3_000
);
let with_long_retry =
anyhow::Error::from(ProviderError::new(429, "test", "rate limit", Some(120_000)));
assert_eq!(
ReliableProvider::compute_backoff(500, &with_long_retry),
30_000
);
let no_retry = anyhow::Error::from(ProviderError::new(500, "test", "error", None));
let backoff = ReliableProvider::compute_backoff(500, &no_retry);
assert!(
(375..625).contains(&backoff),
"expected backoff in [375, 625), got {backoff}"
);
}
#[test]
fn classify_err_typed_path() {
let make_structured = |status: u16, body: &str| -> anyhow::Error {
anyhow::Error::from(ProviderError::new(status, "test", body, None))
};
assert!(matches!(
classify_err(&make_structured(429, "Too Many Requests")),
ErrorClass::Retryable
));
assert!(matches!(
classify_err(&make_structured(429, "rate limit exceeded")),
ErrorClass::Retryable
));
assert!(matches!(
classify_err(&make_structured(429, "insufficient balance")),
ErrorClass::NonRetryable
));
assert_eq!(
classify_err(&make_structured(429, "insufficient balance")),
ErrorClass::NonRetryable
);
assert_eq!(
classify_err(&make_structured(429, "quota exhausted")),
ErrorClass::NonRetryable
);
assert!(matches!(
classify_err(&make_structured(400, "Bad Request")),
ErrorClass::NonRetryable
));
assert!(matches!(
classify_err(&make_structured(403, "Forbidden")),
ErrorClass::NonRetryable
));
assert!(matches!(
classify_err(&make_structured(408, "Request Timeout")),
ErrorClass::Retryable
));
assert!(matches!(
classify_err(&make_structured(500, "Internal Server Error")),
ErrorClass::Retryable
));
assert!(matches!(
classify_err(&make_structured(
400,
"exceeds the context window of this model"
)),
ErrorClass::NonRetryable
));
assert!(matches!(
classify_err(&make_structured(400, "tool call validation failed")),
ErrorClass::ToolSchemaError
));
assert!(matches!(
classify_err(&make_structured(403, "unauthorized")),
ErrorClass::NonRetryable
));
assert!(matches!(
classify_err(&make_structured(404, "model not found")),
ErrorClass::NonRetryable
));
assert_eq!(
classify_err(&make_structured(429, "error code 1113")),
ErrorClass::NonRetryable
);
}
#[test]
fn parse_retry_after_typed_path() {
let with_retry = ProviderError::new(429, "test", "rate limited", Some(5000));
assert_eq!(
parse_retry_after_ms(&anyhow::Error::from(with_retry)),
Some(5000)
);
let no_retry = ProviderError::new(429, "test", "rate limit", None);
assert_eq!(parse_retry_after_ms(&anyhow::Error::from(no_retry)), None);
let structured =
anyhow::Error::from(ProviderError::new(429, "test", "rate limited", Some(3000)));
assert_eq!(ReliableProvider::compute_backoff(500, &structured), 3_000);
let no_header = anyhow::Error::from(ProviderError::new(500, "test", "error", None));
let backoff = ReliableProvider::compute_backoff(500, &no_header);
assert!(
(375..625).contains(&backoff),
"expected backoff in [375, 625), got {backoff}"
);
}
#[tokio::test]
async fn chat_retries_and_recovers() {
let tool_call = crate::ToolCall {
id: "call_1".to_string(),
name: "shell".to_string(),
arguments: serde_json::json!({"command": "date"}),
};
let calls = Arc::new(AtomicUsize::new(0));
let provider = ReliableProvider::new(
"primary".into(),
Box::new(
TestProvider::new("recovered")
.with_fail(2, "temporary failure")
.with_tool_calls(vec![tool_call])
.with_calls(calls.clone()),
) as Box<dyn Provider>,
3,
1,
);
let messages = vec![ChatMessage::user("test")];
let request = ChatRequest {
messages: messages.clone(),
tools: None,
model: "test".to_string(),
allow_image_parts: false,
temperature: 0.1,
reasoning_effort: None,
provider_order: None,
provider_allow_fallbacks: None,
};
let result = provider.chat(request).await.unwrap();
assert_eq!(result.text.as_deref(), Some("recovered"));
assert!(
calls.load(Ordering::SeqCst) > 1,
"should have retried at least once"
);
}
#[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 = ChatRequest {
messages: messages.clone(),
tools: None,
model: "test".to_string(),
allow_image_parts: false,
temperature: 0.1,
reasoning_effort: None,
provider_order: None,
provider_allow_fallbacks: 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"));
}
#[test]
fn context_window_error_classification() {
let is_non_retryable =
|e: &anyhow::Error| matches!(classify_err(e), ErrorClass::NonRetryable);
assert!(is_non_retryable(&anyhow::anyhow!(
"request (8968 tokens) exceeds the available context size (8448 tokens)"
)));
assert!(is_non_retryable(&anyhow::anyhow!(
"This model's maximum context length is 8192 tokens"
)));
assert!(is_non_retryable(&anyhow::anyhow!(
"maximum context length of this model is 128K tokens"
)));
assert!(is_non_retryable(&anyhow::anyhow!("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(ChatRequest {
messages: messages.clone(),
tools: None,
model: "test".to_string(),
allow_image_parts: false,
temperature: 0.1,
reasoning_effort: None,
provider_order: None,
provider_allow_fallbacks: 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"
);
}
#[tokio::test]
async fn chat_context_window_exceeded_eventually_fails() {
let calls = Arc::new(AtomicUsize::new(0));
let provider = ReliableProvider::new(
"primary".into(),
Box::new(
TestProvider::new("never succeeds")
.with_context_overflow(usize::MAX)
.with_calls(calls.clone()),
) as Box<dyn Provider>,
3,
1,
);
let messages = vec![ChatMessage::user("test")];
let result = provider
.chat(ChatRequest {
messages: messages.clone(),
tools: None,
model: "test".to_string(),
allow_image_parts: false,
temperature: 0.1,
reasoning_effort: None,
provider_order: None,
provider_allow_fallbacks: None,
})
.await;
assert!(result.is_err());
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("All attempts failed")
&& err_msg.contains("exceeds the available context size"),
"Should aggregate failures, got: {err_msg}"
);
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"no retries for context window errors"
);
}
#[test]
fn tool_schema_error_detection() {
use ErrorClass::ToolSchemaError;
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(&anyhow::anyhow!("{msg}")), ToolSchemaError),
"should detect: {msg}"
);
}
for msg in ["invalid api key", "model not found"] {
assert!(
!matches!(classify_err(&anyhow::anyhow!("{msg}")), ToolSchemaError),
"should ignore: {msg}"
);
}
}
#[test]
fn non_retryable_400_handling() {
let is_non_retryable =
|e: &anyhow::Error| matches!(classify_err(e), ErrorClass::NonRetryable);
assert!(!is_non_retryable(&anyhow::anyhow!(
"{}",
"400 Bad Request: tool call validation failed: attempted to call tool 'x' which was not in request"
)));
assert!(is_non_retryable(&anyhow::anyhow!(
"{}",
"400 Bad Request: invalid api key provided"
)));
}
#[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(ChatRequest {
messages: messages.clone(),
tools: None,
model: "test".to_string(),
allow_image_parts: false,
temperature: 0.1,
reasoning_effort: None,
provider_order: None,
provider_allow_fallbacks: 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"
);
}
#[tokio::test]
async fn stream_chat_works_when_provider_supports_tool_events() {
let provider =
ReliableProvider::new("primary".into(), Box::new(TestProvider::new("ok")), 0, 1);
let request = ChatRequest {
messages: vec![ChatMessage::user("hello")],
tools: Some(vec![ToolSpec {
name: "test".into(),
description: "A test tool".into(),
parameters: serde_json::json!({}),
}]),
model: "test".to_string(),
allow_image_parts: false,
temperature: 0.1,
reasoning_effort: None,
provider_order: None,
provider_allow_fallbacks: None,
};
let mut stream = provider.stream_chat(request);
let first = stream.next().await.unwrap().unwrap();
if let StreamEvent::ToolCall(tc) = first {
assert_eq!(tc.name, "shell");
} else {
panic!("expected ToolCall event");
}
}
}