mod types;
pub use types::*;
use std::sync::Mutex;
use std::time::{Duration, Instant};
use crate::error::TinyAgentsError;
use crate::harness::model::ProviderError;
fn parse_status_at(text: &str, start: usize) -> Option<u16> {
let digits: String = text
.get(start..)?
.trim_start()
.chars()
.take_while(|ch| ch.is_ascii_digit())
.collect();
if digits.len() == 3 {
digits.parse().ok()
} else {
None
}
}
fn find_case_insensitive(haystack: &str, needle: &str) -> Option<usize> {
haystack
.to_ascii_lowercase()
.find(&needle.to_ascii_lowercase())
}
pub fn structured_http_status(message: &str) -> Option<u16> {
let trimmed = message.trim_start();
if let Some(status) = parse_status_at(trimmed, 0) {
return Some(status);
}
for (idx, _) in message.match_indices('(') {
if let Some(status) = parse_status_at(message, idx + 1) {
return Some(status);
}
}
for marker in ["http ", "status:", "status "] {
if let Some(idx) = find_case_insensitive(message, marker)
&& let Some(status) = parse_status_at(message, idx + marker.len())
{
return Some(status);
}
}
None
}
fn is_retryable_status(status: u16) -> bool {
matches!(status, 408 | 409 | 429) || status >= 500
}
fn is_upstream_unhealthy_status(status: u16) -> bool {
matches!(status, 408 | 409) || status >= 500
}
fn text_indicates_rate_limit(lower: &str) -> bool {
lower.contains("429")
&& (lower.contains("too many") || lower.contains("rate") || lower.contains("limit"))
}
fn text_indicates_upstream_unhealthy(lower: &str) -> bool {
lower.contains("no healthy upstream")
|| lower.contains("upstream unavailable")
|| lower.contains("service unavailable")
|| lower.contains("408 request timeout")
|| lower.contains("409 conflict")
|| lower.contains("500 internal server error")
|| lower.contains("502 bad gateway")
|| lower.contains("503 service unavailable")
|| lower.contains("504 gateway timeout")
}
fn text_indicates_non_retryable(lower: &str) -> bool {
let auth_failure_hints = [
"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",
];
if auth_failure_hints.iter().any(|hint| lower.contains(hint)) {
return true;
}
lower.contains("model")
&& (lower.contains("not found")
|| lower.contains("unknown")
|| lower.contains("unsupported")
|| lower.contains("does not exist")
|| lower.contains("invalid"))
}
fn text_indicates_non_retryable_rate_limit(lower: &str) -> bool {
let business_hints = [
"plan does not include",
"doesn't include",
"not include",
"insufficient balance",
"insufficient_balance",
"insufficient quota",
"insufficient_quota",
"quota exhausted",
"out of credits",
"no available package",
"package not active",
"purchase package",
"model not available for your plan",
];
if business_hints.iter().any(|hint| lower.contains(hint)) {
return true;
}
lower.split(|ch: char| !ch.is_ascii_digit()).any(|token| {
token
.parse::<u16>()
.is_ok_and(|code| matches!(code, 1113 | 1311))
})
}
pub fn classify_provider_failure(
status: Option<u16>,
code: Option<&str>,
message: &str,
) -> ProviderFailureClass {
let status = status.or_else(|| structured_http_status(message));
let lower = match code {
Some(code) if !code.trim().is_empty() => format!("{message} {code}").to_ascii_lowercase(),
_ => message.to_ascii_lowercase(),
};
if status == Some(429) || text_indicates_rate_limit(&lower) {
if text_indicates_non_retryable_rate_limit(&lower) {
return ProviderFailureClass::NonRetryableRateLimit;
}
return ProviderFailureClass::RateLimited;
}
if let Some(status) = status {
if is_upstream_unhealthy_status(status) {
return ProviderFailureClass::UpstreamUnhealthy;
}
if (400..500).contains(&status) && !is_retryable_status(status) {
return ProviderFailureClass::NonRetryable;
}
}
if text_indicates_upstream_unhealthy(&lower) {
return ProviderFailureClass::UpstreamUnhealthy;
}
if text_indicates_non_retryable(&lower) {
return ProviderFailureClass::NonRetryable;
}
ProviderFailureClass::Retryable
}
pub fn classify_provider_error(error: &ProviderError) -> ProviderFailureClass {
classify_provider_failure(error.status, error.code.as_deref(), &error.message)
}
pub fn provider_error_is_retryable(error: &ProviderError) -> bool {
classify_provider_error(error).is_retryable()
}
pub fn parse_retry_after_ms(message: &str) -> Option<u64> {
let lower = message.to_ascii_lowercase();
for prefix in &[
"retry-after:",
"retry_after:",
"retry-after ",
"retry_after ",
] {
if let Some(pos) = lower.find(prefix) {
let after = &message[pos + prefix.len()..];
let number: String = after
.trim_start()
.chars()
.take_while(|ch| ch.is_ascii_digit() || *ch == '.')
.collect();
if let Ok(seconds) = number.parse::<f64>()
&& seconds.is_finite()
&& seconds >= 0.0
{
return u64::try_from(Duration::from_secs_f64(seconds).as_millis()).ok();
}
}
}
None
}
impl RetryPolicy {
pub fn with_max_attempts(mut self, n: usize) -> Self {
self.max_attempts = n;
self
}
pub fn with_initial_backoff_ms(mut self, ms: u64) -> Self {
self.initial_backoff_ms = ms;
self
}
pub fn with_max_backoff_ms(mut self, ms: u64) -> Self {
self.max_backoff_ms = ms;
self
}
pub fn with_multiplier(mut self, m: f64) -> Self {
self.multiplier = m;
self
}
pub fn with_jitter(mut self, jitter: bool) -> Self {
self.jitter = jitter;
self
}
pub fn with_backoff_sleep(mut self, sleep: bool) -> Self {
self.backoff_sleep = sleep;
self
}
pub async fn sleep_backoff(&self, attempt: usize) {
if !self.backoff_sleep {
return;
}
let backoff = self.backoff_for_attempt(attempt);
if backoff > Duration::ZERO {
tokio::time::sleep(backoff).await;
}
}
pub fn should_retry(&self, attempt: usize) -> bool {
attempt + 1 < self.max_attempts
}
pub fn should_retry_error(&self, attempt: usize, error: &TinyAgentsError) -> bool {
is_retryable(error) && self.should_retry(attempt)
}
pub fn max_attempts_capped_at(&self, max_retries_per_call: usize) -> usize {
self.max_attempts
.min(max_retries_per_call.saturating_add(1))
}
pub fn backoff_for_attempt(&self, attempt: usize) -> Duration {
self.backoff_for_attempt_with(attempt, 0.0)
}
pub fn backoff_for_attempt_with(&self, attempt: usize, rand01: f64) -> Duration {
let base = (self.initial_backoff_ms as f64) * self.multiplier.powi(attempt as i32);
let capped = base.min(self.max_backoff_ms as f64);
let effective = if self.jitter {
capped * rand01.clamp(0.0, 1.0)
} else {
capped
};
Duration::from_millis(effective as u64)
}
}
pub fn is_retryable(err: &TinyAgentsError) -> bool {
match err {
TinyAgentsError::Provider(provider_error) => provider_error.retryable,
TinyAgentsError::Model(_) | TinyAgentsError::Tool(_) => true,
_ => false,
}
}
impl FallbackPolicy {
pub fn new(models: impl IntoIterator<Item = impl Into<String>>) -> Self {
Self {
models: models.into_iter().map(Into::into).collect(),
}
}
pub fn next_after<'a>(&'a self, current: &str) -> Option<&'a str> {
let mut iter = self.models.iter();
while let Some(m) = iter.next() {
if m == current {
return iter.next().map(String::as_str);
}
}
None
}
}
impl RateLimiter {
pub fn new(capacity: u64, refill_per_sec: f64) -> Self {
let cap = capacity as f64;
Self {
inner: Mutex::new(types::RateLimiterState {
tokens: cap,
capacity: cap,
refill_per_sec,
last_refill: Instant::now(),
}),
}
}
pub fn try_acquire(&self, tokens: u64, now: Instant) -> bool {
let mut state = self.inner.lock().unwrap();
self.refill(&mut state, now);
if state.tokens >= tokens as f64 {
state.tokens -= tokens as f64;
true
} else {
false
}
}
pub fn available(&self, now: Instant) -> u64 {
let mut state = self.inner.lock().unwrap();
self.refill(&mut state, now);
state.tokens.floor() as u64
}
pub fn capacity(&self) -> u64 {
self.inner.lock().unwrap().capacity as u64
}
pub fn refill_per_sec(&self) -> f64 {
self.inner.lock().unwrap().refill_per_sec
}
pub fn can_ever_acquire(&self, tokens: u64) -> bool {
let state = self.inner.lock().unwrap();
tokens as f64 <= state.capacity && state.refill_per_sec > 0.0
}
fn refill(&self, state: &mut types::RateLimiterState, now: Instant) {
let elapsed = now.duration_since(state.last_refill).as_secs_f64();
if elapsed > 0.0 {
state.tokens = (state.tokens + elapsed * state.refill_per_sec).min(state.capacity);
state.last_refill = now;
}
}
}
#[cfg(test)]
mod test;