1#![doc = include_str!("../README.md")]
2#![warn(missing_docs)]
3
4use std::future;
5use std::net::{SocketAddr, ToSocketAddrs};
6use std::str::FromStr;
7use std::sync::Arc;
8
9use anyhow::anyhow;
10use bytes::Bytes;
11use http::Extensions;
12use http::header::{HeaderName, HeaderValue};
13use http_acl::utils::authority::{Authority, Host};
14use reqwest::{
15 Body, Request, Response, ResponseBuilderExt,
16 dns::{Name, Resolve, Resolving},
17 redirect,
18};
19use reqwest_middleware::{Error, Middleware, Next};
20use thiserror::Error;
21
22pub use http_acl::{
23 self, HttpAcl, HttpAclBuilder, HttpAclHooks, ModifyRequestFn, ModifyResponseFn,
24 RequestMutation, ResponseMutation, ValidateFn,
25};
26
27#[derive(Debug, Clone)]
28pub struct HttpAclMiddleware {
47 acl: Arc<HttpAcl>,
48}
49
50impl HttpAclMiddleware {
51 pub fn new(acl: HttpAcl) -> Self {
53 Self { acl: Arc::new(acl) }
54 }
55
56 pub fn acl(&self) -> Arc<HttpAcl> {
58 self.acl.clone()
59 }
60
61 pub fn dns_resolver(&self) -> Arc<HttpAclDnsResolver> {
69 Arc::new(HttpAclDnsResolver::new(self))
70 }
71
72 pub fn with_dns_resolver(&self, dns_resolver: Arc<dyn Resolve>) -> Arc<HttpAclDnsResolver> {
75 Arc::new(HttpAclDnsResolver::with_dns_resolver(self, dns_resolver))
76 }
77
78 pub fn redirect_policy(&self) -> redirect::Policy {
119 self.redirect_policy_with_max(10)
120 }
121
122 pub fn redirect_policy_with_max(&self, max_redirects: usize) -> redirect::Policy {
124 let acl = self.acl.clone();
125 redirect::Policy::custom(move |attempt| {
126 let deny_reason = 'reason: {
130 if attempt.previous().len() > max_redirects {
131 break 'reason Some("too many redirects".to_string());
132 }
133
134 let url = attempt.url();
135
136 let scheme = url.scheme();
137 if acl.is_scheme_allowed(scheme).is_denied() {
138 break 'reason Some(format!("scheme {scheme} is denied"));
139 }
140
141 let Some(host) = url.host() else {
142 break 'reason Some("missing host".to_string());
143 };
144
145 match host {
146 url::Host::Domain(domain) => {
147 if acl.is_host_allowed(domain).is_denied() {
148 break 'reason Some(format!("host {domain} is denied"));
149 }
150 }
151 url::Host::Ipv4(ip) => {
152 let ip = std::net::IpAddr::V4(ip);
153 if acl.is_ip_allowed(&ip).is_denied() {
154 break 'reason Some(format!("ip {ip} is denied"));
155 }
156 }
157 url::Host::Ipv6(ip) => {
158 let ip = std::net::IpAddr::V6(ip);
159 if acl.is_ip_allowed(&ip).is_denied() {
160 break 'reason Some(format!("ip {ip} is denied"));
161 }
162 }
163 }
164
165 if let Some(port) = url.port_or_known_default()
166 && acl.is_port_allowed(port).is_denied()
167 {
168 break 'reason Some(format!("port {port} is denied"));
169 }
170
171 match percent_encoding::percent_decode_str(url.path()).decode_utf8() {
174 Ok(path) => {
175 if acl.is_url_path_allowed(&path).is_denied() {
176 break 'reason Some(format!("path {path} is denied"));
177 }
178 }
179 Err(_) => break 'reason Some("invalid URL path encoding".to_string()),
180 }
181
182 None
183 };
184
185 match deny_reason {
186 Some(reason) => attempt.error(std::io::Error::other(reason)),
187 None => attempt.follow(),
188 }
189 })
190 }
191}
192
193#[async_trait::async_trait]
194impl Middleware for HttpAclMiddleware {
195 async fn handle(
196 &self,
197 mut req: Request,
198 extensions: &mut Extensions,
199 next: Next<'_>,
200 ) -> std::result::Result<Response, Error> {
201 let scheme = req.url().scheme().to_string();
205 let acl_scheme_match = self.acl.is_scheme_allowed(&scheme);
206 if acl_scheme_match.is_denied() {
207 return Err(Error::Middleware(anyhow!(
208 "scheme {} is denied - {}",
209 scheme,
210 acl_scheme_match
211 )));
212 }
213
214 let method = req.method().as_str();
215 let acl_method_match = self.acl.is_method_allowed(method);
216 if acl_method_match.is_denied() {
217 return Err(Error::Middleware(anyhow!(
218 "method {} is denied - {}",
219 method,
220 acl_method_match
221 )));
222 }
223
224 if let Some(host) = req.url().host_str() {
225 let authority = Authority::parse(host)
226 .map_err(|_| Error::Middleware(anyhow!("invalid host: {}", host)))?;
227
228 match &authority.host {
229 Host::Ip(ip) => {
230 let acl_ip_match = self.acl.is_ip_allowed(ip);
231 if acl_ip_match.is_denied() {
232 return Err(Error::Middleware(anyhow!(
233 "ip {} is denied - {}",
234 ip,
235 acl_ip_match
236 )));
237 }
238 }
239 Host::Domain(domain) => {
240 let acl_host_match = self.acl.is_host_allowed(domain);
241 if acl_host_match.is_denied() {
242 return Err(Error::Middleware(anyhow!(
243 "host {} is denied - {}",
244 domain,
245 acl_host_match
246 )));
247 }
248 }
249 }
250
251 if let Some(port) = req.url().port_or_known_default() {
252 let acl_port_match = self.acl.is_port_allowed(port);
253 if acl_port_match.is_denied() {
254 return Err(Error::Middleware(anyhow!(
255 "port {} is denied - {}",
256 port,
257 acl_port_match
258 )));
259 }
260 }
261
262 for (key, value) in req.headers() {
263 let header_name = key.as_str();
264 let header_value = value.to_str().map_err(|_| {
265 Error::Middleware(anyhow!("invalid header value for {}", header_name))
266 })?;
267 let acl_header_match = self.acl.is_header_allowed(header_name, header_value);
268 if acl_header_match.is_denied() {
269 return Err(Error::Middleware(anyhow!(
270 "header {}: {} is denied - {}",
271 header_name,
272 header_value,
273 acl_header_match
274 )));
275 }
276 }
277
278 let url_path = percent_encoding::percent_decode_str(req.url().path())
281 .decode_utf8()
282 .map_err(|_| Error::Middleware(anyhow!("invalid URL path encoding")))?;
283 let acl_url_path_match = self.acl.is_url_path_allowed(&url_path);
284 if acl_url_path_match.is_denied() {
285 return Err(Error::Middleware(anyhow!(
286 "path {} is denied - {}",
287 url_path,
288 acl_url_path_match
289 )));
290 }
291
292 let valid_match = self.acl.is_valid(
293 &scheme,
294 &authority,
295 req.headers()
296 .iter()
297 .filter_map(|(k, v)| Some((k.as_str(), v.to_str().ok()?))),
298 req.body().and_then(|b| b.as_bytes()),
299 );
300 if valid_match.is_denied() {
301 return Err(Error::Middleware(anyhow!(
302 "request is denied - {}",
303 valid_match
304 )));
305 }
306
307 if self.acl.has_modify_request() {
308 let mut mutation = RequestMutation {
309 headers: req
310 .headers()
311 .iter()
312 .filter_map(|(k, v)| {
313 Some((k.as_str().to_string(), v.to_str().ok()?.to_string()))
314 })
315 .collect(),
316 body: req
317 .body()
318 .and_then(|b| b.as_bytes())
319 .map(Bytes::copy_from_slice),
320 };
321 self.acl.modify_request(&scheme, &authority, &mut mutation);
322
323 req.headers_mut().clear();
324 for (name, value) in &mutation.headers {
325 let header_name = HeaderName::from_str(name).map_err(|e| {
326 Error::Middleware(anyhow!("invalid header name `{name}`: {e}"))
327 })?;
328 let header_value = HeaderValue::from_str(value).map_err(|e| {
329 Error::Middleware(anyhow!("invalid header value for `{name}`: {e}"))
330 })?;
331 req.headers_mut().append(header_name, header_value);
332 }
333 if let Some(body) = mutation.body {
334 *req.body_mut() = Some(Body::from(body));
335 }
336 }
337
338 let mut res = next.run(req, extensions).await?;
339
340 if self.acl.has_modify_response() {
341 let status = res.status();
342 let version = res.version();
343 let url = res.url().clone();
344 let extensions_snapshot = res.extensions().clone();
345 let headers: Vec<(String, String)> = res
346 .headers()
347 .iter()
348 .map(|(k, v)| {
349 (
350 k.as_str().to_string(),
351 String::from_utf8_lossy(v.as_bytes()).into_owned(),
352 )
353 })
354 .collect();
355 let body = res.bytes().await?;
358
359 let mut mutation = ResponseMutation {
360 status: status.as_u16(),
361 headers,
362 body,
363 };
364 self.acl.modify_response(&scheme, &authority, &mut mutation);
365
366 let mut builder = http::Response::builder()
367 .status(mutation.status)
368 .version(version);
369 for (name, value) in &mutation.headers {
370 builder = builder.header(name.as_str(), value.as_str());
371 }
372 if let Some(ext) = builder.extensions_mut() {
379 *ext = extensions_snapshot;
380 }
381 let http_response = builder
382 .url(url)
383 .body(mutation.body)
384 .map_err(|e| Error::Middleware(anyhow!("failed to rebuild response: {e}")))?;
385 res = Response::from(http_response);
386 }
387
388 Ok(res)
389 } else {
390 return Err(Error::Middleware(anyhow!("missing host")));
391 }
392 }
393}
394
395type BoxError = Box<dyn std::error::Error + Send + Sync>;
396
397struct GaiResolver;
398
399impl Resolve for GaiResolver {
400 fn resolve(&self, name: Name) -> Resolving {
401 Box::pin(async move {
402 let addresses = (name.as_str(), 0)
405 .to_socket_addrs()
406 .map_err(|e| Box::new(e) as BoxError)?;
407 Ok(Box::new(addresses.into_iter()) as Box<dyn Iterator<Item = SocketAddr> + Send>)
408 })
409 }
410}
411
412pub struct HttpAclDnsResolver {
424 dns_resolver: Arc<dyn Resolve>,
425 acl: Arc<HttpAcl>,
426}
427
428impl HttpAclDnsResolver {
429 pub fn new(middleware: &HttpAclMiddleware) -> Self {
431 Self {
432 dns_resolver: Arc::new(GaiResolver),
433 acl: middleware.acl(),
434 }
435 }
436
437 pub fn with_dns_resolver(
440 middleware: &HttpAclMiddleware,
441 dns_resolver: Arc<dyn Resolve>,
442 ) -> Self {
443 Self {
444 dns_resolver,
445 acl: middleware.acl(),
446 }
447 }
448}
449
450impl Resolve for HttpAclDnsResolver {
451 fn resolve(&self, name: Name) -> Resolving {
452 if self.acl.is_host_allowed(name.as_str()).is_denied() {
453 let err: BoxError = Box::new(HttpAclError::HostDenied {
454 host: name.as_str().to_string(),
455 });
456 return Box::pin(future::ready(Err(err)));
457 }
458
459 let acl = self.acl.clone();
460 let resolver = self.dns_resolver.clone();
461
462 Box::pin(async move {
463 if let Some(tcp_address) = acl.resolve_trusted_static_dns_mapping(name.as_str()) {
464 Ok(Box::new(std::iter::once(tcp_address))
468 as Box<dyn Iterator<Item = SocketAddr> + Send>)
469 } else if let Some(tcp_address) = acl.resolve_static_dns_mapping(name.as_str()) {
470 if acl.is_ip_allowed(&tcp_address.ip()).is_allowed()
475 && acl.is_port_allowed(tcp_address.port()).is_allowed()
476 {
477 Ok(Box::new(std::iter::once(tcp_address))
478 as Box<dyn Iterator<Item = SocketAddr> + Send>)
479 } else {
480 let err: BoxError =
481 Box::new(std::io::Error::other("Static DNS mapping denied by ACL"));
482 Err(err)
483 }
484 } else {
485 let resolved = resolver.resolve(name).await;
486 match resolved {
487 Ok(addresses) => {
488 let filtered = addresses
489 .into_iter()
490 .filter(|addr| {
491 acl.is_ip_allowed(&addr.ip()).is_allowed()
492 && acl.is_port_allowed(addr.port()).is_allowed()
493 })
494 .collect::<Vec<_>>();
495 Ok(Box::new(filtered.into_iter())
496 as Box<dyn Iterator<Item = SocketAddr> + Send>)
497 }
498 Err(e) => Err(e),
499 }
500 }
501 })
502 }
503}
504
505#[derive(Error, Debug)]
506pub enum HttpAclError {
512 #[error("Host resolution denied by ACL: {host}")]
514 HostDenied {
515 host: String,
517 },
518}
519
520#[cfg(test)]
521mod tests {
522 use super::*;
523
524 #[tokio::test]
525 async fn test_http_acl_middleware() {
526 let acl = HttpAcl::builder()
527 .add_denied_host("example.com".to_string())
528 .unwrap()
529 .build();
530
531 let middleware = HttpAclMiddleware::new(acl);
532
533 let client = reqwest_middleware::ClientBuilder::new(
534 reqwest::Client::builder()
535 .dns_resolver(middleware.dns_resolver())
536 .build()
537 .unwrap(),
538 )
539 .with(middleware)
540 .build();
541
542 let request = client.get("http://example.com/").send().await;
543
544 assert!(request.is_err());
545 assert_eq!(
546 request.unwrap_err().to_string(),
547 "host example.com is denied - The entity is denied according to the denied ACL."
548 );
549 }
550
551 #[tokio::test]
552 async fn test_middleware_decodes_percent_encoded_path() {
553 let acl = HttpAcl::builder()
554 .add_allowed_host("example.com".to_string())
555 .unwrap()
556 .add_denied_url_path("/secret file".to_string())
557 .unwrap()
558 .build();
559
560 let middleware = HttpAclMiddleware::new(acl);
561
562 let client =
563 reqwest_middleware::ClientBuilder::new(reqwest::Client::builder().build().unwrap())
564 .with(middleware)
565 .build();
566
567 let request = client.get("http://example.com/secret%20file").send().await;
572
573 assert!(request.is_err());
574 assert!(
575 request
576 .unwrap_err()
577 .to_string()
578 .contains("path /secret file is denied")
579 );
580 }
581
582 #[tokio::test]
583 async fn test_dns_resolver_returns_typed_error_for_denied_host() {
584 let acl = HttpAcl::builder()
585 .add_denied_host("denied.example.com".to_string())
586 .unwrap()
587 .build();
588
589 let middleware = HttpAclMiddleware::new(acl);
590 let resolver = middleware.dns_resolver();
591
592 let name: reqwest::dns::Name = "denied.example.com".parse().unwrap();
593 let err = match resolver.resolve(name).await {
594 Ok(_) => panic!("expected resolution to be denied"),
595 Err(e) => e,
596 };
597
598 let acl_err = err
599 .downcast_ref::<HttpAclError>()
600 .expect("expected a HttpAclError");
601 assert!(matches!(
602 acl_err,
603 HttpAclError::HostDenied { host } if host == "denied.example.com"
604 ));
605 }
606
607 #[tokio::test]
608 async fn test_dns_resolver_resolves_hostnames() {
609 use tokio::io::{AsyncReadExt, AsyncWriteExt};
610 use tokio::net::TcpListener;
611
612 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
613 let addr = listener.local_addr().unwrap();
614
615 tokio::spawn(async move {
616 if let Ok((mut socket, _)) = listener.accept().await {
617 let mut buf = [0u8; 1024];
618 let _ = socket.read(&mut buf).await;
619 let response = "HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n";
620 let _ = socket.write_all(response.as_bytes()).await;
621 }
622 });
623
624 let acl = HttpAcl::builder()
625 .non_global_ip_ranges(true)
626 .ip_acl_default(true)
627 .port_acl_default(true)
628 .host_acl_default(true)
629 .build();
630
631 let middleware = HttpAclMiddleware::new(acl);
632
633 let client = reqwest_middleware::ClientBuilder::new(
634 reqwest::Client::builder()
635 .dns_resolver(middleware.dns_resolver())
636 .build()
637 .unwrap(),
638 )
639 .with(middleware)
640 .build();
641
642 let request = client
645 .get(format!("http://localhost:{}/", addr.port()))
646 .send()
647 .await;
648
649 assert!(request.is_ok(), "{:?}", request.err());
650 }
651
652 #[tokio::test]
653 async fn test_trusted_static_dns_mapping_bypasses_ip_port_acl() {
654 use tokio::io::{AsyncReadExt, AsyncWriteExt};
655 use tokio::net::TcpListener;
656
657 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
658 let addr = listener.local_addr().unwrap();
659
660 tokio::spawn(async move {
661 if let Ok((mut socket, _)) = listener.accept().await {
662 let mut buf = [0u8; 1024];
663 let _ = socket.read(&mut buf).await;
664 let response = "HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n";
665 let _ = socket.write_all(response.as_bytes()).await;
666 }
667 });
668
669 let acl = HttpAcl::builder()
672 .host_acl_default(true)
673 .add_trusted_static_dns_mapping("trusted.internal".to_string(), addr)
674 .unwrap()
675 .build();
676
677 assert!(acl.is_ip_allowed(&addr.ip()).is_denied());
678 assert!(acl.is_port_allowed(addr.port()).is_denied());
679
680 let middleware = HttpAclMiddleware::new(acl);
681
682 let client = reqwest_middleware::ClientBuilder::new(
683 reqwest::Client::builder()
684 .dns_resolver(middleware.dns_resolver())
685 .build()
686 .unwrap(),
687 )
688 .with(middleware)
689 .build();
690
691 let request = client.get("http://trusted.internal/").send().await;
694
695 assert!(request.is_ok(), "{:?}", request.err());
696 }
697
698 #[tokio::test]
699 async fn test_redirect_policy_blocks_disallowed_target() {
700 use tokio::io::{AsyncReadExt, AsyncWriteExt};
701 use tokio::net::TcpListener;
702
703 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
704 let addr = listener.local_addr().unwrap();
705
706 tokio::spawn(async move {
707 if let Ok((mut socket, _)) = listener.accept().await {
708 let mut buf = [0u8; 1024];
709 let _ = socket.read(&mut buf).await;
710 let response = "HTTP/1.1 302 Found\r\nLocation: http://192.168.1.1/\r\nContent-Length: 0\r\n\r\n";
711 let _ = socket.write_all(response.as_bytes()).await;
712 }
713 });
714
715 let acl = HttpAcl::builder()
718 .non_global_ip_ranges(true)
719 .ip_acl_default(true)
720 .port_acl_default(true)
721 .add_denied_ip_range((
722 "192.168.1.1".parse::<std::net::IpAddr>().unwrap(),
723 "192.168.1.1".parse::<std::net::IpAddr>().unwrap(),
724 ))
725 .unwrap()
726 .build();
727
728 let middleware = HttpAclMiddleware::new(acl);
729
730 let client = reqwest_middleware::ClientBuilder::new(
731 reqwest::Client::builder()
732 .dns_resolver(middleware.dns_resolver())
733 .redirect(middleware.redirect_policy())
734 .build()
735 .unwrap(),
736 )
737 .with(middleware)
738 .build();
739
740 let request = client
741 .get(format!("http://127.0.0.1:{}/", addr.port()))
742 .send()
743 .await;
744
745 assert!(request.is_err());
746 }
747
748 #[test]
749 fn test_no_hooks_configured_leaves_acl_unaffected() {
750 let acl = HttpAcl::builder().build();
751
752 assert!(!acl.has_modify_request());
753 assert!(!acl.has_modify_response());
754 }
755
756 #[tokio::test]
757 async fn test_modify_request_injects_header() {
758 use tokio::io::{AsyncReadExt, AsyncWriteExt};
759 use tokio::net::TcpListener;
760
761 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
762 let addr = listener.local_addr().unwrap();
763
764 let (captured_tx, captured_rx) = tokio::sync::oneshot::channel();
765 tokio::spawn(async move {
766 if let Ok((mut socket, _)) = listener.accept().await {
767 let mut buf = [0u8; 4096];
768 let n = socket.read(&mut buf).await.unwrap_or(0);
769 let _ = captured_tx.send(String::from_utf8_lossy(&buf[..n]).into_owned());
770 let response = "HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n";
771 let _ = socket.write_all(response.as_bytes()).await;
772 }
773 });
774
775 let acl = HttpAcl::builder()
776 .non_global_ip_ranges(true)
777 .ip_acl_default(true)
778 .port_acl_default(true)
779 .host_acl_default(true)
780 .build_full(HttpAclHooks {
781 modify_request_fn: Some(Arc::new(|_scheme, _authority, mutation| {
782 mutation
783 .headers
784 .push(("x-injected-secret".to_string(), "sssh".to_string()));
785 })),
786 ..Default::default()
787 });
788
789 let middleware = HttpAclMiddleware::new(acl);
790
791 let client = reqwest_middleware::ClientBuilder::new(
792 reqwest::Client::builder()
793 .dns_resolver(middleware.dns_resolver())
794 .build()
795 .unwrap(),
796 )
797 .with(middleware)
798 .build();
799
800 let request = client
801 .get(format!("http://127.0.0.1:{}/", addr.port()))
802 .send()
803 .await;
804
805 assert!(request.is_ok(), "{:?}", request.err());
806 let captured = captured_rx.await.unwrap();
807 assert!(
808 captured.to_lowercase().contains("x-injected-secret: sssh"),
809 "captured request did not contain the injected header:\n{captured}"
810 );
811 }
812
813 #[tokio::test]
814 async fn test_modify_request_replaces_body_and_content_length() {
815 use tokio::io::{AsyncReadExt, AsyncWriteExt};
816 use tokio::net::TcpListener;
817
818 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
819 let addr = listener.local_addr().unwrap();
820
821 let (captured_tx, captured_rx) = tokio::sync::oneshot::channel();
822 tokio::spawn(async move {
823 if let Ok((mut socket, _)) = listener.accept().await {
824 let mut buf = [0u8; 4096];
825 let n = socket.read(&mut buf).await.unwrap_or(0);
826 let _ = captured_tx.send(String::from_utf8_lossy(&buf[..n]).into_owned());
827 let response = "HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n";
828 let _ = socket.write_all(response.as_bytes()).await;
829 }
830 });
831
832 let new_body = "a much longer replacement body than the original";
833 let acl = HttpAcl::builder()
834 .non_global_ip_ranges(true)
835 .ip_acl_default(true)
836 .port_acl_default(true)
837 .host_acl_default(true)
838 .build_full(HttpAclHooks {
839 modify_request_fn: Some(Arc::new(move |_scheme, _authority, mutation| {
840 mutation.body = Some(Bytes::from(new_body));
841 })),
842 ..Default::default()
843 });
844
845 let middleware = HttpAclMiddleware::new(acl);
846
847 let client = reqwest_middleware::ClientBuilder::new(
848 reqwest::Client::builder()
849 .dns_resolver(middleware.dns_resolver())
850 .build()
851 .unwrap(),
852 )
853 .with(middleware)
854 .build();
855
856 let request = client
857 .post(format!("http://127.0.0.1:{}/", addr.port()))
858 .body("short")
859 .send()
860 .await;
861
862 assert!(request.is_ok(), "{:?}", request.err());
863 let captured = captured_rx.await.unwrap();
864 assert!(
865 captured.contains(&format!("content-length: {}", new_body.len()))
866 || captured.contains(&format!("Content-Length: {}", new_body.len())),
867 "captured request did not have a Content-Length matching the replaced body:\n{captured}"
868 );
869 assert!(
870 captured.ends_with(new_body),
871 "captured request did not end with the replaced body:\n{captured}"
872 );
873 }
874
875 #[tokio::test]
876 async fn test_modify_response_rewrites_status_headers_and_body() {
877 use tokio::io::{AsyncReadExt, AsyncWriteExt};
878 use tokio::net::TcpListener;
879
880 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
881 let addr = listener.local_addr().unwrap();
882
883 tokio::spawn(async move {
884 if let Ok((mut socket, _)) = listener.accept().await {
885 let mut buf = [0u8; 1024];
886 let _ = socket.read(&mut buf).await;
887 let body = "original body";
888 let response = format!(
889 "HTTP/1.1 200 OK\r\nx-remove-me: yes\r\nContent-Length: {}\r\n\r\n{}",
890 body.len(),
891 body
892 );
893 let _ = socket.write_all(response.as_bytes()).await;
894 }
895 });
896
897 let acl = HttpAcl::builder()
898 .non_global_ip_ranges(true)
899 .ip_acl_default(true)
900 .port_acl_default(true)
901 .host_acl_default(true)
902 .build_full(HttpAclHooks {
903 modify_response_fn: Some(Arc::new(|_scheme, _authority, mutation| {
904 mutation.status = 201;
905 mutation.headers.retain(|(k, _)| k != "x-remove-me");
906 mutation
907 .headers
908 .push(("x-added".to_string(), "yes".to_string()));
909 mutation.body = Bytes::from_static(b"redacted body");
910 })),
911 ..Default::default()
912 });
913
914 let middleware = HttpAclMiddleware::new(acl);
915
916 let client = reqwest_middleware::ClientBuilder::new(
917 reqwest::Client::builder()
918 .dns_resolver(middleware.dns_resolver())
919 .build()
920 .unwrap(),
921 )
922 .with(middleware)
923 .build();
924
925 let response = client
926 .get(format!("http://127.0.0.1:{}/", addr.port()))
927 .send()
928 .await
929 .unwrap();
930
931 assert_eq!(response.status(), 201);
932 assert!(!response.headers().contains_key("x-remove-me"));
933 assert_eq!(response.headers().get("x-added").unwrap(), "yes");
934 let body = response.text().await.unwrap();
935 assert_eq!(body, "redacted body");
936 }
937
938 #[tokio::test]
939 async fn test_modify_response_preserves_url_and_remote_addr() {
940 use tokio::io::{AsyncReadExt, AsyncWriteExt};
941 use tokio::net::TcpListener;
942
943 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
944 let addr = listener.local_addr().unwrap();
945
946 tokio::spawn(async move {
947 if let Ok((mut socket, _)) = listener.accept().await {
948 let mut buf = [0u8; 1024];
949 let _ = socket.read(&mut buf).await;
950 let response = "HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok";
951 let _ = socket.write_all(response.as_bytes()).await;
952 }
953 });
954
955 let acl = HttpAcl::builder()
956 .non_global_ip_ranges(true)
957 .ip_acl_default(true)
958 .port_acl_default(true)
959 .host_acl_default(true)
960 .build_full(HttpAclHooks {
961 modify_response_fn: Some(Arc::new(|_scheme, _authority, mutation| {
962 mutation.body = Bytes::from_static(b"changed");
963 })),
964 ..Default::default()
965 });
966
967 let middleware = HttpAclMiddleware::new(acl);
968
969 let client = reqwest_middleware::ClientBuilder::new(
970 reqwest::Client::builder()
971 .dns_resolver(middleware.dns_resolver())
972 .build()
973 .unwrap(),
974 )
975 .with(middleware)
976 .build();
977
978 let url = format!("http://127.0.0.1:{}/", addr.port());
979 let response = client.get(&url).send().await.unwrap();
980
981 assert_eq!(response.url().as_str(), url);
982 assert!(
983 response.remote_addr().is_some(),
984 "remote_addr() was lost when the response was rebuilt"
985 );
986 }
987
988 #[tokio::test]
989 async fn test_modify_request_header_carries_through_redirect() {
990 use tokio::io::{AsyncReadExt, AsyncWriteExt};
991 use tokio::net::TcpListener;
992
993 let second_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
994 let second_addr = second_listener.local_addr().unwrap();
995
996 let (second_captured_tx, second_captured_rx) = tokio::sync::oneshot::channel();
997 tokio::spawn(async move {
998 if let Ok((mut socket, _)) = second_listener.accept().await {
999 let mut buf = [0u8; 4096];
1000 let n = socket.read(&mut buf).await.unwrap_or(0);
1001 let _ = second_captured_tx.send(String::from_utf8_lossy(&buf[..n]).into_owned());
1002 let response = "HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n";
1003 let _ = socket.write_all(response.as_bytes()).await;
1004 }
1005 });
1006
1007 let first_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1008 let first_addr = first_listener.local_addr().unwrap();
1009 tokio::spawn(async move {
1010 if let Ok((mut socket, _)) = first_listener.accept().await {
1011 let mut buf = [0u8; 1024];
1012 let _ = socket.read(&mut buf).await;
1013 let response = format!(
1014 "HTTP/1.1 302 Found\r\nLocation: http://127.0.0.1:{}/\r\nContent-Length: 0\r\n\r\n",
1015 second_addr.port()
1016 );
1017 let _ = socket.write_all(response.as_bytes()).await;
1018 }
1019 });
1020
1021 let acl = HttpAcl::builder()
1022 .non_global_ip_ranges(true)
1023 .ip_acl_default(true)
1024 .port_acl_default(true)
1025 .host_acl_default(true)
1026 .build_full(HttpAclHooks {
1027 modify_request_fn: Some(Arc::new(|_scheme, _authority, mutation| {
1028 mutation
1029 .headers
1030 .push(("x-marker".to_string(), "hop-one-only".to_string()));
1031 })),
1032 ..Default::default()
1033 });
1034
1035 let middleware = HttpAclMiddleware::new(acl);
1036
1037 let client = reqwest_middleware::ClientBuilder::new(
1038 reqwest::Client::builder()
1039 .dns_resolver(middleware.dns_resolver())
1040 .redirect(middleware.redirect_policy())
1041 .build()
1042 .unwrap(),
1043 )
1044 .with(middleware)
1045 .build();
1046
1047 let request = client
1048 .get(format!("http://127.0.0.1:{}/", first_addr.port()))
1049 .send()
1050 .await;
1051
1052 assert!(request.is_ok(), "{:?}", request.err());
1053 let second_captured = second_captured_rx.await.unwrap();
1054 assert!(
1060 second_captured.to_lowercase().contains("x-marker"),
1061 "expected the header injected for the original request to carry through \
1062 reqwest's redirect handling to the second hop, but it did not:\n{second_captured}"
1063 );
1064 }
1065}