use std::time::Duration;
use reqwest::header::HeaderMap;
use reqwest::{Method, StatusCode};
fn jitter_fraction() -> f64 {
use std::sync::OnceLock;
use std::sync::atomic::{AtomicU64, Ordering};
const GAMMA: u64 = 0x9E37_79B9_7F4A_7C15;
static STATE: OnceLock<AtomicU64> = OnceLock::new();
let state = STATE.get_or_init(|| {
use std::hash::{BuildHasher, Hasher};
let seed = std::collections::hash_map::RandomState::new()
.build_hasher()
.finish();
AtomicU64::new(seed)
});
let mut z = state
.fetch_add(GAMMA, Ordering::Relaxed)
.wrapping_add(GAMMA);
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^= z >> 31;
(z >> 11) as f64 / (1u64 << 53) as f64
}
pub(crate) fn backoff_delay(attempt: u32, base: Duration, max: Duration) -> Duration {
let multiplier = 1u32.checked_shl(attempt).unwrap_or(u32::MAX);
let capped = base.saturating_mul(multiplier).min(max);
capped.mul_f64(jitter_fraction())
}
pub(crate) fn is_retryable_status(status: StatusCode, idempotent: bool) -> bool {
match status.as_u16() {
408 | 429 | 503 => true,
501 | 505 => false,
code => idempotent && (500..=599).contains(&code),
}
}
pub(crate) fn is_proxy_failure_status(status: StatusCode) -> bool {
status == StatusCode::PROXY_AUTHENTICATION_REQUIRED
}
pub(crate) fn is_idempotent(method: &Method) -> bool {
matches!(
method.as_str(),
"GET" | "HEAD" | "OPTIONS" | "PUT" | "DELETE" | "TRACE"
)
}
fn sources(err: &reqwest::Error) -> impl Iterator<Item = &(dyn std::error::Error + 'static)> {
std::iter::successors(std::error::Error::source(err), |inner| inner.source())
}
pub(crate) fn is_transport_error(err: &reqwest::Error) -> bool {
if err.is_connect() || err.is_timeout() {
return true;
}
if !err.is_request() {
return false;
}
sources(err).any(|inner| {
inner
.downcast_ref::<hyper::Error>()
.is_some_and(|e| !e.is_user() && !e.is_parse())
|| inner.downcast_ref::<std::io::Error>().is_some()
|| inner.downcast_ref::<h2::Error>().is_some()
})
}
pub(crate) fn is_never_sent_error(err: &reqwest::Error) -> bool {
if err.is_connect() {
return true;
}
sources(err).any(|inner| {
inner
.downcast_ref::<hyper::Error>()
.is_some_and(hyper::Error::is_canceled)
|| inner
.downcast_ref::<h2::Error>()
.is_some_and(|e| e.reason() == Some(h2::Reason::REFUSED_STREAM))
})
}
pub(crate) fn should_retry_error(err: &reqwest::Error, idempotent: bool) -> bool {
is_never_sent_error(err) || (idempotent && is_transport_error(err))
}
#[cfg_attr(not(feature = "tracing"), allow(dead_code))]
pub(crate) fn error_kind(err: &reqwest::Error) -> &'static str {
if err.is_timeout() {
"timeout"
} else if err.is_connect() {
"connect"
} else if is_never_sent_error(err) {
"never sent"
} else if is_transport_error(err) {
"transport"
} else {
"other"
}
}
pub(crate) fn retry_after(headers: &HeaderMap) -> Option<Duration> {
let value = headers
.get(reqwest::header::RETRY_AFTER)?
.to_str()
.ok()?
.trim();
if let Ok(seconds) = value.parse::<u64>() {
return Some(Duration::from_secs(seconds));
}
let target = httpdate::parse_http_date(value).ok()?;
target.duration_since(std::time::SystemTime::now()).ok()
}
#[cfg(test)]
mod tests {
use super::*;
use reqwest::header::{HeaderValue, RETRY_AFTER};
#[test]
fn backoff_never_exceeds_max() {
let max = Duration::from_secs(10);
for attempt in 0..40 {
let delay = backoff_delay(attempt, Duration::from_millis(100), max);
assert!(delay <= max, "attempt {attempt} produced {delay:?}");
}
}
#[test]
fn backoff_jitter_reaches_the_cap() {
let base = Duration::from_millis(100);
let max = Duration::from_secs(10);
let saw_a_large_sample = (0..1000)
.map(|_| backoff_delay(0, base, max))
.any(|d| d > base.mul_f64(0.9));
assert!(saw_a_large_sample);
}
#[test]
fn jitter_is_uniform_over_the_unit_interval() {
let samples: Vec<f64> = (0..100_000).map(|_| jitter_fraction()).collect();
for &sample in &samples {
assert!((0.0..1.0).contains(&sample), "{sample}");
}
let min = samples.iter().copied().fold(f64::INFINITY, f64::min);
let max = samples.iter().copied().fold(f64::NEG_INFINITY, f64::max);
assert!(min < 0.01, "min {min}");
assert!(max > 0.99, "max {max}");
let mean = samples.iter().sum::<f64>() / samples.len() as f64;
assert!((0.48..=0.52).contains(&mean), "mean {mean}");
}
#[test]
fn idempotent_retries_408_429_and_5xx() {
for status in [
StatusCode::REQUEST_TIMEOUT,
StatusCode::TOO_MANY_REQUESTS,
StatusCode::INTERNAL_SERVER_ERROR,
StatusCode::BAD_GATEWAY,
StatusCode::SERVICE_UNAVAILABLE,
StatusCode::GATEWAY_TIMEOUT,
StatusCode::from_u16(522).unwrap(),
] {
assert!(is_retryable_status(status, true), "{status}");
}
for status in [
StatusCode::NOT_IMPLEMENTED,
StatusCode::HTTP_VERSION_NOT_SUPPORTED,
StatusCode::NOT_FOUND,
StatusCode::FORBIDDEN,
StatusCode::OK,
] {
assert!(!is_retryable_status(status, true), "{status}");
}
}
#[test]
fn non_idempotent_retries_only_unprocessed_statuses() {
for status in [
StatusCode::REQUEST_TIMEOUT,
StatusCode::TOO_MANY_REQUESTS,
StatusCode::SERVICE_UNAVAILABLE,
] {
assert!(is_retryable_status(status, false), "{status}");
}
for status in [
StatusCode::INTERNAL_SERVER_ERROR,
StatusCode::BAD_GATEWAY,
StatusCode::GATEWAY_TIMEOUT,
StatusCode::NOT_IMPLEMENTED,
] {
assert!(!is_retryable_status(status, false), "{status}");
}
}
#[test]
fn proxy_failure_status_is_407_only() {
assert!(is_proxy_failure_status(
StatusCode::PROXY_AUTHENTICATION_REQUIRED
));
assert!(!is_proxy_failure_status(StatusCode::FORBIDDEN));
assert!(!is_proxy_failure_status(StatusCode::NOT_FOUND));
assert!(!is_proxy_failure_status(StatusCode::BAD_GATEWAY));
}
#[test]
fn idempotent_methods() {
for method in [
Method::GET,
Method::HEAD,
Method::OPTIONS,
Method::PUT,
Method::DELETE,
Method::TRACE,
] {
assert!(is_idempotent(&method), "{method}");
}
assert!(!is_idempotent(&Method::POST));
assert!(!is_idempotent(&Method::PATCH));
assert!(!is_idempotent(&Method::CONNECT));
}
#[test]
fn retry_after_parses_delta_seconds() {
let mut headers = HeaderMap::new();
headers.insert(RETRY_AFTER, HeaderValue::from_static("120"));
assert_eq!(retry_after(&headers), Some(Duration::from_secs(120)));
}
#[test]
fn retry_after_parses_http_date() {
let future = std::time::SystemTime::now() + Duration::from_secs(60);
let formatted = httpdate::fmt_http_date(future);
let mut headers = HeaderMap::new();
headers.insert(RETRY_AFTER, HeaderValue::from_str(&formatted).unwrap());
let parsed = retry_after(&headers).unwrap();
assert!((58..=60).contains(&parsed.as_secs()), "{parsed:?}");
}
#[test]
fn retry_after_past_date_is_none() {
let past = std::time::SystemTime::now() - Duration::from_secs(60);
let formatted = httpdate::fmt_http_date(past);
let mut headers = HeaderMap::new();
headers.insert(RETRY_AFTER, HeaderValue::from_str(&formatted).unwrap());
assert_eq!(retry_after(&headers), None);
}
#[test]
fn retry_after_missing_header_is_none() {
let headers = HeaderMap::new();
assert_eq!(retry_after(&headers), None);
}
#[tokio::test]
async fn error_kind_labels_a_connect_failure() {
let err = reqwest::Client::new()
.get("http://127.0.0.1:1/")
.send()
.await
.unwrap_err();
assert_eq!(error_kind(&err), "connect");
let builder_err = reqwest::Proxy::all("http://[").unwrap_err();
assert_eq!(error_kind(&builder_err), "other");
}
#[tokio::test]
async fn error_kind_labels_a_timeout() {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
std::thread::spawn(move || {
if let Ok((stream, _)) = listener.accept() {
std::thread::sleep(Duration::from_secs(30));
drop(stream);
}
});
let client = reqwest::Client::builder()
.timeout(Duration::from_millis(1))
.build()
.unwrap();
let err = client
.get(format!("http://{addr}/"))
.send()
.await
.unwrap_err();
assert_eq!(error_kind(&err), "timeout");
}
#[tokio::test]
async fn error_kind_labels_a_truncated_body_as_other() {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
std::thread::spawn(move || {
if let Ok((mut stream, _)) = listener.accept() {
use std::io::{Read, Write};
let mut buf = [0u8; 1024];
let _ = stream.read(&mut buf);
let _ = stream.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 10\r\n\r\n");
drop(stream);
}
});
let response = reqwest::Client::new()
.get(format!("http://{addr}/"))
.send()
.await
.unwrap();
let err = response.bytes().await.unwrap_err();
assert_eq!(error_kind(&err), "other");
}
#[tokio::test]
async fn error_kind_labels_a_dropped_connection_as_transport() {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
std::thread::spawn(move || {
if let Ok((mut stream, _)) = listener.accept() {
use std::io::Read;
let mut buf = [0u8; 1024];
let _ = stream.read(&mut buf);
drop(stream);
}
});
let err = reqwest::Client::new()
.get(format!("http://{addr}/"))
.send()
.await
.unwrap_err();
assert_eq!(error_kind(&err), "transport");
}
#[test]
fn retry_after_garbage_value_is_none() {
let mut headers = HeaderMap::new();
headers.insert(RETRY_AFTER, HeaderValue::from_static("not-a-date"));
assert_eq!(retry_after(&headers), None);
headers.insert(RETRY_AFTER, HeaderValue::from_static("-5"));
assert_eq!(retry_after(&headers), None);
}
}