use std::pin::pin;
use std::sync::atomic::{AtomicBool, Ordering};
use crate::backends::stream_timeout::CANCEL_POLL_MS;
use crate::error_codes::{BACKEND_NETWORK, BACKEND_SEND, BACKEND_SERVER, BACKEND_TIMEOUT};
pub const MAX_STREAM_ATTEMPTS: u32 = 3;
pub const STREAM_RETRY_BACKOFF_MS: u32 = 300;
pub const SEND_MAX_ATTEMPTS: u32 = 2;
pub const SEND_RETRY_BACKOFF_MS: u32 = 500;
pub fn is_transient(code: u16) -> bool {
matches!(code, BACKEND_NETWORK | BACKEND_SERVER | BACKEND_TIMEOUT | BACKEND_SEND)
}
pub fn max_attempts(code: u16) -> u32 {
if code == BACKEND_SEND {
SEND_MAX_ATTEMPTS
} else if is_transient(code) {
MAX_STREAM_ATTEMPTS
} else {
1
}
}
pub fn should_retry(code: u16, attempt: u32) -> bool {
attempt < max_attempts(code)
}
pub fn backoff_ms(code: u16, attempt: u32) -> u32 {
if code == BACKEND_SEND {
SEND_RETRY_BACKOFF_MS
} else {
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(e.code(), attempt)).await;
}
Err(e) => return Err(e),
}
}
}
pub(crate) enum OpenOutcome<S> {
Opened(S),
Cancelled,
Failed(crate::error::Error),
}
async fn await_or_cancel<S, Fut>(fut: Fut, cancel: &AtomicBool) -> Option<crate::error::Result<S>>
where
Fut: core::future::Future<Output = crate::error::Result<S>>,
{
let mut fut = pin!(fut);
loop {
if cancel.load(Ordering::Acquire) {
return None;
}
let sleep = pin!(crate::runtime::sleep_ms(CANCEL_POLL_MS));
match futures_util::future::select(fut.as_mut(), sleep).await {
futures_util::future::Either::Left((res, _sleep)) => return Some(res),
futures_util::future::Either::Right((_elapsed, _fut)) => {}
}
}
}
async fn sleep_or_cancel(ms: u32, cancel: &AtomicBool) -> bool {
let mut waited = 0u32;
while waited < ms {
if cancel.load(Ordering::Acquire) {
return false;
}
let slice = CANCEL_POLL_MS.min(ms - waited).max(1);
crate::runtime::sleep_ms(slice).await;
waited = waited.saturating_add(slice);
}
!cancel.load(Ordering::Acquire)
}
pub(crate) async fn open_stream_with_retry_or_cancel<S, F, Fut>(
mut open: F,
cancel: &AtomicBool,
) -> OpenOutcome<S>
where
F: FnMut() -> Fut,
Fut: core::future::Future<Output = crate::error::Result<S>>,
{
let mut attempt = 0u32;
loop {
if cancel.load(Ordering::Acquire) {
return OpenOutcome::Cancelled;
}
attempt += 1;
match await_or_cancel(open(), cancel).await {
None => return OpenOutcome::Cancelled,
Some(Ok(s)) => return OpenOutcome::Opened(s),
Some(Err(e)) if should_retry(e.code(), attempt) => {
if !sleep_or_cancel(backoff_ms(e.code(), attempt), cancel).await {
return OpenOutcome::Cancelled;
}
}
Some(Err(e)) => return OpenOutcome::Failed(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, BACKEND_SEND] {
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(BACKEND_SERVER, 2), STREAM_RETRY_BACKOFF_MS * 2);
}
#[test]
fn send_class_retries_once_with_flat_backoff() {
assert_eq!(max_attempts(BACKEND_SEND), SEND_MAX_ATTEMPTS);
assert!(should_retry(BACKEND_SEND, 1));
assert!(!should_retry(BACKEND_SEND, SEND_MAX_ATTEMPTS));
assert_eq!(backoff_ms(BACKEND_SEND, 1), SEND_RETRY_BACKOFF_MS);
assert_eq!(max_attempts(BACKEND_NETWORK), MAX_STREAM_ATTEMPTS);
assert_eq!(max_attempts(BACKEND_AUTH), 1);
}
#[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_retries_bare_send_failure_once() {
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 < 2 {
Err(crate::error::Error::other("gemini POST: error sending request"))
} else {
Ok("stream")
}
}
})
.await
.expect("2nd attempt succeeds");
assert_eq!(out, "stream");
assert_eq!(calls.load(Ordering::SeqCst), 2);
let calls = AtomicU32::new(0);
let err = open_stream_with_retry(|| {
calls.fetch_add(1, Ordering::SeqCst);
async { Err::<(), _>(crate::error::Error::other("gemini POST: error sending request")) }
})
.await
.expect_err("still-dead network surfaces the error");
assert_eq!(err.code(), crate::error_codes::BACKEND_SEND);
assert_eq!(calls.load(Ordering::SeqCst), SEND_MAX_ATTEMPTS, "exactly one retry");
}
#[tokio::test]
async fn cancel_during_a_pending_open_returns_cancelled_promptly() {
use std::sync::atomic::AtomicU32;
use std::sync::Arc;
let cancel = Arc::new(AtomicBool::new(false));
let flip = cancel.clone();
tokio::spawn(async move {
crate::runtime::sleep_ms(30).await;
flip.store(true, Ordering::Release);
});
let attempts = AtomicU32::new(0);
let t0 = std::time::Instant::now();
let out = open_stream_with_retry_or_cancel(
|| {
attempts.fetch_add(1, Ordering::SeqCst);
async {
std::future::pending::<()>().await;
Ok::<_, crate::error::Error>("never")
}
},
&cancel,
)
.await;
assert!(matches!(out, OpenOutcome::Cancelled), "cancel must break a pending open");
assert!(
t0.elapsed() < std::time::Duration::from_secs(5),
"cancel latency is bounded by the poll slice"
);
assert_eq!(attempts.load(Ordering::SeqCst), 1, "the open is never re-attempted");
}
#[tokio::test]
async fn cancel_suppresses_the_transient_retry() {
use std::sync::atomic::AtomicU32;
use std::sync::Arc;
let cancel = Arc::new(AtomicBool::new(false));
let attempts = AtomicU32::new(0);
let out = open_stream_with_retry_or_cancel(
|| {
attempts.fetch_add(1, Ordering::SeqCst);
let cancel = cancel.clone();
async move {
cancel.store(true, Ordering::Release);
Err::<(), _>(crate::error::Error::other("HTTP 503 internal server error"))
}
},
&cancel,
)
.await;
assert!(matches!(out, OpenOutcome::Cancelled), "cancelled, not retried/failed");
assert_eq!(attempts.load(Ordering::SeqCst), 1, "attempt count stays 1 on cancel");
}
#[tokio::test]
async fn uncancelled_open_or_cancel_matches_the_plain_policy() {
use std::sync::atomic::AtomicU32;
let cancel = AtomicBool::new(false);
let attempts = AtomicU32::new(0);
let out = open_stream_with_retry_or_cancel(
|| {
let n = attempts.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")
}
}
},
&cancel,
)
.await;
match out {
OpenOutcome::Opened(s) => assert_eq!(s, "stream"),
_ => panic!("transient failures must still retry to success"),
}
assert_eq!(attempts.load(Ordering::SeqCst), MAX_STREAM_ATTEMPTS);
let attempts = AtomicU32::new(0);
let out = open_stream_with_retry_or_cancel(
|| {
attempts.fetch_add(1, Ordering::SeqCst);
async {
Err::<(), _>(crate::error::Error::other("HTTP 401 Unauthorized: bad API key"))
}
},
&cancel,
)
.await;
match out {
OpenOutcome::Failed(e) => assert_eq!(e.code(), BACKEND_AUTH),
_ => panic!("auth must fail fast, not cancel/retry"),
}
assert_eq!(attempts.load(Ordering::SeqCst), 1);
}
#[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);
}
}