use std::time::{Duration, SystemTime, UNIX_EPOCH};
use async_trait::async_trait;
use polyc_llm::{Chunk, CompletionRequest, LlmError, LlmErrorKind, LlmProvider};
#[async_trait]
pub trait Clock: std::fmt::Debug + Send + Sync {
fn now(&self) -> SystemTime;
fn jitter_frac(&self) -> f64 {
let nanos = self
.now()
.duration_since(UNIX_EPOCH)
.map_or(0, |d| d.subsec_nanos());
f64::from(nanos % 1_000_000) / 1_000_000.0
}
async fn sleep(&self, dur: Duration);
}
#[derive(Debug, Default, Clone, Copy)]
pub struct RealClock;
#[async_trait]
impl Clock for RealClock {
fn now(&self) -> SystemTime {
SystemTime::now()
}
async fn sleep(&self, dur: Duration) {
tokio::time::sleep(dur).await;
}
}
#[derive(Debug, Clone, Copy)]
pub struct RetryConfig {
pub max_retries: u32,
pub base_delay: Duration,
pub max_delay: Duration,
}
impl Default for RetryConfig {
fn default() -> Self {
Self {
max_retries: 4,
base_delay: Duration::from_millis(500),
max_delay: Duration::from_secs(30),
}
}
}
impl RetryConfig {
#[must_use]
pub fn from_env() -> Self {
let d = Self::default();
Self {
max_retries: env_parse("POLYCHROME_LLM_MAX_RETRIES").unwrap_or(d.max_retries),
base_delay: env_parse("POLYCHROME_LLM_RETRY_BASE_MS")
.map_or(d.base_delay, Duration::from_millis),
max_delay: env_parse("POLYCHROME_LLM_RETRY_MAX_MS")
.map_or(d.max_delay, Duration::from_millis),
}
}
}
pub(crate) fn env_parse<T: std::str::FromStr>(key: &str) -> Option<T> {
std::env::var(key).ok()?.parse().ok()
}
const fn is_retryable(kind: LlmErrorKind) -> bool {
matches!(
kind,
LlmErrorKind::RateLimit | LlmErrorKind::Timeout | LlmErrorKind::Unavailable
)
}
#[must_use]
pub fn backoff_delay(attempt: u32, base: Duration, cap: Duration, jitter_frac: f64) -> Duration {
let factor = 1u32.checked_shl(attempt.min(16)).unwrap_or(u32::MAX);
let exp = base.saturating_mul(factor).min(cap);
let scale = 0.5_f64.mul_add(jitter_frac.clamp(0.0, 1.0), 0.5);
exp.mul_f64(scale)
}
pub async fn complete_with_retry<P>(
provider: &P,
req: CompletionRequest,
cfg: &RetryConfig,
clock: &dyn Clock,
) -> Result<futures::stream::BoxStream<'static, Result<Chunk, P::Error>>, P::Error>
where
P: LlmProvider + ?Sized,
{
let mut attempt = 0u32;
loop {
match provider.complete(req.clone()).await {
Ok(stream) => return Ok(stream),
Err(err) => {
let kind = err.kind();
if !is_retryable(kind) || attempt >= cfg.max_retries {
return Err(err);
}
let delay = err.retry_after().filter(|d| !d.is_zero()).map_or_else(
|| backoff_delay(attempt, cfg.base_delay, cfg.max_delay, clock.jitter_frac()),
|d| d.min(cfg.max_delay),
);
attempt += 1;
tracing::warn!(
attempt,
?kind,
delay_ms = u64::try_from(delay.as_millis()).unwrap_or(u64::MAX),
"model call failed; retrying"
);
clock.sleep(delay).await;
}
}
}
}
#[cfg(test)]
#[allow(clippy::pedantic, clippy::nursery, missing_docs)]
mod tests {
use std::sync::atomic::{AtomicUsize, Ordering};
use async_trait::async_trait;
use futures::StreamExt as _;
use polyc_llm::error::DummyError;
use polyc_llm::{Chunk, CompletionRequest, LlmProvider, StopReason};
use super::*;
struct FlakyProvider {
calls: AtomicUsize,
fail_n: usize,
err: fn() -> DummyError,
}
#[async_trait]
impl LlmProvider for FlakyProvider {
type Error = DummyError;
async fn complete(
&self,
_req: CompletionRequest,
) -> Result<futures::stream::BoxStream<'static, Result<Chunk, Self::Error>>, Self::Error>
{
let n = self.calls.fetch_add(1, Ordering::SeqCst);
if n < self.fail_n {
return Err((self.err)());
}
Ok(futures::stream::iter(vec![Ok(Chunk::Stop(StopReason::EndTurn))]).boxed())
}
}
fn fast_cfg(max_retries: u32) -> RetryConfig {
RetryConfig {
max_retries,
base_delay: Duration::from_millis(0),
max_delay: Duration::from_millis(0),
}
}
fn unavailable() -> DummyError {
DummyError::Transport("reset".to_owned())
}
fn bad_request() -> DummyError {
DummyError::Provider {
status: 400,
body: "nope".to_owned(),
}
}
fn ambiguous() -> DummyError {
DummyError::Ambiguous("outcome unknown".to_owned())
}
#[tokio::test]
async fn an_ambiguous_attempt_is_never_retried() {
let provider = FlakyProvider {
calls: AtomicUsize::new(0),
fail_n: 1,
err: ambiguous,
};
let error = complete_with_retry(
&provider,
CompletionRequest::new("model"),
&fast_cfg(5),
&RealClock,
)
.await;
let Err(error) = error else {
panic!("an ambiguous attempt is terminal, not a stream");
};
assert_eq!(error.kind(), polyc_llm::LlmErrorKind::Ambiguous);
assert_eq!(
provider.calls.load(Ordering::SeqCst),
1,
"an ambiguous attempt is tried once and never again"
);
}
#[tokio::test]
async fn retries_then_succeeds() {
let p = FlakyProvider {
calls: AtomicUsize::new(0),
fail_n: 2,
err: unavailable,
};
let out =
complete_with_retry(&p, CompletionRequest::new("m"), &fast_cfg(4), &RealClock).await;
assert!(out.is_ok(), "should succeed after 2 retries");
assert_eq!(p.calls.load(Ordering::SeqCst), 3, "2 failures + 1 success");
}
#[tokio::test]
async fn gives_up_after_budget() {
let p = FlakyProvider {
calls: AtomicUsize::new(0),
fail_n: 99,
err: unavailable,
};
let out =
complete_with_retry(&p, CompletionRequest::new("m"), &fast_cfg(3), &RealClock).await;
assert!(out.is_err(), "exhausts the budget");
assert_eq!(p.calls.load(Ordering::SeqCst), 4);
}
#[tokio::test]
async fn terminal_error_is_not_retried() {
let p = FlakyProvider {
calls: AtomicUsize::new(0),
fail_n: 99,
err: bad_request,
};
let out =
complete_with_retry(&p, CompletionRequest::new("m"), &fast_cfg(4), &RealClock).await;
assert!(out.is_err());
assert_eq!(p.calls.load(Ordering::SeqCst), 1, "bad-request is terminal");
}
#[test]
fn backoff_grows_and_caps() {
let base = Duration::from_millis(100);
let cap = Duration::from_millis(1000);
assert_eq!(backoff_delay(0, base, cap, 1.0), Duration::from_millis(100));
assert_eq!(backoff_delay(1, base, cap, 1.0), Duration::from_millis(200));
assert_eq!(backoff_delay(2, base, cap, 1.0), Duration::from_millis(400));
assert_eq!(
backoff_delay(4, base, cap, 1.0),
Duration::from_millis(1000)
);
assert_eq!(backoff_delay(1, base, cap, 0.0), Duration::from_millis(100));
}
}