#![doc = include_str!("../README.md")]
#![warn(missing_docs)]
use std::future;
use std::net::{SocketAddr, ToSocketAddrs};
use std::str::FromStr;
use std::sync::Arc;
use anyhow::anyhow;
use bytes::Bytes;
use http::Extensions;
use http::header::{HeaderName, HeaderValue};
use http_acl::utils::authority::{Authority, Host};
use reqwest::{
Body, Request, Response, ResponseBuilderExt,
dns::{Name, Resolve, Resolving},
redirect,
};
use reqwest_middleware::{Error, Middleware, Next};
use thiserror::Error;
pub use http_acl::{
self, HttpAcl, HttpAclBuilder, HttpAclHooks, ModifyRequestFn, ModifyResponseFn,
RequestMutation, ResponseMutation, ValidateFn,
};
#[derive(Debug, Clone)]
pub struct HttpAclMiddleware {
acl: Arc<HttpAcl>,
}
impl HttpAclMiddleware {
pub fn new(acl: HttpAcl) -> Self {
Self { acl: Arc::new(acl) }
}
pub fn acl(&self) -> Arc<HttpAcl> {
self.acl.clone()
}
pub fn dns_resolver(&self) -> Arc<HttpAclDnsResolver> {
Arc::new(HttpAclDnsResolver::new(self))
}
pub fn with_dns_resolver(&self, dns_resolver: Arc<dyn Resolve>) -> Arc<HttpAclDnsResolver> {
Arc::new(HttpAclDnsResolver::with_dns_resolver(self, dns_resolver))
}
pub fn redirect_policy(&self) -> redirect::Policy {
self.redirect_policy_with_max(10)
}
pub fn redirect_policy_with_max(&self, max_redirects: usize) -> redirect::Policy {
let acl = self.acl.clone();
redirect::Policy::custom(move |attempt| {
let deny_reason = 'reason: {
if attempt.previous().len() > max_redirects {
break 'reason Some("too many redirects".to_string());
}
let url = attempt.url();
let scheme = url.scheme();
if acl.is_scheme_allowed(scheme).is_denied() {
break 'reason Some(format!("scheme {scheme} is denied"));
}
let Some(host) = url.host() else {
break 'reason Some("missing host".to_string());
};
match host {
url::Host::Domain(domain) => {
if acl.is_host_allowed(domain).is_denied() {
break 'reason Some(format!("host {domain} is denied"));
}
}
url::Host::Ipv4(ip) => {
let ip = std::net::IpAddr::V4(ip);
if acl.is_ip_allowed(&ip).is_denied() {
break 'reason Some(format!("ip {ip} is denied"));
}
}
url::Host::Ipv6(ip) => {
let ip = std::net::IpAddr::V6(ip);
if acl.is_ip_allowed(&ip).is_denied() {
break 'reason Some(format!("ip {ip} is denied"));
}
}
}
if let Some(port) = url.port_or_known_default()
&& acl.is_port_allowed(port).is_denied()
{
break 'reason Some(format!("port {port} is denied"));
}
match percent_encoding::percent_decode_str(url.path()).decode_utf8() {
Ok(path) => {
if acl.is_url_path_allowed(&path).is_denied() {
break 'reason Some(format!("path {path} is denied"));
}
}
Err(_) => break 'reason Some("invalid URL path encoding".to_string()),
}
None
};
match deny_reason {
Some(reason) => attempt.error(std::io::Error::other(reason)),
None => attempt.follow(),
}
})
}
}
#[async_trait::async_trait]
impl Middleware for HttpAclMiddleware {
async fn handle(
&self,
mut req: Request,
extensions: &mut Extensions,
next: Next<'_>,
) -> std::result::Result<Response, Error> {
let scheme = req.url().scheme().to_string();
let acl_scheme_match = self.acl.is_scheme_allowed(&scheme);
if acl_scheme_match.is_denied() {
return Err(Error::Middleware(anyhow!(
"scheme {} is denied - {}",
scheme,
acl_scheme_match
)));
}
let method = req.method().as_str();
let acl_method_match = self.acl.is_method_allowed(method);
if acl_method_match.is_denied() {
return Err(Error::Middleware(anyhow!(
"method {} is denied - {}",
method,
acl_method_match
)));
}
if let Some(host) = req.url().host_str() {
let authority = Authority::parse(host)
.map_err(|_| Error::Middleware(anyhow!("invalid host: {}", host)))?;
match &authority.host {
Host::Ip(ip) => {
let acl_ip_match = self.acl.is_ip_allowed(ip);
if acl_ip_match.is_denied() {
return Err(Error::Middleware(anyhow!(
"ip {} is denied - {}",
ip,
acl_ip_match
)));
}
}
Host::Domain(domain) => {
let acl_host_match = self.acl.is_host_allowed(domain);
if acl_host_match.is_denied() {
return Err(Error::Middleware(anyhow!(
"host {} is denied - {}",
domain,
acl_host_match
)));
}
}
}
if let Some(port) = req.url().port_or_known_default() {
let acl_port_match = self.acl.is_port_allowed(port);
if acl_port_match.is_denied() {
return Err(Error::Middleware(anyhow!(
"port {} is denied - {}",
port,
acl_port_match
)));
}
}
for (key, value) in req.headers() {
let header_name = key.as_str();
let header_value = value.to_str().map_err(|_| {
Error::Middleware(anyhow!("invalid header value for {}", header_name))
})?;
let acl_header_match = self.acl.is_header_allowed(header_name, header_value);
if acl_header_match.is_denied() {
return Err(Error::Middleware(anyhow!(
"header {}: {} is denied - {}",
header_name,
header_value,
acl_header_match
)));
}
}
let url_path = percent_encoding::percent_decode_str(req.url().path())
.decode_utf8()
.map_err(|_| Error::Middleware(anyhow!("invalid URL path encoding")))?;
let acl_url_path_match = self.acl.is_url_path_allowed(&url_path);
if acl_url_path_match.is_denied() {
return Err(Error::Middleware(anyhow!(
"path {} is denied - {}",
url_path,
acl_url_path_match
)));
}
let valid_match = self.acl.is_valid(
&scheme,
&authority,
req.headers()
.iter()
.filter_map(|(k, v)| Some((k.as_str(), v.to_str().ok()?))),
req.body().and_then(|b| b.as_bytes()),
);
if valid_match.is_denied() {
return Err(Error::Middleware(anyhow!(
"request is denied - {}",
valid_match
)));
}
if self.acl.has_modify_request() {
let mut mutation = RequestMutation {
headers: req
.headers()
.iter()
.filter_map(|(k, v)| {
Some((k.as_str().to_string(), v.to_str().ok()?.to_string()))
})
.collect(),
body: req
.body()
.and_then(|b| b.as_bytes())
.map(Bytes::copy_from_slice),
};
self.acl.modify_request(&scheme, &authority, &mut mutation);
req.headers_mut().clear();
for (name, value) in &mutation.headers {
let header_name = HeaderName::from_str(name).map_err(|e| {
Error::Middleware(anyhow!("invalid header name `{name}`: {e}"))
})?;
let header_value = HeaderValue::from_str(value).map_err(|e| {
Error::Middleware(anyhow!("invalid header value for `{name}`: {e}"))
})?;
req.headers_mut().append(header_name, header_value);
}
if let Some(body) = mutation.body {
*req.body_mut() = Some(Body::from(body));
}
}
let mut res = next.run(req, extensions).await?;
if self.acl.has_modify_response() {
let status = res.status();
let version = res.version();
let url = res.url().clone();
let extensions_snapshot = res.extensions().clone();
let headers: Vec<(String, String)> = res
.headers()
.iter()
.map(|(k, v)| {
(
k.as_str().to_string(),
String::from_utf8_lossy(v.as_bytes()).into_owned(),
)
})
.collect();
let body = res.bytes().await?;
let mut mutation = ResponseMutation {
status: status.as_u16(),
headers,
body,
};
self.acl.modify_response(&scheme, &authority, &mut mutation);
let mut builder = http::Response::builder()
.status(mutation.status)
.version(version);
for (name, value) in &mutation.headers {
builder = builder.header(name.as_str(), value.as_str());
}
if let Some(ext) = builder.extensions_mut() {
*ext = extensions_snapshot;
}
let http_response = builder
.url(url)
.body(mutation.body)
.map_err(|e| Error::Middleware(anyhow!("failed to rebuild response: {e}")))?;
res = Response::from(http_response);
}
Ok(res)
} else {
return Err(Error::Middleware(anyhow!("missing host")));
}
}
}
type BoxError = Box<dyn std::error::Error + Send + Sync>;
struct GaiResolver;
impl Resolve for GaiResolver {
fn resolve(&self, name: Name) -> Resolving {
Box::pin(async move {
let addresses = (name.as_str(), 0)
.to_socket_addrs()
.map_err(|e| Box::new(e) as BoxError)?;
Ok(Box::new(addresses.into_iter()) as Box<dyn Iterator<Item = SocketAddr> + Send>)
})
}
}
pub struct HttpAclDnsResolver {
dns_resolver: Arc<dyn Resolve>,
acl: Arc<HttpAcl>,
}
impl HttpAclDnsResolver {
pub fn new(middleware: &HttpAclMiddleware) -> Self {
Self {
dns_resolver: Arc::new(GaiResolver),
acl: middleware.acl(),
}
}
pub fn with_dns_resolver(
middleware: &HttpAclMiddleware,
dns_resolver: Arc<dyn Resolve>,
) -> Self {
Self {
dns_resolver,
acl: middleware.acl(),
}
}
}
impl Resolve for HttpAclDnsResolver {
fn resolve(&self, name: Name) -> Resolving {
if self.acl.is_host_allowed(name.as_str()).is_denied() {
let err: BoxError = Box::new(HttpAclError::HostDenied {
host: name.as_str().to_string(),
});
return Box::pin(future::ready(Err(err)));
}
let acl = self.acl.clone();
let resolver = self.dns_resolver.clone();
Box::pin(async move {
if let Some(tcp_address) = acl.resolve_trusted_static_dns_mapping(name.as_str()) {
Ok(Box::new(std::iter::once(tcp_address))
as Box<dyn Iterator<Item = SocketAddr> + Send>)
} else if let Some(tcp_address) = acl.resolve_static_dns_mapping(name.as_str()) {
if acl.is_ip_allowed(&tcp_address.ip()).is_allowed()
&& acl.is_port_allowed(tcp_address.port()).is_allowed()
{
Ok(Box::new(std::iter::once(tcp_address))
as Box<dyn Iterator<Item = SocketAddr> + Send>)
} else {
let err: BoxError =
Box::new(std::io::Error::other("Static DNS mapping denied by ACL"));
Err(err)
}
} else {
let resolved = resolver.resolve(name).await;
match resolved {
Ok(addresses) => {
let filtered = addresses
.into_iter()
.filter(|addr| {
acl.is_ip_allowed(&addr.ip()).is_allowed()
&& acl.is_port_allowed(addr.port()).is_allowed()
})
.collect::<Vec<_>>();
Ok(Box::new(filtered.into_iter())
as Box<dyn Iterator<Item = SocketAddr> + Send>)
}
Err(e) => Err(e),
}
}
})
}
}
#[derive(Error, Debug)]
pub enum HttpAclError {
#[error("Host resolution denied by ACL: {host}")]
HostDenied {
host: String,
},
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_http_acl_middleware() {
let acl = HttpAcl::builder()
.add_denied_host("example.com".to_string())
.unwrap()
.build();
let middleware = HttpAclMiddleware::new(acl);
let client = reqwest_middleware::ClientBuilder::new(
reqwest::Client::builder()
.dns_resolver(middleware.dns_resolver())
.build()
.unwrap(),
)
.with(middleware)
.build();
let request = client.get("http://example.com/").send().await;
assert!(request.is_err());
assert_eq!(
request.unwrap_err().to_string(),
"host example.com is denied - The entity is denied according to the denied ACL."
);
}
#[tokio::test]
async fn test_middleware_decodes_percent_encoded_path() {
let acl = HttpAcl::builder()
.add_allowed_host("example.com".to_string())
.unwrap()
.add_denied_url_path("/secret file".to_string())
.unwrap()
.build();
let middleware = HttpAclMiddleware::new(acl);
let client =
reqwest_middleware::ClientBuilder::new(reqwest::Client::builder().build().unwrap())
.with(middleware)
.build();
let request = client.get("http://example.com/secret%20file").send().await;
assert!(request.is_err());
assert!(
request
.unwrap_err()
.to_string()
.contains("path /secret file is denied")
);
}
#[tokio::test]
async fn test_dns_resolver_returns_typed_error_for_denied_host() {
let acl = HttpAcl::builder()
.add_denied_host("denied.example.com".to_string())
.unwrap()
.build();
let middleware = HttpAclMiddleware::new(acl);
let resolver = middleware.dns_resolver();
let name: reqwest::dns::Name = "denied.example.com".parse().unwrap();
let err = match resolver.resolve(name).await {
Ok(_) => panic!("expected resolution to be denied"),
Err(e) => e,
};
let acl_err = err
.downcast_ref::<HttpAclError>()
.expect("expected a HttpAclError");
assert!(matches!(
acl_err,
HttpAclError::HostDenied { host } if host == "denied.example.com"
));
}
#[tokio::test]
async fn test_dns_resolver_resolves_hostnames() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
if let Ok((mut socket, _)) = listener.accept().await {
let mut buf = [0u8; 1024];
let _ = socket.read(&mut buf).await;
let response = "HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n";
let _ = socket.write_all(response.as_bytes()).await;
}
});
let acl = HttpAcl::builder()
.non_global_ip_ranges(true)
.ip_acl_default(true)
.port_acl_default(true)
.host_acl_default(true)
.build();
let middleware = HttpAclMiddleware::new(acl);
let client = reqwest_middleware::ClientBuilder::new(
reqwest::Client::builder()
.dns_resolver(middleware.dns_resolver())
.build()
.unwrap(),
)
.with(middleware)
.build();
let request = client
.get(format!("http://localhost:{}/", addr.port()))
.send()
.await;
assert!(request.is_ok(), "{:?}", request.err());
}
#[tokio::test]
async fn test_trusted_static_dns_mapping_bypasses_ip_port_acl() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
if let Ok((mut socket, _)) = listener.accept().await {
let mut buf = [0u8; 1024];
let _ = socket.read(&mut buf).await;
let response = "HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n";
let _ = socket.write_all(response.as_bytes()).await;
}
});
let acl = HttpAcl::builder()
.host_acl_default(true)
.add_trusted_static_dns_mapping("trusted.internal".to_string(), addr)
.unwrap()
.build();
assert!(acl.is_ip_allowed(&addr.ip()).is_denied());
assert!(acl.is_port_allowed(addr.port()).is_denied());
let middleware = HttpAclMiddleware::new(acl);
let client = reqwest_middleware::ClientBuilder::new(
reqwest::Client::builder()
.dns_resolver(middleware.dns_resolver())
.build()
.unwrap(),
)
.with(middleware)
.build();
let request = client.get("http://trusted.internal/").send().await;
assert!(request.is_ok(), "{:?}", request.err());
}
#[tokio::test]
async fn test_redirect_policy_blocks_disallowed_target() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
if let Ok((mut socket, _)) = listener.accept().await {
let mut buf = [0u8; 1024];
let _ = socket.read(&mut buf).await;
let response = "HTTP/1.1 302 Found\r\nLocation: http://192.168.1.1/\r\nContent-Length: 0\r\n\r\n";
let _ = socket.write_all(response.as_bytes()).await;
}
});
let acl = HttpAcl::builder()
.non_global_ip_ranges(true)
.ip_acl_default(true)
.port_acl_default(true)
.add_denied_ip_range((
"192.168.1.1".parse::<std::net::IpAddr>().unwrap(),
"192.168.1.1".parse::<std::net::IpAddr>().unwrap(),
))
.unwrap()
.build();
let middleware = HttpAclMiddleware::new(acl);
let client = reqwest_middleware::ClientBuilder::new(
reqwest::Client::builder()
.dns_resolver(middleware.dns_resolver())
.redirect(middleware.redirect_policy())
.build()
.unwrap(),
)
.with(middleware)
.build();
let request = client
.get(format!("http://127.0.0.1:{}/", addr.port()))
.send()
.await;
assert!(request.is_err());
}
#[test]
fn test_no_hooks_configured_leaves_acl_unaffected() {
let acl = HttpAcl::builder().build();
assert!(!acl.has_modify_request());
assert!(!acl.has_modify_response());
}
#[tokio::test]
async fn test_modify_request_injects_header() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let (captured_tx, captured_rx) = tokio::sync::oneshot::channel();
tokio::spawn(async move {
if let Ok((mut socket, _)) = listener.accept().await {
let mut buf = [0u8; 4096];
let n = socket.read(&mut buf).await.unwrap_or(0);
let _ = captured_tx.send(String::from_utf8_lossy(&buf[..n]).into_owned());
let response = "HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n";
let _ = socket.write_all(response.as_bytes()).await;
}
});
let acl = HttpAcl::builder()
.non_global_ip_ranges(true)
.ip_acl_default(true)
.port_acl_default(true)
.host_acl_default(true)
.build_full(HttpAclHooks {
modify_request_fn: Some(Arc::new(|_scheme, _authority, mutation| {
mutation
.headers
.push(("x-injected-secret".to_string(), "sssh".to_string()));
})),
..Default::default()
});
let middleware = HttpAclMiddleware::new(acl);
let client = reqwest_middleware::ClientBuilder::new(
reqwest::Client::builder()
.dns_resolver(middleware.dns_resolver())
.build()
.unwrap(),
)
.with(middleware)
.build();
let request = client
.get(format!("http://127.0.0.1:{}/", addr.port()))
.send()
.await;
assert!(request.is_ok(), "{:?}", request.err());
let captured = captured_rx.await.unwrap();
assert!(
captured.to_lowercase().contains("x-injected-secret: sssh"),
"captured request did not contain the injected header:\n{captured}"
);
}
#[tokio::test]
async fn test_modify_request_replaces_body_and_content_length() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let (captured_tx, captured_rx) = tokio::sync::oneshot::channel();
tokio::spawn(async move {
if let Ok((mut socket, _)) = listener.accept().await {
let mut buf = [0u8; 4096];
let n = socket.read(&mut buf).await.unwrap_or(0);
let _ = captured_tx.send(String::from_utf8_lossy(&buf[..n]).into_owned());
let response = "HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n";
let _ = socket.write_all(response.as_bytes()).await;
}
});
let new_body = "a much longer replacement body than the original";
let acl = HttpAcl::builder()
.non_global_ip_ranges(true)
.ip_acl_default(true)
.port_acl_default(true)
.host_acl_default(true)
.build_full(HttpAclHooks {
modify_request_fn: Some(Arc::new(move |_scheme, _authority, mutation| {
mutation.body = Some(Bytes::from(new_body));
})),
..Default::default()
});
let middleware = HttpAclMiddleware::new(acl);
let client = reqwest_middleware::ClientBuilder::new(
reqwest::Client::builder()
.dns_resolver(middleware.dns_resolver())
.build()
.unwrap(),
)
.with(middleware)
.build();
let request = client
.post(format!("http://127.0.0.1:{}/", addr.port()))
.body("short")
.send()
.await;
assert!(request.is_ok(), "{:?}", request.err());
let captured = captured_rx.await.unwrap();
assert!(
captured.contains(&format!("content-length: {}", new_body.len()))
|| captured.contains(&format!("Content-Length: {}", new_body.len())),
"captured request did not have a Content-Length matching the replaced body:\n{captured}"
);
assert!(
captured.ends_with(new_body),
"captured request did not end with the replaced body:\n{captured}"
);
}
#[tokio::test]
async fn test_modify_response_rewrites_status_headers_and_body() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
if let Ok((mut socket, _)) = listener.accept().await {
let mut buf = [0u8; 1024];
let _ = socket.read(&mut buf).await;
let body = "original body";
let response = format!(
"HTTP/1.1 200 OK\r\nx-remove-me: yes\r\nContent-Length: {}\r\n\r\n{}",
body.len(),
body
);
let _ = socket.write_all(response.as_bytes()).await;
}
});
let acl = HttpAcl::builder()
.non_global_ip_ranges(true)
.ip_acl_default(true)
.port_acl_default(true)
.host_acl_default(true)
.build_full(HttpAclHooks {
modify_response_fn: Some(Arc::new(|_scheme, _authority, mutation| {
mutation.status = 201;
mutation.headers.retain(|(k, _)| k != "x-remove-me");
mutation
.headers
.push(("x-added".to_string(), "yes".to_string()));
mutation.body = Bytes::from_static(b"redacted body");
})),
..Default::default()
});
let middleware = HttpAclMiddleware::new(acl);
let client = reqwest_middleware::ClientBuilder::new(
reqwest::Client::builder()
.dns_resolver(middleware.dns_resolver())
.build()
.unwrap(),
)
.with(middleware)
.build();
let response = client
.get(format!("http://127.0.0.1:{}/", addr.port()))
.send()
.await
.unwrap();
assert_eq!(response.status(), 201);
assert!(!response.headers().contains_key("x-remove-me"));
assert_eq!(response.headers().get("x-added").unwrap(), "yes");
let body = response.text().await.unwrap();
assert_eq!(body, "redacted body");
}
#[tokio::test]
async fn test_modify_response_preserves_url_and_remote_addr() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
if let Ok((mut socket, _)) = listener.accept().await {
let mut buf = [0u8; 1024];
let _ = socket.read(&mut buf).await;
let response = "HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok";
let _ = socket.write_all(response.as_bytes()).await;
}
});
let acl = HttpAcl::builder()
.non_global_ip_ranges(true)
.ip_acl_default(true)
.port_acl_default(true)
.host_acl_default(true)
.build_full(HttpAclHooks {
modify_response_fn: Some(Arc::new(|_scheme, _authority, mutation| {
mutation.body = Bytes::from_static(b"changed");
})),
..Default::default()
});
let middleware = HttpAclMiddleware::new(acl);
let client = reqwest_middleware::ClientBuilder::new(
reqwest::Client::builder()
.dns_resolver(middleware.dns_resolver())
.build()
.unwrap(),
)
.with(middleware)
.build();
let url = format!("http://127.0.0.1:{}/", addr.port());
let response = client.get(&url).send().await.unwrap();
assert_eq!(response.url().as_str(), url);
assert!(
response.remote_addr().is_some(),
"remote_addr() was lost when the response was rebuilt"
);
}
#[tokio::test]
async fn test_modify_request_header_carries_through_redirect() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
let second_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let second_addr = second_listener.local_addr().unwrap();
let (second_captured_tx, second_captured_rx) = tokio::sync::oneshot::channel();
tokio::spawn(async move {
if let Ok((mut socket, _)) = second_listener.accept().await {
let mut buf = [0u8; 4096];
let n = socket.read(&mut buf).await.unwrap_or(0);
let _ = second_captured_tx.send(String::from_utf8_lossy(&buf[..n]).into_owned());
let response = "HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n";
let _ = socket.write_all(response.as_bytes()).await;
}
});
let first_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let first_addr = first_listener.local_addr().unwrap();
tokio::spawn(async move {
if let Ok((mut socket, _)) = first_listener.accept().await {
let mut buf = [0u8; 1024];
let _ = socket.read(&mut buf).await;
let response = format!(
"HTTP/1.1 302 Found\r\nLocation: http://127.0.0.1:{}/\r\nContent-Length: 0\r\n\r\n",
second_addr.port()
);
let _ = socket.write_all(response.as_bytes()).await;
}
});
let acl = HttpAcl::builder()
.non_global_ip_ranges(true)
.ip_acl_default(true)
.port_acl_default(true)
.host_acl_default(true)
.build_full(HttpAclHooks {
modify_request_fn: Some(Arc::new(|_scheme, _authority, mutation| {
mutation
.headers
.push(("x-marker".to_string(), "hop-one-only".to_string()));
})),
..Default::default()
});
let middleware = HttpAclMiddleware::new(acl);
let client = reqwest_middleware::ClientBuilder::new(
reqwest::Client::builder()
.dns_resolver(middleware.dns_resolver())
.redirect(middleware.redirect_policy())
.build()
.unwrap(),
)
.with(middleware)
.build();
let request = client
.get(format!("http://127.0.0.1:{}/", first_addr.port()))
.send()
.await;
assert!(request.is_ok(), "{:?}", request.err());
let second_captured = second_captured_rx.await.unwrap();
assert!(
second_captured.to_lowercase().contains("x-marker"),
"expected the header injected for the original request to carry through \
reqwest's redirect handling to the second hop, but it did not:\n{second_captured}"
);
}
}