use std::time::Duration;
use reqwest::{RequestBuilder, Response, StatusCode};
use tracing::debug;
pub const ATTEMPTS: usize = 3;
pub fn backoff(attempt: usize) -> Duration {
Duration::from_millis(500 * attempt as u64)
}
pub fn retryable_status(status: StatusCode) -> bool {
status.is_server_error() || status == StatusCode::TOO_MANY_REQUESTS
}
pub fn retryable_error(e: &reqwest::Error) -> bool {
e.is_timeout() || e.is_connect() || e.is_request()
}
pub async fn send(rb: RequestBuilder) -> reqwest::Result<Response> {
let mut attempt = 1;
let mut current = rb;
loop {
let next = (attempt < ATTEMPTS).then(|| current.try_clone()).flatten();
let result = current.send().await;
let retry = match &result {
Ok(r) => retryable_status(r.status()),
Err(e) => retryable_error(e),
};
let Some(next_rb) = next.filter(|_| retry) else {
return result;
};
match &result {
Ok(r) => debug!(attempt, status = %r.status(), url = %r.url(), "服务器错误,重试"),
Err(e) => debug!(attempt, "请求失败,重试: {e}"),
}
tokio::time::sleep(backoff(attempt)).await;
attempt += 1;
current = next_rb;
}
}
pub fn host_of(url: &str) -> Option<&str> {
let rest = url.split_once("://")?.1;
let host = rest.split(['/', '?', '#']).next()?;
host.rsplit_once('@')
.map_or(host, |(_, h)| h)
.split(':')
.next()
}
pub fn host_in(url: &str, hosts: &[&str]) -> bool {
host_of(url).is_some_and(|h| {
let h = h.to_ascii_lowercase();
hosts
.iter()
.any(|q| h == *q || h.ends_with(&format!(".{q}")))
})
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
#[test]
fn status_classification() {
assert!(retryable_status(StatusCode::INTERNAL_SERVER_ERROR));
assert!(retryable_status(StatusCode::BAD_GATEWAY));
assert!(retryable_status(StatusCode::TOO_MANY_REQUESTS));
assert!(!retryable_status(StatusCode::OK));
assert!(!retryable_status(StatusCode::NOT_FOUND));
assert!(!retryable_status(StatusCode::FORBIDDEN));
}
#[test]
fn backoff_grows() {
assert_eq!(backoff(1), Duration::from_millis(500));
assert_eq!(backoff(2), Duration::from_millis(1000));
}
async fn server(fail: usize, mode: &'static str) -> (String, Arc<AtomicUsize>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let hits = Arc::new(AtomicUsize::new(0));
let counter = hits.clone();
tokio::spawn(async move {
loop {
let (mut sock, _) = listener.accept().await.unwrap();
let n = counter.fetch_add(1, Ordering::SeqCst);
let mut buf = [0u8; 1024];
let _ = sock.read(&mut buf).await;
let resp: &[u8] = if n < fail && mode == "drop" {
continue; } else if n < fail {
b"HTTP/1.1 503 Service Unavailable\r\ncontent-length: 0\r\nconnection: close\r\n\r\n"
} else {
b"HTTP/1.1 200 OK\r\ncontent-length: 2\r\nconnection: close\r\n\r\nok"
};
let _ = sock.write_all(resp).await;
}
});
(format!("http://{addr}/"), hits)
}
fn client() -> reqwest::Client {
reqwest::Client::builder()
.no_proxy()
.timeout(Duration::from_secs(5))
.build()
.unwrap()
}
#[tokio::test]
async fn retries_dropped_connection() {
let (url, hits) = server(2, "drop").await;
let r = send(client().get(&url)).await.unwrap();
assert_eq!(r.text().await.unwrap(), "ok");
assert_eq!(hits.load(Ordering::SeqCst), 3);
}
#[tokio::test]
async fn retries_503_then_succeeds() {
let (url, hits) = server(1, "503").await;
let r = send(client().post(&url).body("x")).await.unwrap();
assert_eq!(r.status(), StatusCode::OK);
assert_eq!(hits.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn gives_up_after_attempts() {
let (url, hits) = server(10, "503").await;
let r = send(client().get(&url)).await.unwrap();
assert_eq!(r.status(), StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(hits.load(Ordering::SeqCst), ATTEMPTS);
}
#[tokio::test]
async fn no_retry_on_404() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let hits = Arc::new(AtomicUsize::new(0));
let counter = hits.clone();
tokio::spawn(async move {
loop {
let (mut sock, _) = listener.accept().await.unwrap();
counter.fetch_add(1, Ordering::SeqCst);
let mut buf = [0u8; 1024];
let _ = sock.read(&mut buf).await;
let _ = sock
.write_all(
b"HTTP/1.1 404 Not Found\r\ncontent-length: 0\r\nconnection: close\r\n\r\n",
)
.await;
}
});
let r = send(client().get(format!("http://{addr}/"))).await.unwrap();
assert_eq!(r.status(), StatusCode::NOT_FOUND);
assert_eq!(hits.load(Ordering::SeqCst), 1);
}
}