//! Inlined `RetryAfterMiddleware`: parses the `Retry-After` header on //! 429/503 and sleeps before the next request to that URL. //! //! Inlined (MIT, from `melotic/reqwest-retry-after`) so the upstream's //! unbounded `HashMap` storage can be bounded for a //! long-running process. //! //! Policy (review-001 FWD-05 / FWD-11): //! //! - Every parsed deadline is clamped to a ceiling //! ([`RetryAfterMiddleware::with_capacity_and_ceiling`]; the shared //! client wires `HttpClientConfig::retry_after_ceiling`, default //! 300 s) — a hostile upstream cannot park the client on a 10-year //! deadline. //! - Deadlines are recorded against the *effective* (post-redirect) //! URL — the host actually being throttled — not the pre-redirect //! request URL. //! - At capacity, expired entries go first, then the Farthest-future //! deadline is dropped (the entry least likely to matter to a caller //! whose own timeouts are in seconds). //! - Waking is jittered: each waiter sleeps the remaining time minus a //! random slice (up to 25% of the remaining wait, clamped to 2 s), so //! a herd of queued calls does not all fire at the same instant. use std::collections::HashMap; use std::sync::Mutex; use std::time::{Duration, 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 { let now = SystemTime::now(); let capped = now.checked_add(ceiling)?; Some(deadline.min(capped)) } fn parse_retry_after_with_ceiling(value: &str, ceiling: Duration) -> Option { let trimmed = value.trim(); let parsed = if let Ok(secs) = trimmed.parse::() { 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) } /// Per-URL `Retry-After` memory (FWD-05): records the backlog deadline /// an upstream declared and holds back subsequent requests to that URL /// until it elapses. LRU-bounded by capacity; deadlines clamped to the /// configured ceiling. pub struct RetryAfterMiddleware { deadlines: Mutex>, capacity: usize, ceiling: Duration, } impl RetryAfterMiddleware { /// A middleware with the default 300 s `Retry-After` ceiling. /// /// # Panics /// /// Panics when `capacity` is 0 (a rate-limit memory that remembers /// nothing is a construction error, not a runtime fallback). pub fn with_capacity(capacity: usize) -> Self { Self::with_capacity_and_ceiling(capacity, Duration::from_secs(300)) } /// A middleware with an explicit `capacity` bound and a ceiling /// clamping hostile upstream deadlines. /// /// # Panics /// /// Panics when `capacity` is 0. pub fn with_capacity_and_ceiling(capacity: usize, ceiling: Duration) -> Self { Self { deadlines: Mutex::new(HashMap::with_capacity(capacity.min(128))), capacity, ceiling, } } fn record(&self, url: Url, deadline: 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); } deadlines.insert(url, deadline); } fn evict(&self, deadlines: &mut HashMap) { 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 { let deadlines = self.deadlines.lock().unwrap_or_else(|e| e.into_inner()); deadlines.get(url).copied() } async fn maybe_sleep_for(&self, url: &Url) { let Some(deadline) = self.deadline_for(url) else { return; }; let Ok(remaining) = deadline.duration_since(SystemTime::now()) else { return; }; if remaining.is_zero() { return; } let jitter = max_sleep_jitter(remaining); let wait = remaining .checked_sub(jitter) .unwrap_or(Duration::from_millis(1)); tokio::time::sleep(wait).await; } fn record_if_throttled(&self, url: Url, response: &Response) { 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) { self.record(url, deadline); } } } } #[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 { self.deadline_for(url) } #[cfg(test)] fn record_test(&self, url: Url, deadline: SystemTime) { self.record(url, deadline); } } 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 { let req_url = req.url().clone(); self.maybe_sleep_for(&req_url).await; let response = next.run(req, extensions).await?; self.record_if_throttled(response.url().clone(), &response); Ok(response) } } #[cfg(test)] mod tests { use super::*; fn url(s: &str) -> Url { Url::parse(s).unwrap() } 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); 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); 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); 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); 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); 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); 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); 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); 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).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).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" ); } }