use std::time::{Duration, Instant};
use super::{
FallbackPolicy, ProviderFailureClass, RateLimiter, RetryPolicy, classify_provider_error,
classify_provider_failure, is_retryable, parse_retry_after_ms, structured_http_status,
};
use crate::error::TinyAgentsError;
#[test]
fn smoke_retry_policy_compiles() {
let policy = RetryPolicy::default();
assert!(policy.should_retry(0));
assert!(!policy.should_retry(3));
assert!(is_retryable(&TinyAgentsError::Model("timeout".into())));
assert!(!is_retryable(&TinyAgentsError::Validation(
"bad input".into()
)));
}
#[test]
fn backoff_grows_exponentially_then_caps() {
let policy = RetryPolicy::default();
assert_eq!(policy.backoff_for_attempt(0), Duration::from_millis(200));
assert_eq!(policy.backoff_for_attempt(1), Duration::from_millis(400));
assert_eq!(policy.backoff_for_attempt(2), Duration::from_millis(800));
assert_eq!(policy.backoff_for_attempt(3), Duration::from_millis(1_600));
let mut prev = Duration::ZERO;
for attempt in 0..20 {
let cur = policy.backoff_for_attempt(attempt);
assert!(cur >= prev, "backoff must be monotonic non-decreasing");
assert!(
cur <= Duration::from_millis(30_000),
"must never exceed cap"
);
prev = cur;
}
assert_eq!(
policy.backoff_for_attempt(50),
Duration::from_millis(30_000)
);
}
#[test]
fn backoff_jitter_scales_by_rand01() {
let policy = RetryPolicy::default().with_jitter(true);
assert_eq!(
policy.backoff_for_attempt_with(2, 0.0),
Duration::from_millis(0)
);
assert_eq!(
policy.backoff_for_attempt_with(2, 0.5),
Duration::from_millis(400)
);
assert_eq!(
policy.backoff_for_attempt_with(2, 5.0),
Duration::from_millis(800)
);
assert_eq!(
policy.backoff_for_attempt_with(2, -3.0),
Duration::from_millis(0)
);
}
#[test]
fn backoff_without_jitter_ignores_rand01() {
let policy = RetryPolicy::default(); assert_eq!(
policy.backoff_for_attempt_with(1, 0.99),
Duration::from_millis(400)
);
}
#[test]
fn should_retry_boundary_at_max_attempts() {
let policy = RetryPolicy::default().with_max_attempts(3);
assert!(policy.should_retry(0));
assert!(policy.should_retry(1));
assert!(!policy.should_retry(2)); assert!(!policy.should_retry(3));
let no_retry = RetryPolicy::default().with_max_attempts(1);
assert!(!no_retry.should_retry(0));
}
#[test]
fn max_attempts_capped_at_takes_the_stricter_of_the_two_caps() {
let policy = RetryPolicy::default().with_max_attempts(3);
assert_eq!(policy.max_attempts_capped_at(10), 3);
assert_eq!(policy.max_attempts_capped_at(1), 2);
assert_eq!(policy.max_attempts_capped_at(0), 1);
}
#[test]
fn is_retryable_classification() {
assert!(is_retryable(&TinyAgentsError::Model("5xx".into())));
assert!(is_retryable(&TinyAgentsError::Tool("transient".into())));
assert!(!is_retryable(&TinyAgentsError::Validation("bad".into())));
assert!(!is_retryable(&TinyAgentsError::RecursionLimit(10)));
let serde_err = serde_json::from_str::<i32>("not-json").unwrap_err();
assert!(!is_retryable(&TinyAgentsError::Serialization(serde_err)));
}
#[test]
fn provider_error_retryability_is_read_from_the_structured_flag_not_assumed() {
use crate::harness::model::ProviderError;
let rate_limited = ProviderError {
provider: "openai".to_string(),
status: Some(429),
retryable: true,
message: "rate limited".to_string(),
..ProviderError::default()
};
assert!(is_retryable(&TinyAgentsError::Provider(Box::new(
rate_limited
))));
let unauthorized = ProviderError {
provider: "openai".to_string(),
status: Some(401),
retryable: false,
message: "invalid api key".to_string(),
..ProviderError::default()
};
assert!(!is_retryable(&TinyAgentsError::Provider(Box::new(
unauthorized
))));
}
#[test]
fn structured_http_status_uses_only_anchored_positions() {
assert_eq!(
structured_http_status("custom_openai API error (403 Forbidden): nope"),
Some(403)
);
assert_eq!(structured_http_status("HTTP 404 Not Found"), Some(404));
assert_eq!(structured_http_status("status: 401"), Some(401));
assert_eq!(structured_http_status("408 Request Timeout"), Some(408));
assert_eq!(
structured_http_status("upstream took 450ms to respond, retrying"),
None
);
assert_eq!(
structured_http_status("gpt-4-0409 returned an empty completion"),
None
);
assert_eq!(
structured_http_status("received 412 partial bytes before reset"),
None
);
}
#[test]
fn provider_failure_classifies_generic_http_statuses() {
assert_eq!(
classify_provider_failure(Some(401), None, "invalid api key"),
ProviderFailureClass::NonRetryable
);
assert_eq!(
classify_provider_failure(Some(404), None, "model not found"),
ProviderFailureClass::NonRetryable
);
assert_eq!(
classify_provider_failure(Some(429), None, "too many requests"),
ProviderFailureClass::RateLimited
);
assert_eq!(
classify_provider_failure(Some(408), None, "request timeout"),
ProviderFailureClass::UpstreamUnhealthy
);
assert_eq!(
classify_provider_failure(Some(502), None, "bad gateway"),
ProviderFailureClass::UpstreamUnhealthy
);
}
#[test]
fn provider_failure_classifies_message_hints_without_status() {
assert_eq!(
classify_provider_failure(None, None, "authentication failed"),
ProviderFailureClass::NonRetryable
);
assert_eq!(
classify_provider_failure(None, None, "model glm-4.7 is unsupported"),
ProviderFailureClass::NonRetryable
);
assert_eq!(
classify_provider_failure(None, None, "no healthy upstream available"),
ProviderFailureClass::UpstreamUnhealthy
);
assert_eq!(
classify_provider_failure(None, None, "429 Too Many Requests: rate limit exceeded"),
ProviderFailureClass::RateLimited
);
}
#[test]
fn provider_failure_classifies_non_retryable_rate_limits() {
assert_eq!(
classify_provider_failure(
Some(429),
Some("1311"),
"the current account plan does not include glm-5"
),
ProviderFailureClass::NonRetryableRateLimit
);
assert_eq!(
classify_provider_failure(Some(429), None, "insufficient balance"),
ProviderFailureClass::NonRetryableRateLimit
);
}
#[test]
fn provider_failure_class_controls_retryability_and_reason_labels() {
assert!(ProviderFailureClass::RateLimited.is_retryable());
assert!(ProviderFailureClass::UpstreamUnhealthy.is_retryable());
assert!(!ProviderFailureClass::NonRetryable.is_retryable());
assert!(!ProviderFailureClass::NonRetryableRateLimit.is_retryable());
assert_eq!(ProviderFailureClass::Retryable.reason(), "retryable");
assert_eq!(ProviderFailureClass::NonRetryable.reason(), "non_retryable");
assert_eq!(ProviderFailureClass::RateLimited.reason(), "rate_limited");
assert_eq!(
ProviderFailureClass::NonRetryableRateLimit.reason(),
"rate_limited_non_retryable"
);
assert_eq!(
ProviderFailureClass::UpstreamUnhealthy.reason(),
"upstream_unhealthy"
);
}
#[test]
fn classify_provider_error_reads_structured_error_fields() {
use crate::harness::model::ProviderError;
let provider_error = ProviderError {
provider: "openai".to_string(),
model: Some("gpt-4o".to_string()),
status: Some(429),
code: Some("insufficient_quota".to_string()),
message: "insufficient quota".to_string(),
..ProviderError::default()
};
assert_eq!(
classify_provider_error(&provider_error),
ProviderFailureClass::NonRetryableRateLimit
);
}
#[test]
fn retry_after_parser_accepts_integer_float_and_space_separators() {
assert_eq!(
parse_retry_after_ms("429 Too Many Requests, Retry-After: 5"),
Some(5_000)
);
assert_eq!(
parse_retry_after_ms("Rate limited. retry_after: 2.5 seconds"),
Some(2_500)
);
assert_eq!(parse_retry_after_ms("Retry-After 7"), Some(7_000));
assert_eq!(parse_retry_after_ms("500 Internal Server Error"), None);
}
#[test]
fn fallback_next_after_semantics() {
let policy = FallbackPolicy::new(["a", "b", "c"]);
assert_eq!(policy.next_after("a"), Some("b"));
assert_eq!(policy.next_after("b"), Some("c"));
assert_eq!(policy.next_after("c"), None);
assert_eq!(policy.next_after("missing"), None);
let empty = FallbackPolicy::default();
assert_eq!(empty.next_after("a"), None);
}
#[test]
fn rate_limiter_acquire_until_empty() {
let limiter = RateLimiter::new(3, 1.0);
let now = Instant::now();
assert_eq!(limiter.available(now), 3);
assert!(limiter.try_acquire(1, now));
assert!(limiter.try_acquire(2, now));
assert_eq!(limiter.available(now), 0);
assert!(!limiter.try_acquire(1, now));
}
#[test]
fn rate_limiter_refills_over_time() {
let limiter = RateLimiter::new(10, 5.0); let start = Instant::now();
assert!(limiter.try_acquire(10, start));
assert_eq!(limiter.available(start), 0);
let after_1s = start + Duration::from_secs(1);
assert_eq!(limiter.available(after_1s), 5);
let limiter2 = RateLimiter::new(10, 5.0);
let s2 = Instant::now();
assert!(limiter2.try_acquire(10, s2));
let after_half = s2 + Duration::from_millis(500);
assert_eq!(limiter2.available(after_half), 2);
}
#[test]
fn rate_limiter_refill_caps_at_capacity() {
let limiter = RateLimiter::new(5, 100.0);
let start = Instant::now();
let later = start + Duration::from_secs(60);
assert_eq!(limiter.available(later), 5);
}
#[test]
fn backoff_sleep_defaults_off_and_is_opt_in() {
assert!(!RetryPolicy::default().backoff_sleep);
assert!(
RetryPolicy::default()
.with_backoff_sleep(true)
.backoff_sleep
);
}
#[tokio::test(start_paused = true)]
async fn sleep_backoff_waits_only_when_enabled() {
use tokio::time::Instant as TokioInstant;
let policy = RetryPolicy::default();
let t0 = TokioInstant::now();
policy.sleep_backoff(1).await;
assert_eq!(t0.elapsed(), Duration::ZERO);
let sleeping = RetryPolicy::default().with_backoff_sleep(true);
let expected = sleeping.backoff_for_attempt(1);
let t1 = TokioInstant::now();
sleeping.sleep_backoff(1).await;
assert_eq!(t1.elapsed(), expected);
assert!(expected > Duration::ZERO);
}
#[test]
fn should_retry_error_combines_classification_and_attempt_cap() {
let policy = RetryPolicy::default().with_max_attempts(3);
let retryable = TinyAgentsError::Model("5xx".into());
assert!(policy.should_retry_error(0, &retryable));
assert!(policy.should_retry_error(1, &retryable));
assert!(!policy.should_retry_error(2, &retryable));
let non_retryable = TinyAgentsError::Validation("bad".into());
assert!(!policy.should_retry_error(0, &non_retryable));
}