use super::trait_::{Governor, GovernorError};
use crate::error::OpsError;
use async_trait::async_trait;
use http::Extensions;
use reqwest_middleware::{ClientBuilder, ClientWithMiddleware, Middleware, Next};
use std::collections::HashSet;
use std::sync::Arc;
use std::time::Duration;
const REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
const CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
pub fn governed_client(
governor: Arc<dyn Governor>,
hosts: &[&str],
) -> Result<ClientWithMiddleware, OpsError> {
let base = reqwest::Client::builder()
.timeout(REQUEST_TIMEOUT)
.connect_timeout(CONNECT_TIMEOUT)
.build()
.map_err(|e| OpsError::Unavailable {
component: format!("reqwest::Client::builder: {e}"),
})?;
let allowlist: HashSet<String> = hosts.iter().map(|h| h.to_ascii_lowercase()).collect();
let middleware = GovernorMiddleware {
governor,
allowlist,
};
Ok(ClientBuilder::new(base).with(middleware).build())
}
struct GovernorMiddleware {
governor: Arc<dyn Governor>,
allowlist: HashSet<String>,
}
#[async_trait]
impl Middleware for GovernorMiddleware {
async fn handle(
&self,
req: reqwest::Request,
extensions: &mut Extensions,
next: Next<'_>,
) -> reqwest_middleware::Result<reqwest::Response> {
let host = req
.url()
.host_str()
.ok_or_else(|| {
reqwest_middleware::Error::Middleware(anyhow::anyhow!(
"governed_client: URL has no host: {}",
req.url()
))
})?
.to_string();
if !self.allowlist.contains(&host) {
return Err(reqwest_middleware::Error::Middleware(anyhow::anyhow!(
"governed_client: host `{host}` not in allowlist"
)));
}
let _permit = self
.governor
.acquire_egress(host.clone())
.await
.map_err(governor_err_to_mw)?;
next.run(req, extensions).await
}
}
fn governor_err_to_mw(e: GovernorError) -> reqwest_middleware::Error {
reqwest_middleware::Error::Middleware(anyhow::anyhow!(
"governed_client: governor refused egress: {e}"
))
}