use http::{Extensions, Method};
use reqwest::{Request, Response};
use reqwest_middleware::{Middleware, Next, Result};
#[cfg(feature = "dns")]
mod stale_address;
#[cfg(feature = "dns")]
pub use stale_address::StaleAddressRetry;
const MAX_REPLAYS: usize = 5;
#[derive(Debug, Clone, Copy, Default)]
pub struct DeadConnectionRetry;
fn is_idempotent(method: &Method) -> bool {
matches!(
*method,
Method::GET | Method::HEAD | Method::OPTIONS | Method::TRACE | Method::PUT | Method::DELETE
) || method.as_str() == crate::request::QUERY
}
fn died_before_response(err: &reqwest_middleware::Error) -> bool {
let mut source: Option<&(dyn std::error::Error + 'static)> = Some(err);
while let Some(err) = source {
if let Some(err) = err.downcast_ref::<hyper::Error>() {
if err.is_incomplete_message() || err.is_closed() {
return true;
}
}
source = err.source();
}
false
}
#[async_trait::async_trait]
impl Middleware for DeadConnectionRetry {
async fn handle(
&self,
req: Request,
extensions: &mut Extensions,
next: Next<'_>,
) -> Result<Response> {
let mut replay = is_idempotent(req.method())
.then(|| req.try_clone())
.flatten();
let mut outcome = next.clone().run(req, extensions).await;
for _ in 0..MAX_REPLAYS {
match &outcome {
Err(err) if died_before_response(err) => {}
_ => return outcome,
}
let Some(request) = replay.take() else {
return outcome;
};
replay = request.try_clone();
outcome = next.clone().run(request, extensions).await;
}
outcome
}
}
#[cfg(feature = "dns")]
pub(super) fn failed_to_connect(err: &reqwest_middleware::Error) -> bool {
match err {
reqwest_middleware::Error::Reqwest(err) => err.is_connect(),
reqwest_middleware::Error::Middleware(_) => false,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn idempotent_methods_are_the_replayable_ones() {
for method in [
Method::GET,
Method::HEAD,
Method::OPTIONS,
Method::TRACE,
Method::PUT,
Method::DELETE,
Method::from_bytes(b"QUERY").unwrap(),
] {
assert!(is_idempotent(&method), "{method} should be replayable");
}
for method in [Method::POST, Method::PATCH, Method::CONNECT] {
assert!(!is_idempotent(&method), "{method} should not be replayable");
}
}
}