use std::collections::HashMap;
use std::sync::Mutex;
use std::time::{Duration, Instant, SystemTime};
use http::Extensions;
use reqwest::{Request, Response, StatusCode};
use reqwest_middleware::{Middleware, Next, Result};
use url::Url;
const RETRY_AFTER_HEADER: &str = "retry-after";
const THROTTLED_STATUS: &[u16] = &[StatusCode::TOO_MANY_REQUESTS.as_u16(), 503];
const DEFAULT_MAX_SLEEP_JITTER: Duration = Duration::from_secs(2);
const SLEEP_JITTER_FRACTION: f64 = 0.25;
fn is_throttled(status: u16) -> bool {
THROTTLED_STATUS.contains(&status)
}
fn clamp_deadline_to_ceiling(deadline: SystemTime, ceiling: Duration) -> Option<SystemTime> {
let now = SystemTime::now();
let capped = now.checked_add(ceiling)?;
Some(deadline.min(capped))
}
#[derive(Clone, Debug)]
pub(crate) struct BudgetClock {
hard_stop: SystemTime,
anchored: Instant,
}
impl BudgetClock {
pub(crate) fn start(max_total_retry_duration: Duration) -> Self {
Self {
hard_stop: SystemTime::now()
.checked_add(max_total_retry_duration)
.unwrap_or(SystemTime::now()),
anchored: Instant::now(),
}
}
pub(crate) fn for_duration(anchor: SystemTime, max_total_retry_duration: Duration) -> Self {
Self {
hard_stop: anchor
.checked_add(max_total_retry_duration)
.unwrap_or(anchor),
anchored: Instant::now(),
}
}
fn hard_stop(&self) -> SystemTime {
let offset = self
.hard_stop
.duration_since(self.anchor_wall())
.unwrap_or(Duration::ZERO);
self.anchor_wall()
.checked_add(offset)
.unwrap_or(self.hard_stop)
}
fn anchor_wall(&self) -> SystemTime {
let now = SystemTime::now();
now.checked_sub(self.anchored.elapsed()).unwrap_or(now)
}
}
fn budget_from_extensions(extensions: &Extensions) -> Option<BudgetClock> {
if let Some(clock) = extensions.get::<BudgetClock>() {
return Some(clock.clone());
}
extensions
.get::<RequestStartTime>()
.map(|start| BudgetClock::for_duration(start.0, Duration::ZERO))
}
pub(crate) fn anchor_budget(extensions: &mut Extensions, budget: Duration) {
if extensions.get::<BudgetClock>().is_none() {
extensions.insert(BudgetClock::start(budget));
}
}
#[derive(Clone, Debug)]
struct RequestStartTime(SystemTime);
fn parse_retry_after_with_ceiling(value: &str, ceiling: Duration) -> Option<SystemTime> {
let trimmed = value.trim();
let parsed = if let Ok(secs) = trimmed.parse::<u64>() {
SystemTime::now().checked_add(Duration::from_secs(secs))?
} else {
httpdate::parse_http_date(trimmed).ok()?
};
let clamped = clamp_deadline_to_ceiling(parsed, ceiling)?;
if clamped <= SystemTime::now() {
return None;
}
Some(clamped)
}
pub struct RetryAfterMiddleware {
deadlines: Mutex<HashMap<Url, SystemTime>>,
capacity: usize,
ceiling: Duration,
budget: Duration,
}
impl RetryAfterMiddleware {
pub fn with_capacity(capacity: usize) -> Self {
Self::with_capacity_and_ceiling(capacity, Duration::from_secs(300))
}
pub fn with_capacity_and_ceiling(capacity: usize, ceiling: Duration) -> Self {
Self::with_capacity_ceiling_and_budget(capacity, ceiling, Duration::ZERO)
}
pub fn with_capacity_ceiling_and_budget(
capacity: usize,
ceiling: Duration,
budget: Duration,
) -> Self {
Self {
deadlines: Mutex::new(HashMap::with_capacity(capacity.min(128))),
capacity,
ceiling,
budget,
}
}
fn record(&self, url: Url, deadline: Option<SystemTime>) {
let mut deadlines = self.deadlines.lock().unwrap_or_else(|e| e.into_inner());
if !deadlines.contains_key(&url) && deadlines.len() >= self.capacity {
self.evict(&mut deadlines);
}
match deadline {
Some(deadline) => {
let earliest = deadlines
.get(&url)
.map(|existing| (*existing).min(deadline))
.unwrap_or(deadline);
deadlines.insert(url, earliest);
}
None => {
deadlines.remove(&url);
}
}
}
fn evict(&self, deadlines: &mut HashMap<Url, SystemTime>) {
let now = SystemTime::now();
let expired = deadlines
.iter()
.filter(|(_, deadline)| **deadline <= now)
.map(|(url, _)| url.clone())
.next();
if let Some(expired_url) = expired {
deadlines.remove(&expired_url);
return;
}
if let Some(farthest) = deadlines
.iter()
.max_by_key(|(_, deadline)| deadline.duration_since(now).unwrap_or_default())
.map(|(url, _)| url.clone())
{
deadlines.remove(&farthest);
}
}
fn deadline_for(&self, url: &Url) -> Option<SystemTime> {
let deadlines = self.deadlines.lock().unwrap_or_else(|e| e.into_inner());
deadlines.get(url).copied()
}
async fn maybe_sleep_for(&self, url: &Url, extensions: &Extensions) {
let Some(deadline) = self.deadline_for(url) else {
return;
};
let now = SystemTime::now();
let Ok(remaining) = deadline.duration_since(now) else {
return;
};
let hard_stop = match budget_from_extensions(extensions) {
Some(clock) => clock.hard_stop(),
None if self.budget.is_zero() => {
let jitter = max_sleep_jitter(remaining);
let wait = remaining
.checked_sub(jitter)
.unwrap_or(Duration::from_millis(1));
tokio::time::sleep(wait).await;
return;
}
None => SystemTime::now() + self.budget,
};
let budget_left = hard_stop.duration_since(now).unwrap_or(Duration::ZERO);
let remaining = remaining.min(budget_left);
if remaining.is_zero() {
return;
}
let wait = remaining
.checked_sub(max_sleep_jitter(remaining))
.unwrap_or(Duration::from_millis(1));
tokio::time::sleep(wait).await;
}
fn record_if_throttled(&self, url: Url, response: &Response, extensions: &Extensions) {
let status = response.status();
if is_throttled(status.as_u16()) {
if let Some(retry_after) = response
.headers()
.get(RETRY_AFTER_HEADER)
.and_then(|value| value.to_str().ok())
{
if let Some(deadline) = parse_retry_after_with_ceiling(retry_after, self.ceiling) {
let stored = match budget_from_extensions(extensions) {
Some(clock) => {
let clamped = deadline.min(clock.hard_stop());
(clamped > SystemTime::now()).then_some(clamped)
}
None if self.budget.is_zero() => Some(deadline),
None => {
let clamped = deadline.min(SystemTime::now() + self.budget);
(clamped > SystemTime::now()).then_some(clamped)
}
};
self.record(url, stored);
}
}
}
}
#[cfg(test)]
fn len(&self) -> usize {
self.deadlines
.lock()
.unwrap_or_else(|e| e.into_inner())
.len()
}
#[cfg(test)]
fn deadline_for_test(&self, url: &Url) -> Option<SystemTime> {
self.deadline_for(url)
}
#[cfg(test)]
fn record_test(&self, url: Url, deadline: SystemTime) {
self.record(url, Some(deadline));
}
#[cfg(test)]
fn record_with_budget_test(&self, url: Url, budget: Duration, header: Option<&str>) {
let mut extensions = Extensions::new();
extensions.insert(BudgetClock::start(budget));
let response = crate::client::retry_after::tests::synthetic_response(
StatusCode::TOO_MANY_REQUESTS,
header,
);
self.record_if_throttled(url, &response, &extensions);
}
}
fn max_sleep_jitter(remaining: Duration) -> Duration {
if remaining.is_zero() {
return Duration::ZERO;
}
let fractional = remaining.as_secs_f64() * SLEEP_JITTER_FRACTION;
let fractional = if fractional.is_finite() && fractional > 0.0 {
fractional
} else {
0.0
};
let capped = fractional.min(DEFAULT_MAX_SLEEP_JITTER.as_secs_f64());
Duration::try_from_secs_f64(capped).unwrap_or(Duration::ZERO)
}
#[async_trait::async_trait]
impl Middleware for RetryAfterMiddleware {
async fn handle(
&self,
req: Request,
extensions: &mut Extensions,
next: Next<'_>,
) -> Result<Response> {
let req_url = req.url().clone();
self.maybe_sleep_for(&req_url, extensions).await;
let response = next.run(req, extensions).await?;
self.record_if_throttled(response.url().clone(), &response, extensions);
Ok(response)
}
}
#[cfg(test)]
pub(crate) mod tests {
use super::*;
fn url(s: &str) -> Url {
Url::parse(s).unwrap()
}
pub(crate) fn synthetic_response(status: StatusCode, retry_after: Option<&str>) -> Response {
let mut builder = http::Response::builder().status(status);
if let Some(value) = retry_after {
builder = builder.header(RETRY_AFTER_HEADER, value);
}
builder.body("").unwrap().into()
}
#[test]
fn parse_retry_after_seconds() {
let mw = RetryAfterMiddleware::with_capacity_and_ceiling(8, Duration::from_secs(300));
let target = url("https://api.example.com/v1/chat");
let response = synthetic_response(StatusCode::TOO_MANY_REQUESTS, Some("5"));
mw.record_if_throttled(target.clone(), &response, &Extensions::new());
let deadline = mw.deadline_for_test(&target).expect("seconds parse");
let now = SystemTime::now();
assert!(deadline > now);
assert!(deadline < now + Duration::from_secs(6));
}
#[test]
fn parse_retry_after_http_date() {
let mw = RetryAfterMiddleware::with_capacity_and_ceiling(8, Duration::from_secs(300));
let target = url("https://api.example.com/v1/chat");
let response = synthetic_response(
StatusCode::SERVICE_UNAVAILABLE,
Some("Wed, 21 Oct 2099 07:28:00 GMT"),
);
mw.record_if_throttled(target.clone(), &response, &Extensions::new());
let deadline = mw.deadline_for_test(&target).expect("HTTP-date parses");
let ceiling = SystemTime::now() + Duration::from_secs(300);
assert!(
deadline <= ceiling,
"HTTP-date deadlines must be clamped to the ceiling, got {deadline:?}"
);
assert!(deadline > SystemTime::now());
}
#[test]
fn parse_retry_after_past_http_date_yields_none() {
let mw = RetryAfterMiddleware::with_capacity_and_ceiling(8, Duration::from_secs(300));
let target = url("https://api.example.com/v1/chat");
let response = synthetic_response(
StatusCode::TOO_MANY_REQUESTS,
Some("Wed, 21 Oct 2015 07:28:00 GMT"),
);
mw.record_if_throttled(target.clone(), &response, &Extensions::new());
assert!(
mw.deadline_for_test(&target).is_none(),
"a deadline already in the past must not be recorded"
);
}
#[test]
fn parse_retry_after_invalid_yields_none() {
let ceiling = Duration::from_secs(300);
assert!(parse_retry_after_with_ceiling("not-a-date", ceiling).is_none());
assert!(parse_retry_after_with_ceiling("", ceiling).is_none());
}
#[test]
fn retry_after_seconds_are_clamped_to_the_ceiling() {
let mw = RetryAfterMiddleware::with_capacity_and_ceiling(8, Duration::from_secs(300));
let target = url("https://api.example.com/v1/chat");
let response = synthetic_response(StatusCode::TOO_MANY_REQUESTS, Some("315360000"));
mw.record_if_throttled(target.clone(), &response, &Extensions::new());
let deadline = mw
.deadline_for_test(&target)
.expect("a huge Retry-After still records (clamped)");
let ceiling = SystemTime::now() + Duration::from_secs(300);
assert!(
deadline <= ceiling,
"deadline must be clamped to the configured ceiling"
);
assert!(deadline > SystemTime::now());
}
#[test]
fn ceiling_is_configurable() {
let mw = RetryAfterMiddleware::with_capacity_and_ceiling(8, Duration::from_secs(5));
let target = url("https://api.example.com/v1/chat");
let response = synthetic_response(StatusCode::TOO_MANY_REQUESTS, Some("3600"));
mw.record_if_throttled(target.clone(), &response, &Extensions::new());
let deadline = mw
.deadline_for_test(&target)
.expect("clamped deadline kept");
let upper = SystemTime::now() + Duration::from_secs(5);
assert!(deadline <= upper, "custom ceiling must be honored");
assert!(deadline > SystemTime::now());
}
#[test]
fn record_stores_deadline_for_url() {
let mw = RetryAfterMiddleware::with_capacity(8);
let u = url("https://api.example.com/v1/chat");
let deadline = SystemTime::now() + Duration::from_secs(10);
mw.record_test(u.clone(), deadline);
assert_eq!(mw.deadline_for_test(&u), Some(deadline));
assert_eq!(mw.len(), 1);
}
#[test]
fn record_evicts_expired_entries_first() {
let mw = RetryAfterMiddleware::with_capacity(2);
let expired = url("https://expired.example.com");
let u1 = url("https://a.example.com");
let u2 = url("https://b.example.com");
mw.record_test(expired.clone(), SystemTime::now() - Duration::from_secs(1));
mw.record_test(u1.clone(), SystemTime::now() + Duration::from_secs(100));
mw.record_test(u2.clone(), SystemTime::now() + Duration::from_secs(50));
assert_eq!(mw.len(), 2, "capacity must be enforced");
assert!(
mw.deadline_for_test(&expired).is_none(),
"an expired entry must be evicted before live ones"
);
assert!(mw.deadline_for_test(&u1).is_some());
assert!(mw.deadline_for_test(&u2).is_some());
}
#[test]
fn record_evicts_the_farthest_deadline_when_none_expired() {
let mw = RetryAfterMiddleware::with_capacity(2);
let u1 = url("https://a.example.com");
let u2 = url("https://b.example.com");
let u3 = url("https://c.example.com");
mw.record_test(u1.clone(), SystemTime::now() + Duration::from_secs(100));
mw.record_test(u2.clone(), SystemTime::now() + Duration::from_secs(1));
mw.record_test(u3.clone(), SystemTime::now() + Duration::from_secs(50));
assert_eq!(mw.len(), 2, "capacity must be enforced");
assert!(
mw.deadline_for_test(&u1).is_none(),
"the farthest-future entry must be evicted when nothing has expired"
);
assert!(mw.deadline_for_test(&u2).is_some());
assert!(mw.deadline_for_test(&u3).is_some());
}
#[test]
fn record_overwrites_existing_url_deadline_without_evicting() {
let mw = RetryAfterMiddleware::with_capacity(2);
let u = url("https://api.example.com/v1/chat");
let far = url("https://far.example.com");
mw.record_test(u.clone(), SystemTime::now() + Duration::from_secs(10));
mw.record_test(far.clone(), SystemTime::now() + Duration::from_secs(20));
mw.record_test(u.clone(), SystemTime::now() + Duration::from_secs(30));
assert_eq!(
mw.len(),
2,
"overwriting an existing URL must not evict another entry"
);
assert!(mw.deadline_for_test(&far).is_some());
}
#[tokio::test]
async fn middleware_records_under_the_effective_url() {
let mw = std::sync::Arc::new(RetryAfterMiddleware::with_capacity(8));
let origin = url("https://api.example.com/v1/chat");
let redirector = url("https://redirector.example.com/429");
let response = synthetic_response(StatusCode::TOO_MANY_REQUESTS, Some("5"));
mw.record_if_throttled(origin.clone(), &response, &Extensions::new());
assert!(
mw.deadline_for_test(&origin).is_some(),
"the effective (post-redirect) URL carries the deadline"
);
assert_eq!(redirector.host_str(), Some("redirector.example.com"));
}
#[tokio::test]
async fn middleware_does_not_record_on_non_throttled_status() {
let mw = std::sync::Arc::new(RetryAfterMiddleware::with_capacity(8));
let target = url("https://api.example.com/v1/chat");
let response = synthetic_response(StatusCode::OK, Some("5"));
mw.record_if_throttled(target.clone(), &response, &Extensions::new());
assert!(mw.deadline_for_test(&target).is_none());
}
#[tokio::test]
async fn middleware_does_not_record_when_header_absent() {
let mw = std::sync::Arc::new(RetryAfterMiddleware::with_capacity(8));
let target = url("https://api.example.com/v1/chat");
let response = synthetic_response(StatusCode::TOO_MANY_REQUESTS, None);
mw.record_if_throttled(target.clone(), &response, &Extensions::new());
assert!(mw.deadline_for_test(&target).is_none());
}
#[tokio::test]
async fn middleware_sleeps_before_request_with_active_deadline() {
let mw = std::sync::Arc::new(RetryAfterMiddleware::with_capacity(8));
let target = url("https://api.example.com/v1/chat");
mw.record_test(
target.clone(),
SystemTime::now() + Duration::from_millis(50),
);
let started = SystemTime::now();
mw.maybe_sleep_for(&target, &Extensions::new()).await;
let elapsed = SystemTime::now().duration_since(started).unwrap();
assert!(
elapsed >= Duration::from_millis(37),
"middleware must sleep (minus jitter) until the deadline elapses"
);
assert!(elapsed < Duration::from_secs(2));
}
#[tokio::test]
async fn sleep_wakes_before_the_deadline_within_the_jitter_bound() {
let mw = std::sync::Arc::new(RetryAfterMiddleware::with_capacity(8));
let target = url("https://api.example.com/v1/chat");
let remaining = Duration::from_secs(4);
mw.record_test(target.clone(), SystemTime::now() + remaining);
let started = SystemTime::now();
mw.maybe_sleep_for(&target, &Extensions::new()).await;
let elapsed = SystemTime::now().duration_since(started).unwrap();
let max_jitter = max_sleep_jitter(remaining);
assert!(
elapsed
<= remaining.checked_sub(max_jitter).unwrap_or(remaining)
+ Duration::from_millis(50),
"wake must happen roughly a jitter-slice before the deadline, took {elapsed:?}"
);
assert!(max_jitter > Duration::ZERO, "jitter must be non-zero");
assert!(max_jitter <= Duration::from_secs(2));
}
#[test]
fn jitter_is_bounded_by_a_fraction_of_the_remaining_wait() {
let short = max_sleep_jitter(Duration::from_millis(200));
assert!(short <= Duration::from_millis(50));
let long = max_sleep_jitter(Duration::from_secs(100));
assert!(
long <= Duration::from_secs(2),
"jitter is capped at 2 s even for long waits"
);
}
#[tokio::test]
async fn sleep_is_truncated_to_the_retry_budget() {
let mw = std::sync::Arc::new(RetryAfterMiddleware::with_capacity_ceiling_and_budget(
8,
Duration::from_secs(300),
Duration::from_secs(1),
));
let target = url("https://api.example.com/v1/chat");
mw.record_test(target.clone(), SystemTime::now() + Duration::from_secs(60));
let mut extensions = Extensions::new();
extensions.insert(BudgetClock::start(Duration::from_secs(1)));
let started = Instant::now();
mw.maybe_sleep_for(&target, &extensions).await;
let elapsed = started.elapsed();
assert!(
elapsed <= Duration::from_millis(970),
"the sleep must be truncated to the remaining budget, took {elapsed:?}"
);
assert!(elapsed >= Duration::from_millis(700));
}
#[tokio::test]
async fn budget_spent_skips_the_sleep_entirely() {
let mw = std::sync::Arc::new(RetryAfterMiddleware::with_capacity_ceiling_and_budget(
8,
Duration::from_secs(300),
Duration::from_secs(1),
));
let target = url("https://api.example.com/v1/chat");
mw.record_test(target.clone(), SystemTime::now() + Duration::from_secs(60));
let expired = BudgetClock::for_duration(
SystemTime::now() - Duration::from_secs(5),
Duration::from_secs(1),
);
let mut extensions = Extensions::new();
extensions.insert(expired);
let started = Instant::now();
mw.maybe_sleep_for(&target, &extensions).await;
assert!(
started.elapsed() <= Duration::from_millis(60),
"a spent budget must skip the Retry-After sleep, took {:?}",
started.elapsed()
);
}
#[test]
fn throttle_refresh_is_clamped_to_the_budget_hard_stop() {
let mw = std::sync::Arc::new(RetryAfterMiddleware::with_capacity_ceiling_and_budget(
8,
Duration::from_secs(300),
Duration::from_secs(1),
));
let target = url("https://api.example.com/v1/chat");
let response = synthetic_response(StatusCode::TOO_MANY_REQUESTS, Some("300"));
let mut extensions = Extensions::new();
extensions.insert(BudgetClock::start(Duration::from_secs(1)));
mw.record_if_throttled(target.clone(), &response, &extensions);
let deadline = mw
.deadline_for_test(&target)
.expect("a live deadline records (clamped to the hard stop)");
assert!(
deadline <= SystemTime::now() + Duration::from_millis(1100),
"a retry-storm refresh must not push the deadline past the budget, got {deadline:?}"
);
}
#[test]
fn throttle_storm_cannot_extend_the_first_seen_deadline() {
let mw = std::sync::Arc::new(RetryAfterMiddleware::with_capacity_ceiling_and_budget(
8,
Duration::from_secs(300),
Duration::from_secs(300),
));
let target = url("https://api.example.com/v1/chat");
let first = synthetic_response(StatusCode::TOO_MANY_REQUESTS, Some("10"));
mw.record_if_throttled(target.clone(), &first, &Extensions::new());
let first_deadline = mw.deadline_for_test(&target).expect("first deadline");
let storm = synthetic_response(StatusCode::TOO_MANY_REQUESTS, Some("300"));
mw.record_if_throttled(target.clone(), &storm, &Extensions::new());
assert_eq!(
mw.deadline_for_test(&target),
Some(first_deadline),
"a later higher Retry-After must not extend the first-seen deadline"
);
}
#[test]
fn budget_exhaustion_drops_the_throttle_entry() {
let mw = std::sync::Arc::new(RetryAfterMiddleware::with_capacity_ceiling_and_budget(
8,
Duration::from_secs(300),
Duration::from_secs(300),
));
let target = url("https://api.example.com/v1/chat");
mw.record_with_budget_test(target.clone(), Duration::from_secs(300), None);
assert!(
mw.deadline_for_test(&target).is_none(),
"a refresh arriving after the budget is spent must drop the entry"
);
}
#[test]
fn missing_header_refresh_after_budget_exhaustion_removes_the_deadline() {
let mw = std::sync::Arc::new(RetryAfterMiddleware::with_capacity(8));
let target = url("https://api.example.com/v1/chat");
mw.record_test(target.clone(), SystemTime::now() + Duration::from_secs(10));
mw.record_with_budget_test(target.clone(), Duration::from_secs(300), None);
assert!(mw.deadline_for_test(&target).is_some());
}
#[test]
fn budget_clock_hard_stop_survives_a_wall_clock_step() {
let clock = BudgetClock::start(Duration::from_secs(30));
std::thread::sleep(Duration::from_millis(30));
let projected = clock.hard_stop();
assert!(
projected <= SystemTime::now() + Duration::from_secs(30),
"hard_stop must track monotonic progress against a stepped wall clock, got {projected:?}"
);
}
}