use std::time::{Duration, Instant};
use reqwest::{Response, ResponseBuilderExt as _};
pub(crate) const DEFAULT_CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
pub(crate) const DEFAULT_READ_TIMEOUT: Duration = Duration::from_secs(120);
pub(crate) const CONNECT_TIMEOUT_ENV_VAR: &str = "OMNI_DEV_HTTP_CONNECT_TIMEOUT_SECS";
pub(crate) const READ_TIMEOUT_ENV_VAR: &str = "OMNI_DEV_HTTP_READ_TIMEOUT_SECS";
pub(crate) fn connect_timeout() -> Duration {
duration_from_secs(
crate::utils::settings::get_env_var(CONNECT_TIMEOUT_ENV_VAR).ok(),
DEFAULT_CONNECT_TIMEOUT,
)
}
pub(crate) fn read_timeout() -> Duration {
duration_from_secs(
crate::utils::settings::get_env_var(READ_TIMEOUT_ENV_VAR).ok(),
DEFAULT_READ_TIMEOUT,
)
}
fn duration_from_secs(raw: Option<String>, default: Duration) -> Duration {
raw.and_then(|v| v.parse::<u64>().ok())
.filter(|&secs| secs > 0)
.map_or(default, Duration::from_secs)
}
const MAX_RETRIES: u32 = 3;
const DEFAULT_RETRY_DELAY_SECS: u64 = 2;
pub(crate) async fn retry_429<B, L>(build: B, log: L) -> reqwest::Result<Response>
where
B: Fn() -> reqwest::RequestBuilder,
L: Fn(Instant, &reqwest::Result<Response>),
{
retry_if(build, log, |status, _body| status == 429).await
}
pub(crate) async fn retry_if<B, L, P>(
build: B,
log: L,
is_retryable: P,
) -> reqwest::Result<Response>
where
B: Fn() -> reqwest::RequestBuilder,
L: Fn(Instant, &reqwest::Result<Response>),
P: Fn(u16, &[u8]) -> bool,
{
let mut attempt = 0;
loop {
let started = Instant::now();
let result = build().send().await;
log(started, &result);
let response = result?;
if response.status().is_success() {
return Ok(response);
}
let status = response.status();
let version = response.version();
let url = response.url().clone();
let headers = response.headers().clone();
let body = response.bytes().await?;
if is_retryable(status.as_u16(), &body) && attempt < MAX_RETRIES {
wait_for_retry(&headers, status.as_u16(), attempt).await;
attempt += 1;
continue;
}
let mut builder = http::Response::builder()
.status(status)
.version(version)
.url(url);
if let Some(header_map) = builder.headers_mut() {
header_map.extend(
headers
.iter()
.map(|(name, value)| (name.clone(), value.clone())),
);
}
#[allow(clippy::expect_used)]
let rebuilt = builder
.body(body)
.expect("rebuilding a response from its own already-valid parts cannot fail");
return Ok(Response::from(rebuilt));
}
}
async fn wait_for_retry(headers: &reqwest::header::HeaderMap, status: u16, attempt: u32) {
let delay = header_u64(headers, "Retry-After")
.or_else(|| header_u64(headers, "X-RateLimit-Reset"))
.unwrap_or_else(|| DEFAULT_RETRY_DELAY_SECS.pow(attempt + 1));
eprintln!(
"Rate limited ({status}). Retrying in {delay}s (attempt {})...",
attempt + 1
);
tokio::time::sleep(Duration::from_secs(delay)).await;
}
fn header_u64(headers: &reqwest::header::HeaderMap, name: &str) -> Option<u64> {
headers
.get(name)
.and_then(|v| v.to_str().ok())
.and_then(|s| s.parse::<u64>().ok())
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
#[test]
fn duration_from_secs_parses_valid_override() {
assert_eq!(
duration_from_secs(Some("45".to_string()), DEFAULT_CONNECT_TIMEOUT),
Duration::from_secs(45)
);
}
#[test]
fn duration_from_secs_falls_back_for_absent_zero_or_garbage() {
for raw in [
None,
Some(String::new()),
Some("0".to_string()),
Some("abc".to_string()),
Some("-5".to_string()),
] {
assert_eq!(
duration_from_secs(raw.clone(), DEFAULT_READ_TIMEOUT),
DEFAULT_READ_TIMEOUT,
"expected default for {raw:?}"
);
}
}
#[test]
fn connect_and_read_timeouts_default_to_documented_values_when_unset() {
assert_eq!(DEFAULT_CONNECT_TIMEOUT, Duration::from_secs(10));
assert_eq!(DEFAULT_READ_TIMEOUT, Duration::from_secs(120));
}
#[tokio::test]
async fn retries_429_then_succeeds_and_logs_each_attempt() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/x"))
.respond_with(ResponseTemplate::new(429).append_header("Retry-After", "0"))
.up_to_n_times(1)
.with_priority(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/x"))
.respond_with(ResponseTemplate::new(200))
.with_priority(2)
.mount(&server)
.await;
let client = reqwest::Client::new();
let url = format!("{}/x", server.uri());
let calls = AtomicUsize::new(0);
let resp = retry_429(
|| client.get(&url),
|_started, _result| {
calls.fetch_add(1, Ordering::SeqCst);
},
)
.await
.unwrap();
assert_eq!(resp.status().as_u16(), 200);
assert_eq!(calls.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn returns_429_after_max_retries() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/x"))
.respond_with(ResponseTemplate::new(429).append_header("Retry-After", "0"))
.mount(&server)
.await;
let client = reqwest::Client::new();
let url = format!("{}/x", server.uri());
let calls = AtomicUsize::new(0);
let resp = retry_429(
|| client.get(&url),
|_s, _r| {
calls.fetch_add(1, Ordering::SeqCst);
},
)
.await
.unwrap();
assert_eq!(resp.status().as_u16(), 429);
assert_eq!(calls.load(Ordering::SeqCst), (MAX_RETRIES + 1) as usize);
}
#[tokio::test]
async fn honours_x_ratelimit_reset() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/x"))
.respond_with(ResponseTemplate::new(429).append_header("X-RateLimit-Reset", "0"))
.up_to_n_times(1)
.with_priority(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/x"))
.respond_with(ResponseTemplate::new(200))
.with_priority(2)
.mount(&server)
.await;
let client = reqwest::Client::new();
let url = format!("{}/x", server.uri());
let resp = retry_429(|| client.get(&url), |_s, _r| {}).await.unwrap();
assert_eq!(resp.status().as_u16(), 200);
}
#[tokio::test]
async fn does_not_retry_non_429() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/x"))
.respond_with(ResponseTemplate::new(500))
.expect(1)
.mount(&server)
.await;
let client = reqwest::Client::new();
let url = format!("{}/x", server.uri());
let resp = retry_429(|| client.get(&url), |_s, _r| {}).await.unwrap();
assert_eq!(resp.status().as_u16(), 500);
}
#[tokio::test]
async fn transport_error_is_returned_without_retry() {
let client = reqwest::Client::builder()
.timeout(Duration::from_millis(200))
.build()
.unwrap();
let url = "http://127.0.0.1:1/x".to_string();
let calls = AtomicUsize::new(0);
let result = retry_429(
|| client.get(&url),
|_s, _r| {
calls.fetch_add(1, Ordering::SeqCst);
},
)
.await;
assert!(result.is_err());
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn retry_if_retries_a_custom_status_when_predicate_says_so() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/x"))
.respond_with(ResponseTemplate::new(403).set_body_string("quota exceeded"))
.up_to_n_times(1)
.with_priority(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/x"))
.respond_with(ResponseTemplate::new(200))
.with_priority(2)
.mount(&server)
.await;
let client = reqwest::Client::new();
let url = format!("{}/x", server.uri());
let resp = retry_if(
|| client.get(&url),
|_s, _r| {},
|status, _body| status == 403,
)
.await
.unwrap();
assert_eq!(resp.status().as_u16(), 200);
}
#[tokio::test]
async fn retry_if_does_not_retry_when_predicate_says_no() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/x"))
.respond_with(ResponseTemplate::new(403).set_body_string("insufficientPermissions"))
.expect(1)
.mount(&server)
.await;
let client = reqwest::Client::new();
let url = format!("{}/x", server.uri());
let resp = retry_if(
|| client.get(&url),
|_s, _r| {},
|status, body| status == 403 && body == b"rateLimitExceeded",
)
.await
.unwrap();
assert_eq!(resp.status().as_u16(), 403);
}
#[tokio::test]
async fn retry_if_preserves_headers_and_body_through_reconstruction_on_give_up() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/x"))
.respond_with(
ResponseTemplate::new(429)
.append_header("X-RateLimit-Remaining", "0")
.set_body_string("too many requests"),
)
.mount(&server)
.await;
let client = reqwest::Client::new();
let url = format!("{}/x", server.uri());
let resp = retry_429(|| client.get(&url), |_s, _r| {}).await.unwrap();
assert_eq!(resp.status().as_u16(), 429);
assert_eq!(resp.headers().get("X-RateLimit-Remaining").unwrap(), "0");
let body = resp.text().await.unwrap();
assert_eq!(body, "too many requests");
}
#[tokio::test]
async fn retry_if_preserves_body_on_first_attempt_give_up() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/x"))
.respond_with(ResponseTemplate::new(403).set_body_string("insufficientPermissions"))
.expect(1)
.mount(&server)
.await;
let client = reqwest::Client::new();
let url = format!("{}/x", server.uri());
let resp = retry_if(|| client.get(&url), |_s, _r| {}, |_s, _b| false)
.await
.unwrap();
let body = resp.text().await.unwrap();
assert_eq!(body, "insufficientPermissions");
}
}