use crate::breaker::ProviderBreaker;
use crate::error::{Result, ShimError};
use crate::provider::{Provider, ProviderRequest};
use bytes::Bytes;
use chrono::{DateTime, Utc};
use eventsource_stream::Eventsource;
use futures::{Stream, StreamExt};
use reqwest::header::HeaderMap;
use reqwest::Client;
use std::pin::Pin;
use std::sync::Arc;
use std::time::Duration;
#[derive(Clone, Copy, Debug)]
struct RetryConfig {
max_retries: u32,
base: Duration,
cap: Duration,
}
impl Default for RetryConfig {
fn default() -> Self {
Self {
max_retries: 3,
base: Duration::from_secs(1),
cap: Duration::from_secs(60),
}
}
}
impl RetryConfig {
fn from_env() -> Self {
let d = Self::default();
let max_retries = env_parse("LLMSHIM_MAX_RETRIES").unwrap_or(d.max_retries);
let cap_secs = env_parse::<u64>("LLMSHIM_MAX_BACKOFF_SECS").unwrap_or(d.cap.as_secs());
Self {
max_retries,
base: d.base,
cap: Duration::from_secs(cap_secs),
}
}
}
fn env_parse<T: std::str::FromStr>(key: &str) -> Option<T> {
std::env::var(key).ok()?.trim().parse().ok()
}
#[derive(Clone)]
pub struct ShimClient {
http: Client,
retry: RetryConfig,
breaker: Option<Arc<ProviderBreaker>>,
}
impl Default for ShimClient {
fn default() -> Self {
Self::new()
}
}
impl ShimClient {
pub fn new() -> Self {
Self {
http: Client::builder()
.redirect(reqwest::redirect::Policy::none())
.pool_idle_timeout(Duration::from_secs(90))
.pool_max_idle_per_host(4)
.tcp_keepalive(Duration::from_secs(30))
.tcp_nodelay(true)
.build()
.expect("failed to build HTTP client"),
retry: RetryConfig::from_env(),
breaker: None,
}
}
pub fn with_breaker(mut self, breaker: Arc<ProviderBreaker>) -> Self {
self.breaker = Some(breaker);
self
}
async fn observe(&self, provider: &dyn Provider, outcome: std::result::Result<(), &ShimError>) {
if let Some(breaker) = &self.breaker {
breaker.observe(provider.name(), outcome).await;
}
}
pub async fn warmup(&self, urls: &[&str]) {
let futs: Vec<_> = urls
.iter()
.map(|url| {
let client = self.http.clone();
let url = url.to_string();
tokio::spawn(async move {
let _ = client
.head(&url)
.timeout(Duration::from_secs(5))
.send()
.await;
})
})
.collect();
for f in futs {
let _ = f.await;
}
}
const RETRYABLE_STATUSES: &'static [u16] = &[429, 500, 502, 503, 504, 529];
pub async fn send(&self, req: &ProviderRequest) -> Result<reqwest::Response> {
let max_retries = self.retry.max_retries;
for attempt in 0..=max_retries {
let mut builder = self.http.post(&req.url);
for (k, v) in &req.headers {
builder = builder.header(k, v);
}
builder = builder.json(&req.body);
match builder.send().await {
Ok(resp) => {
let status = resp.status();
if status.is_success() {
return Ok(resp);
}
let status_code = status.as_u16();
if Self::RETRYABLE_STATUSES.contains(&status_code) && attempt < max_retries {
let wait =
retry_after_wait(resp.headers(), self.retry.cap).unwrap_or_else(|| {
backoff_with_jitter(attempt, self.retry.base, self.retry.cap)
});
let _ = resp.text().await;
tokio::time::sleep(wait).await;
continue;
}
let body = resp.text().await.unwrap_or_default();
return Err(ShimError::ProviderError {
status: status_code,
body,
});
}
Err(e) if Self::is_retryable_transport(&e) && attempt < max_retries => {
tokio::time::sleep(backoff_with_jitter(
attempt,
self.retry.base,
self.retry.cap,
))
.await;
continue;
}
Err(e) => return Err(ShimError::Http(e)),
}
}
unreachable!()
}
fn is_retryable_transport(err: &reqwest::Error) -> bool {
err.is_connect() || err.is_timeout() || err.is_request() || err.is_body()
}
pub async fn completion(
&self,
provider: &dyn Provider,
model: &str,
request: &serde_json::Value,
) -> Result<serde_json::Value> {
let result = self.completion_unobserved(provider, model, request).await;
self.observe(provider, result.as_ref().map(|_| ())).await;
result
}
async fn completion_unobserved(
&self,
provider: &dyn Provider,
model: &str,
request: &serde_json::Value,
) -> Result<serde_json::Value> {
let plan = crate::shim::Plan::new(
provider.name(),
model,
provider.replay_target(model).wire,
request,
)?;
let mut rendered = plan.render()?;
rendered["stream"] = serde_json::json!(false);
let mut usage = serde_json::json!({});
for attempt in 0..2 {
let (mut result, target) = self
.completion_once(provider, model, &rendered)
.await
.map_err(|e| plan.dispatch_error(e))?;
crate::shim::add_usage(&mut usage, &result);
match plan.finish(&mut result, &target) {
Ok(()) => {
if attempt > 0 {
result["usage"] = usage;
}
crate::cost::stamp(provider.name(), model, &mut result);
return Ok(result);
}
Err(feedback) if attempt == 0 && plan.can_repair(&result) => {
rendered = plan.repair(&feedback)?;
rendered["stream"] = serde_json::json!(false);
}
Err(_) => return Err(crate::shim::failed()),
}
}
unreachable!()
}
async fn completion_once(
&self,
provider: &dyn Provider,
model: &str,
request: &serde_json::Value,
) -> Result<(serde_json::Value, crate::reasoning::ReplayTarget)> {
let provider_req = provider.prepare_request(model, request).await?;
let target = provider.request_replay_target(model, &provider_req);
let resp = self.send(&provider_req).await?;
if provider.name() == "chatgpt"
&& target.wire == crate::reasoning::WireFormat::OpenAiResponses
{
let mut result = crate::providers::chatgpt::collect_response(model, resp).await?;
crate::reasoning::bind_response_context(&mut result, &target);
crate::toolcall::bind_response_context(&mut result, &target);
return Ok((result, target));
}
let body: serde_json::Value = resp.json().await?;
let mut result = provider.transform_response(model, body)?;
crate::reasoning::bind_response_context(&mut result, &target);
crate::toolcall::bind_response_context(&mut result, &target);
Ok((result, target))
}
pub async fn stream(
&self,
provider: &dyn Provider,
model: &str,
request: &serde_json::Value,
) -> Result<Pin<Box<dyn Stream<Item = Result<String>> + Send>>> {
let opened = self.stream_unobserved(provider, model, request).await;
self.observe(provider, opened.as_ref().map(|_| ())).await;
opened
}
async fn stream_unobserved(
&self,
provider: &dyn Provider,
model: &str,
request: &serde_json::Value,
) -> Result<Pin<Box<dyn Stream<Item = Result<String>> + Send>>> {
let plan = crate::shim::Plan::new(
provider.name(),
model,
provider.replay_target(model).wire,
request,
)?;
let mut rendered = plan.render()?;
if !plan.buffered() {
return self
.stream_once(provider, model, &rendered)
.await
.map(|(stream, _)| stream);
}
let mut usage = serde_json::json!({});
for attempt in 0..2 {
let (stream, target) = self
.stream_once(provider, model, &rendered)
.await
.map_err(|e| plan.dispatch_error(e))?;
let mut result = crate::shim::collect(stream)
.await
.map_err(|e| plan.dispatch_error(e))?;
crate::shim::add_usage(&mut usage, &result);
match plan.finish(&mut result, &target) {
Ok(()) => {
if attempt > 0 {
result["usage"] = usage;
}
crate::cost::stamp(provider.name(), model, &mut result);
return Ok(Box::pin(futures::stream::iter(crate::shim::chunks(result))));
}
Err(feedback) if attempt == 0 && plan.can_repair(&result) => {
rendered = plan.repair(&feedback)?
}
Err(_) => return Err(crate::shim::failed()),
}
}
unreachable!()
}
pub async fn stream_owned(
&self,
provider: Arc<dyn Provider>,
model: &str,
request: &serde_json::Value,
) -> Result<Pin<Box<dyn Stream<Item = Result<String>> + Send>>> {
let opened = self
.stream_owned_unobserved(provider.clone(), model, request)
.await;
self.observe(provider.as_ref(), opened.as_ref().map(|_| ()))
.await;
opened
}
async fn stream_owned_unobserved(
&self,
provider: Arc<dyn Provider>,
model: &str,
request: &serde_json::Value,
) -> Result<Pin<Box<dyn Stream<Item = Result<String>> + Send>>> {
let plan = crate::shim::Plan::new(
provider.name(),
model,
provider.replay_target(model).wire,
request,
)?;
let rendered = plan.render()?;
let (first, target) = self
.stream_once(provider.as_ref(), model, &rendered)
.await
.map_err(|e| plan.dispatch_error(e))?;
if !plan.buffered() {
return Ok(first);
}
let client = self.clone();
let model = model.to_owned();
Ok(Box::pin(futures::stream::once(async move {
let mut result = crate::shim::collect(first)
.await
.map_err(|e| plan.dispatch_error(e))?;
let mut usage = serde_json::json!({});
crate::shim::add_usage(&mut usage, &result);
if let Err(feedback) = plan.finish(&mut result, &target) {
if !plan.can_repair(&result) {
return Err(crate::shim::failed());
}
let rendered = plan.repair(&feedback)?;
let (second, target) = client
.stream_once(provider.as_ref(), &model, &rendered)
.await
.map_err(|e| plan.dispatch_error(e))?;
result = crate::shim::collect(second)
.await
.map_err(|e| plan.dispatch_error(e))?;
crate::shim::add_usage(&mut usage, &result);
plan.finish(&mut result, &target)
.map_err(|_| crate::shim::failed())?;
result["usage"] = usage;
}
crate::cost::stamp(provider.name(), &model, &mut result);
crate::shim::chunks(result).pop().unwrap()
})))
}
async fn stream_once(
&self,
provider: &dyn Provider,
model: &str,
request: &serde_json::Value,
) -> Result<(
Pin<Box<dyn Stream<Item = Result<String>> + Send>>,
crate::reasoning::ReplayTarget,
)> {
let mut req_value = request.clone();
req_value["stream"] = serde_json::Value::Bool(true);
let provider_req = provider.prepare_request(model, &req_value).await?;
let target = provider.request_replay_target(model, &provider_req);
let resp = self.send(&provider_req).await?;
let events = native_events(resp.bytes_stream());
let sse = SseStream {
inner: events,
normalizer: crate::streaming::StreamNormalizer::new(target.clone()),
};
let (name, model) = (provider.name().to_owned(), model.to_owned());
let priced =
sse.map(move |item| item.map(|chunk| crate::cost::stamp_chunk(&name, &model, chunk)));
Ok((Box::pin(priced), target))
}
}
fn retry_after_wait(headers: &HeaderMap, cap: Duration) -> Option<Duration> {
let base = parse_retry_after(headers).or_else(|| parse_provider_reset(headers))?;
let capped = base.min(cap);
Some(capped + small_jitter())
}
fn parse_retry_after(headers: &HeaderMap) -> Option<Duration> {
parse_retry_after_at(headers, Utc::now())
}
fn parse_retry_after_at(headers: &HeaderMap, now: DateTime<Utc>) -> Option<Duration> {
let raw = header_str(headers, "retry-after")?.trim();
if let Ok(secs) = raw.parse::<u64>() {
return Some(Duration::from_secs(secs));
}
let when = DateTime::parse_from_rfc2822(raw).ok()?.with_timezone(&Utc);
duration_until(when, now)
}
fn parse_provider_reset(headers: &HeaderMap) -> Option<Duration> {
parse_provider_reset_at(headers, Utc::now())
}
fn parse_provider_reset_at(headers: &HeaderMap, now: DateTime<Utc>) -> Option<Duration> {
let mut best: Option<Duration> = None;
let mut consider = |d: Option<Duration>| {
if let Some(d) = d {
best = Some(best.map_or(d, |b| b.max(d)));
}
};
for name in ["x-ratelimit-reset-tokens", "x-ratelimit-reset-requests"] {
if let Some(v) = header_str(headers, name) {
consider(parse_go_duration(v));
}
}
for (name, value) in headers.iter() {
let name = name.as_str();
if name.starts_with("anthropic-ratelimit-") && name.ends_with("-reset") {
if let Ok(v) = value.to_str() {
if let Ok(when) = DateTime::parse_from_rfc3339(v.trim()) {
consider(duration_until(when.with_timezone(&Utc), now));
}
}
}
}
best
}
fn parse_go_duration(s: &str) -> Option<Duration> {
let s = s.trim();
if s.is_empty() {
return None;
}
let bytes = s.as_bytes();
let mut i = 0;
let mut total = Duration::ZERO;
let mut saw_unit = false;
while i < bytes.len() {
let num_start = i;
while i < bytes.len() && (bytes[i].is_ascii_digit() || bytes[i] == b'.') {
i += 1;
}
if i == num_start {
return None; }
let value: f64 = s[num_start..i].parse().ok()?;
let unit_start = i;
while i < bytes.len() && !(bytes[i].is_ascii_digit() || bytes[i] == b'.') {
i += 1;
}
let unit = &s[unit_start..i];
let secs = match unit {
"h" => value * 3600.0,
"m" => value * 60.0,
"s" => value,
"ms" => value / 1_000.0,
"us" | "µs" | "μs" => value / 1_000_000.0,
"ns" => value / 1_000_000_000.0,
_ => return None,
};
total += Duration::from_secs_f64(secs);
saw_unit = true;
}
saw_unit.then_some(total)
}
fn duration_until(when: DateTime<Utc>, now: DateTime<Utc>) -> Option<Duration> {
(when - now).to_std().ok()
}
fn header_str<'a>(headers: &'a HeaderMap, name: &str) -> Option<&'a str> {
headers.get(name)?.to_str().ok()
}
fn backoff_bound(attempt: u32, base: Duration, cap: Duration) -> Duration {
let mult = 1u64.checked_shl(attempt).unwrap_or(u64::MAX);
let ms = (base.as_millis() as u64).saturating_mul(mult);
Duration::from_millis(ms).min(cap)
}
fn full_jitter(bound: Duration, rand: u64) -> Duration {
let ms = bound.as_millis() as u64;
if ms == 0 {
return Duration::ZERO;
}
Duration::from_millis(rand % (ms + 1))
}
fn backoff_with_jitter(attempt: u32, base: Duration, cap: Duration) -> Duration {
full_jitter(backoff_bound(attempt, base, cap), rand_u64())
}
fn small_jitter() -> Duration {
Duration::from_millis(rand_u64() % 251)
}
fn rand_u64() -> u64 {
use std::time::{SystemTime, UNIX_EPOCH};
let seed = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_nanos() as u64)
.unwrap_or(0);
let mut x = seed.wrapping_add(0x9E37_79B9_7F4A_7C15);
x = (x ^ (x >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
x = (x ^ (x >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
x ^ (x >> 31)
}
fn native_events(
stream: impl Stream<Item = std::result::Result<Bytes, reqwest::Error>> + Send + 'static,
) -> Pin<Box<dyn Stream<Item = Result<String>> + Send>> {
Box::pin(stream.eventsource().map(|event| {
event
.map(|e| e.data)
.map_err(|_| ShimError::Stream("could not read upstream SSE".into()))
}))
}
struct SseStream {
inner: Pin<Box<dyn Stream<Item = Result<String>> + Send>>,
normalizer: crate::streaming::StreamNormalizer,
}
impl Stream for SseStream {
type Item = Result<String>;
fn poll_next(
mut self: Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
use std::task::Poll;
loop {
if self.normalizer.is_finished() {
return Poll::Ready(None);
}
let data = match self.inner.as_mut().poll_next(cx) {
Poll::Ready(Some(Ok(data))) => data,
Poll::Ready(Some(Err(error))) => {
self.normalizer.abort();
return Poll::Ready(Some(Err(error)));
}
Poll::Ready(None) => return Poll::Ready(self.normalizer.finish().transpose()),
Poll::Pending => return Poll::Pending,
};
match self.normalizer.push(&data) {
Ok(Some(chunk)) => return Poll::Ready(Some(Ok(chunk))),
Ok(None) => continue,
Err(error) => {
self.normalizer.abort();
return Poll::Ready(Some(Err(error)));
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::TimeZone;
use reqwest::header::{HeaderMap, HeaderValue};
fn fragmented_sse(
provider: &dyn Provider,
model: &str,
wire: String,
keep_open: bool,
) -> SseStream {
let bytes: Vec<_> = wire
.as_bytes()
.iter()
.map(|b| Ok(Bytes::copy_from_slice(&[*b])))
.collect();
let tail = if keep_open {
Box::pin(futures::stream::pending())
as Pin<Box<dyn Stream<Item = std::result::Result<Bytes, reqwest::Error>> + Send>>
} else {
Box::pin(futures::stream::empty())
};
SseStream {
inner: native_events(futures::stream::iter(bytes).chain(tail)),
normalizer: provider.stream_normalizer(model),
}
}
#[tokio::test]
async fn signed_reasoning_survives_utf8_byte_splits_multiline_sse_and_crlf() {
let p = crate::providers::anthropic::Anthropic::new("key".into());
let events = [
serde_json::json!({"type":"content_block_delta","index":0,"delta":{"type":"thinking_delta","thinking":"é雪🙂"}}),
serde_json::json!({"type":"content_block_delta","index":0,"delta":{"type":"signature_delta","signature":"opaque+/="}}),
serde_json::json!({"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":3}}),
];
let mut wire = String::from(": keepalive\r\n\r\n");
for e in events {
let json = e.to_string();
let (first, rest) = json.split_once(',').unwrap();
wire.push_str(&format!("data: {first},\r\ndata: {rest}\r\n\r\n"));
}
let chunks: Vec<_> = fragmented_sse(&p, "claude-sonnet-4-6", wire, false)
.collect()
.await;
let mut acc = crate::reasoning::ReasoningAccumulator::default();
for c in chunks {
acc.push(
&serde_json::from_str::<serde_json::Value>(&c.unwrap()).unwrap()["choices"][0]
["delta"],
);
}
assert_eq!(acc.blocks()[0]["text"], "é雪🙂");
assert_eq!(acc.blocks()[0]["signature"], "opaque+/=");
}
#[tokio::test]
async fn done_closes_chat_stream_after_late_usage_without_waiting_for_http_eof() {
let p = crate::providers::openai_compat::OpenAiCompatible::new("custom", "", None);
let wire="data: {\"choices\":[{\"delta\":{\"content\":\"ok\"},\"finish_reason\":\"stop\"}]}\n\ndata: {\"choices\":[],\"usage\":{\"prompt_cache_hit_tokens\":7}}\n\ndata: [DONE]\n\n";
let chunks: Vec<_> = tokio::time::timeout(
Duration::from_secs(1),
fragmented_sse(&p, "model", wire.into(), true).collect(),
)
.await
.unwrap();
let last: serde_json::Value =
serde_json::from_str(chunks.last().unwrap().as_ref().unwrap()).unwrap();
assert_eq!(last["usage"]["cache_read_tokens"], 7);
assert_eq!(last["choices"][0]["finish_reason"], "stop");
}
#[tokio::test]
async fn done_without_a_terminal_chunk_is_an_error() {
let p = crate::providers::openai_compat::OpenAiCompatible::new("custom", "", None);
let chunks: Vec<_> = fragmented_sse(&p, "model", "data: [DONE]\n\n".into(), true)
.collect()
.await;
assert_eq!(chunks.len(), 1);
assert!(matches!(chunks[0], Err(ShimError::Stream(_))));
}
fn headers(pairs: &[(&'static str, &str)]) -> HeaderMap {
let mut h = HeaderMap::new();
for (k, v) in pairs {
h.insert(*k, HeaderValue::from_str(v).unwrap());
}
h
}
#[test]
fn retry_after_integer_seconds() {
let h = headers(&[("retry-after", "5")]);
let now = Utc.with_ymd_and_hms(2026, 1, 1, 0, 0, 0).unwrap();
assert_eq!(parse_retry_after_at(&h, now), Some(Duration::from_secs(5)));
}
#[test]
fn retry_after_zero_seconds() {
let h = headers(&[("retry-after", "0")]);
let now = Utc.with_ymd_and_hms(2026, 1, 1, 0, 0, 0).unwrap();
assert_eq!(parse_retry_after_at(&h, now), Some(Duration::ZERO));
}
#[test]
fn retry_after_http_date_future() {
let now = Utc.with_ymd_and_hms(2015, 10, 21, 7, 28, 0).unwrap();
let h = headers(&[("retry-after", "Wed, 21 Oct 2015 07:28:30 GMT")]);
assert_eq!(parse_retry_after_at(&h, now), Some(Duration::from_secs(30)));
}
#[test]
fn retry_after_http_date_in_past_is_none() {
let now = Utc.with_ymd_and_hms(2015, 10, 21, 7, 29, 0).unwrap();
let h = headers(&[("retry-after", "Wed, 21 Oct 2015 07:28:00 GMT")]);
assert_eq!(parse_retry_after_at(&h, now), None);
}
#[test]
fn retry_after_absent_or_garbage_is_none() {
let now = Utc::now();
assert_eq!(parse_retry_after_at(&HeaderMap::new(), now), None);
let h = headers(&[("retry-after", "soon-ish")]);
assert_eq!(parse_retry_after_at(&h, now), None);
}
#[test]
fn openai_reset_go_duration_takes_max() {
let now = Utc::now();
let h = headers(&[
("x-ratelimit-reset-requests", "1s"),
("x-ratelimit-reset-tokens", "6m0s"),
]);
assert_eq!(
parse_provider_reset_at(&h, now),
Some(Duration::from_secs(360))
);
}
#[test]
fn openai_reset_millis() {
let now = Utc::now();
let h = headers(&[("x-ratelimit-reset-tokens", "100ms")]);
assert_eq!(
parse_provider_reset_at(&h, now),
Some(Duration::from_millis(100))
);
}
#[test]
fn anthropic_reset_rfc3339() {
let now = Utc.with_ymd_and_hms(2026, 1, 1, 0, 0, 0).unwrap();
let h = headers(&[("anthropic-ratelimit-requests-reset", "2026-01-01T00:00:10Z")]);
assert_eq!(
parse_provider_reset_at(&h, now),
Some(Duration::from_secs(10))
);
}
#[test]
fn anthropic_reset_takes_max_across_resources() {
let now = Utc.with_ymd_and_hms(2026, 1, 1, 0, 0, 0).unwrap();
let h = headers(&[
("anthropic-ratelimit-requests-reset", "2026-01-01T00:00:05Z"),
("anthropic-ratelimit-tokens-reset", "2026-01-01T00:00:20Z"),
]);
assert_eq!(
parse_provider_reset_at(&h, now),
Some(Duration::from_secs(20))
);
}
#[test]
fn provider_reset_unknown_format_is_none() {
let now = Utc::now();
let h = headers(&[("x-ratelimit-reset-tokens", "not-a-duration")]);
assert_eq!(parse_provider_reset_at(&h, now), None);
assert_eq!(parse_provider_reset_at(&HeaderMap::new(), now), None);
}
#[test]
fn go_duration_variants() {
assert_eq!(parse_go_duration("1s"), Some(Duration::from_secs(1)));
assert_eq!(parse_go_duration("6m0s"), Some(Duration::from_secs(360)));
assert_eq!(parse_go_duration("100ms"), Some(Duration::from_millis(100)));
assert_eq!(parse_go_duration("1h2m3s"), Some(Duration::from_secs(3723)));
assert_eq!(parse_go_duration("1.5s"), Some(Duration::from_millis(1500)));
assert_eq!(parse_go_duration(""), None);
assert_eq!(parse_go_duration("abc"), None);
assert_eq!(parse_go_duration("10"), None); assert_eq!(parse_go_duration("5x"), None); }
#[test]
fn backoff_bound_doubles_and_caps() {
let base = Duration::from_secs(1);
let cap = Duration::from_secs(60);
assert_eq!(backoff_bound(0, base, cap), Duration::from_secs(1));
assert_eq!(backoff_bound(1, base, cap), Duration::from_secs(2));
assert_eq!(backoff_bound(2, base, cap), Duration::from_secs(4));
assert_eq!(backoff_bound(10, base, cap), cap);
assert_eq!(backoff_bound(200, base, cap), cap);
}
#[test]
fn full_jitter_stays_within_bound() {
let bound = Duration::from_millis(1000);
for rand in [0u64, 1, 500, 1000, 1001, u64::MAX] {
let j = full_jitter(bound, rand);
assert!(j <= bound, "jitter {j:?} exceeded bound {bound:?}");
}
assert_eq!(full_jitter(bound, 0), Duration::ZERO);
assert_eq!(full_jitter(Duration::ZERO, u64::MAX), Duration::ZERO);
}
#[test]
fn backoff_with_jitter_within_bound_over_many_draws() {
let base = Duration::from_secs(1);
let cap = Duration::from_secs(60);
for attempt in 0..4 {
let bound = backoff_bound(attempt, base, cap);
for _ in 0..200 {
let d = backoff_with_jitter(attempt, base, cap);
assert!(d <= bound, "{d:?} exceeded bound {bound:?}");
}
}
}
#[test]
fn retry_after_wait_caps_bogus_header() {
let h = headers(&[("retry-after", "999999")]);
let cap = Duration::from_secs(60);
let w = retry_after_wait(&h, cap).unwrap();
assert!(w >= cap && w < cap + Duration::from_millis(251));
}
#[test]
fn retry_after_wait_none_without_hints() {
assert_eq!(
retry_after_wait(&HeaderMap::new(), Duration::from_secs(60)),
None
);
}
#[test]
fn retry_config_defaults() {
let d = RetryConfig::default();
assert_eq!(d.max_retries, 3);
assert_eq!(d.base, Duration::from_secs(1));
assert_eq!(d.cap, Duration::from_secs(60));
}
#[test]
fn env_parse_valid_and_invalid() {
std::env::set_var("LLMSHIM_TEST_ENV_PARSE_OK", "7");
std::env::set_var("LLMSHIM_TEST_ENV_PARSE_BAD", "not-a-number");
assert_eq!(env_parse::<u32>("LLMSHIM_TEST_ENV_PARSE_OK"), Some(7));
assert_eq!(env_parse::<u32>("LLMSHIM_TEST_ENV_PARSE_BAD"), None);
assert_eq!(env_parse::<u32>("LLMSHIM_TEST_ENV_PARSE_MISSING"), None);
std::env::remove_var("LLMSHIM_TEST_ENV_PARSE_OK");
std::env::remove_var("LLMSHIM_TEST_ENV_PARSE_BAD");
}
#[tokio::test]
async fn redirects_do_not_forward_prompts_or_credentials() {
for status in [301, 302, 303, 307, 308] {
for same_origin in [false, true] {
let mut origin = mockito::Server::new_async().await;
let mut other = mockito::Server::new_async().await;
let destination = if same_origin { &mut origin } else { &mut other };
let location = format!("{}/moved", destination.url());
let forwarded = destination
.mock(if status <= 303 { "GET" } else { "POST" }, "/moved")
.with_status(200)
.with_body("unexpected forwarding")
.expect(0)
.create_async()
.await;
let redirect = origin
.mock("POST", "/v1/chat/completions")
.match_header("x-api-key", "test-secret")
.match_body(mockito::Matcher::Json(serde_json::json!({
"messages": [{"role": "user", "content": "private transcript"}]
})))
.with_status(status)
.with_header("location", &location)
.expect(1)
.create_async()
.await;
let request = ProviderRequest {
url: format!("{}/v1/chat/completions", origin.url()),
headers: vec![("x-api-key".into(), "test-secret".into())],
body: serde_json::json!({
"messages": [{"role": "user", "content": "private transcript"}]
}),
};
let result = ShimClient::new().send(&request).await;
redirect.assert_async().await;
forwarded.assert_async().await;
assert!(
matches!(result, Err(ShimError::ProviderError { status: actual, .. }) if actual == status as u16)
);
}
}
}
#[tokio::test]
async fn honors_retry_after_then_succeeds() {
let mut server = mockito::Server::new_async().await;
let m429 = server
.mock("POST", "/v1/chat")
.with_status(429)
.with_header("retry-after", "1")
.with_body("rate limited")
.expect(1)
.create_async()
.await;
let m200 = server
.mock("POST", "/v1/chat")
.with_status(200)
.with_body("ok")
.expect(1)
.create_async()
.await;
let client = ShimClient::new();
let req = ProviderRequest {
url: format!("{}/v1/chat", server.url()),
headers: vec![],
body: serde_json::json!({"hello": "world"}),
};
let start = std::time::Instant::now();
let resp = client.send(&req).await.expect("should succeed after retry");
let elapsed = start.elapsed();
assert!(resp.status().is_success());
assert!(
elapsed >= Duration::from_millis(900),
"expected ~1s Retry-After wait, got {elapsed:?}"
);
assert_eq!(resp.text().await.unwrap(), "ok");
m429.assert_async().await;
m200.assert_async().await;
}
#[tokio::test]
async fn falls_back_to_jittered_backoff_without_header() {
let mut server = mockito::Server::new_async().await;
let m500 = server
.mock("POST", "/v1/chat")
.with_status(500)
.with_body("boom")
.expect(1)
.create_async()
.await;
let m200 = server
.mock("POST", "/v1/chat")
.with_status(200)
.with_body("ok")
.expect(1)
.create_async()
.await;
let client = ShimClient::new();
let req = ProviderRequest {
url: format!("{}/v1/chat", server.url()),
headers: vec![],
body: serde_json::json!({}),
};
let start = std::time::Instant::now();
let resp = client.send(&req).await.expect("should succeed after retry");
let elapsed = start.elapsed();
assert!(resp.status().is_success());
assert!(
elapsed < Duration::from_secs(3),
"backoff should be sub-cap jitter, got {elapsed:?}"
);
m500.assert_async().await;
m200.assert_async().await;
}
}