use http::Extensions;
use reqwest::{Request, Response};
use reqwest_middleware::{Middleware, Next, Result};
use web_faith_dns::FaithResolver;
use super::failed_to_connect;
#[derive(Debug, Clone)]
pub struct StaleAddressRetry {
resolver: Option<FaithResolver>,
}
impl StaleAddressRetry {
pub fn new(resolver: Option<FaithResolver>) -> Self {
Self { resolver }
}
}
#[async_trait::async_trait]
impl Middleware for StaleAddressRetry {
async fn handle(
&self,
req: Request,
extensions: &mut Extensions,
next: Next<'_>,
) -> Result<Response> {
let Some(resolver) = self.resolver.clone() else {
return next.run(req, extensions).await;
};
let host = req.url().host_str().map(str::to_owned);
let replay = match &host {
Some(host) if resolver.served_stale(host) => req.try_clone(),
_ => None,
};
let outcome = next.clone().run(req, extensions).await;
match &outcome {
Err(err) if failed_to_connect(err) => {}
_ => return outcome,
}
let (Some(host), Some(request)) = (host, replay) else {
return outcome;
};
resolver.invalidate_stale(&host);
next.run(request, extensions).await
}
}