use crate::error_codes::{BACKEND_NETWORK, BACKEND_SERVER, BACKEND_TIMEOUT};
pub const MAX_STREAM_ATTEMPTS: u32 = 3;
pub const STREAM_RETRY_BACKOFF_MS: u32 = 300;
pub fn is_transient(code: u16) -> bool {
matches!(code, BACKEND_NETWORK | BACKEND_SERVER | BACKEND_TIMEOUT)
}
pub fn should_retry(code: u16, attempt: u32) -> bool {
is_transient(code) && attempt < MAX_STREAM_ATTEMPTS
}
pub fn backoff_ms(attempt: u32) -> u32 {
STREAM_RETRY_BACKOFF_MS * attempt
}
pub(crate) async fn open_stream_with_retry<S, F, Fut>(mut open: F) -> crate::error::Result<S>
where
F: FnMut() -> Fut,
Fut: core::future::Future<Output = crate::error::Result<S>>,
{
let mut attempt = 0u32;
loop {
attempt += 1;
match open().await {
Ok(s) => return Ok(s),
Err(e) if should_retry(e.code(), attempt) => {
crate::runtime::sleep_ms(backoff_ms(attempt)).await;
}
Err(e) => return Err(e),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error_codes::{BACKEND_AUTH, BACKEND_CREDITS, BACKEND_RATE_LIMIT};
#[test]
fn only_transient_classes_retry() {
for c in [BACKEND_NETWORK, BACKEND_SERVER, BACKEND_TIMEOUT] {
assert!(is_transient(c), "code {c} should be transient");
}
for c in [BACKEND_AUTH, BACKEND_CREDITS, BACKEND_RATE_LIMIT, 0] {
assert!(!is_transient(c), "code {c} must NOT retry");
}
}
#[test]
fn should_retry_stops_at_the_attempt_cap() {
assert!(should_retry(BACKEND_SERVER, 1));
assert!(should_retry(BACKEND_SERVER, MAX_STREAM_ATTEMPTS - 1));
assert!(!should_retry(BACKEND_SERVER, MAX_STREAM_ATTEMPTS)); assert!(!should_retry(BACKEND_RATE_LIMIT, 1)); assert_eq!(backoff_ms(2), STREAM_RETRY_BACKOFF_MS * 2);
}
#[tokio::test]
async fn open_stream_with_retry_retries_transient_then_succeeds() {
use std::sync::atomic::{AtomicU32, Ordering};
let calls = AtomicU32::new(0);
let out = open_stream_with_retry(|| {
let n = calls.fetch_add(1, Ordering::SeqCst) + 1;
async move {
if n < MAX_STREAM_ATTEMPTS {
Err(crate::error::Error::other("HTTP 503 internal server error"))
} else {
Ok("stream")
}
}
})
.await
.expect("succeeds within the attempt cap");
assert_eq!(out, "stream");
assert_eq!(calls.load(Ordering::SeqCst), MAX_STREAM_ATTEMPTS);
}
#[tokio::test]
async fn open_stream_with_retry_fails_fast_on_non_transient() {
use std::sync::atomic::{AtomicU32, Ordering};
let calls = AtomicU32::new(0);
let err = open_stream_with_retry(|| {
calls.fetch_add(1, Ordering::SeqCst);
async { Err::<(), _>(crate::error::Error::other("HTTP 401 Unauthorized: bad API key")) }
})
.await
.expect_err("auth must not retry");
assert_eq!(err.code(), BACKEND_AUTH);
assert_eq!(calls.load(Ordering::SeqCst), 1, "exactly one attempt");
}
#[tokio::test]
async fn open_stream_with_retry_gives_up_at_the_cap() {
use std::sync::atomic::{AtomicU32, Ordering};
let calls = AtomicU32::new(0);
let err = open_stream_with_retry(|| {
calls.fetch_add(1, Ordering::SeqCst);
async { Err::<(), _>(crate::error::Error::other("HTTP 503 internal server error")) }
})
.await
.expect_err("all attempts failed");
assert_eq!(err.code(), BACKEND_SERVER);
assert_eq!(calls.load(Ordering::SeqCst), MAX_STREAM_ATTEMPTS);
}
}