1#![doc = include_str!("../README.md")]
2#![warn(missing_docs)]
3
4use std::future;
5use std::net::{SocketAddr, ToSocketAddrs};
6use std::sync::Arc;
7
8use anyhow::anyhow;
9use http::Extensions;
10use http_acl::utils::authority::{Authority, Host};
11use reqwest::{
12 Request, Response,
13 dns::{Name, Resolve, Resolving},
14 redirect,
15};
16use reqwest_middleware::{Error, Middleware, Next};
17use thiserror::Error;
18
19pub use http_acl::{self, HttpAcl, HttpAclBuilder};
20
21#[derive(Debug, Clone)]
22pub struct HttpAclMiddleware {
32 acl: Arc<HttpAcl>,
33}
34
35impl HttpAclMiddleware {
36 pub fn new(acl: HttpAcl) -> Self {
38 Self { acl: Arc::new(acl) }
39 }
40
41 pub fn acl(&self) -> Arc<HttpAcl> {
43 self.acl.clone()
44 }
45
46 pub fn dns_resolver(&self) -> Arc<HttpAclDnsResolver> {
54 Arc::new(HttpAclDnsResolver::new(self))
55 }
56
57 pub fn with_dns_resolver(&self, dns_resolver: Arc<dyn Resolve>) -> Arc<HttpAclDnsResolver> {
60 Arc::new(HttpAclDnsResolver::with_dns_resolver(self, dns_resolver))
61 }
62
63 pub fn redirect_policy(&self) -> redirect::Policy {
94 self.redirect_policy_with_max(10)
95 }
96
97 pub fn redirect_policy_with_max(&self, max_redirects: usize) -> redirect::Policy {
99 let acl = self.acl.clone();
100 redirect::Policy::custom(move |attempt| {
101 let deny_reason = 'reason: {
105 if attempt.previous().len() > max_redirects {
106 break 'reason Some("too many redirects".to_string());
107 }
108
109 let url = attempt.url();
110
111 let scheme = url.scheme();
112 if acl.is_scheme_allowed(scheme).is_denied() {
113 break 'reason Some(format!("scheme {scheme} is denied"));
114 }
115
116 let Some(host) = url.host() else {
117 break 'reason Some("missing host".to_string());
118 };
119
120 match host {
121 url::Host::Domain(domain) => {
122 if acl.is_host_allowed(domain).is_denied() {
123 break 'reason Some(format!("host {domain} is denied"));
124 }
125 }
126 url::Host::Ipv4(ip) => {
127 let ip = std::net::IpAddr::V4(ip);
128 if acl.is_ip_allowed(&ip).is_denied() {
129 break 'reason Some(format!("ip {ip} is denied"));
130 }
131 }
132 url::Host::Ipv6(ip) => {
133 let ip = std::net::IpAddr::V6(ip);
134 if acl.is_ip_allowed(&ip).is_denied() {
135 break 'reason Some(format!("ip {ip} is denied"));
136 }
137 }
138 }
139
140 if let Some(port) = url.port_or_known_default()
141 && acl.is_port_allowed(port).is_denied()
142 {
143 break 'reason Some(format!("port {port} is denied"));
144 }
145
146 match percent_encoding::percent_decode_str(url.path()).decode_utf8() {
149 Ok(path) => {
150 if acl.is_url_path_allowed(&path).is_denied() {
151 break 'reason Some(format!("path {path} is denied"));
152 }
153 }
154 Err(_) => break 'reason Some("invalid URL path encoding".to_string()),
155 }
156
157 None
158 };
159
160 match deny_reason {
161 Some(reason) => attempt.error(std::io::Error::other(reason)),
162 None => attempt.follow(),
163 }
164 })
165 }
166}
167
168#[async_trait::async_trait]
169impl Middleware for HttpAclMiddleware {
170 async fn handle(
171 &self,
172 req: Request,
173 extensions: &mut Extensions,
174 next: Next<'_>,
175 ) -> std::result::Result<Response, Error> {
176 let scheme = req.url().scheme();
177 let acl_scheme_match = self.acl.is_scheme_allowed(scheme);
178 if acl_scheme_match.is_denied() {
179 return Err(Error::Middleware(anyhow!(
180 "scheme {} is denied - {}",
181 scheme,
182 acl_scheme_match
183 )));
184 }
185
186 let method = req.method().as_str();
187 let acl_method_match = self.acl.is_method_allowed(method);
188 if acl_method_match.is_denied() {
189 return Err(Error::Middleware(anyhow!(
190 "method {} is denied - {}",
191 method,
192 acl_method_match
193 )));
194 }
195
196 if let Some(host) = req.url().host_str() {
197 let authority = Authority::parse(host)
198 .map_err(|_| Error::Middleware(anyhow!("invalid host: {}", host)))?;
199
200 match &authority.host {
201 Host::Ip(ip) => {
202 let acl_ip_match = self.acl.is_ip_allowed(ip);
203 if acl_ip_match.is_denied() {
204 return Err(Error::Middleware(anyhow!(
205 "ip {} is denied - {}",
206 ip,
207 acl_ip_match
208 )));
209 }
210 }
211 Host::Domain(domain) => {
212 let acl_host_match = self.acl.is_host_allowed(domain);
213 if acl_host_match.is_denied() {
214 return Err(Error::Middleware(anyhow!(
215 "host {} is denied - {}",
216 domain,
217 acl_host_match
218 )));
219 }
220 }
221 }
222
223 if let Some(port) = req.url().port_or_known_default() {
224 let acl_port_match = self.acl.is_port_allowed(port);
225 if acl_port_match.is_denied() {
226 return Err(Error::Middleware(anyhow!(
227 "port {} is denied - {}",
228 port,
229 acl_port_match
230 )));
231 }
232 }
233
234 for (key, value) in req.headers() {
235 let header_name = key.as_str();
236 let header_value = value.to_str().map_err(|_| {
237 Error::Middleware(anyhow!("invalid header value for {}", header_name))
238 })?;
239 let acl_header_match = self.acl.is_header_allowed(header_name, header_value);
240 if acl_header_match.is_denied() {
241 return Err(Error::Middleware(anyhow!(
242 "header {}: {} is denied - {}",
243 header_name,
244 header_value,
245 acl_header_match
246 )));
247 }
248 }
249
250 let url_path = percent_encoding::percent_decode_str(req.url().path())
253 .decode_utf8()
254 .map_err(|_| Error::Middleware(anyhow!("invalid URL path encoding")))?;
255 let acl_url_path_match = self.acl.is_url_path_allowed(&url_path);
256 if acl_url_path_match.is_denied() {
257 return Err(Error::Middleware(anyhow!(
258 "path {} is denied - {}",
259 url_path,
260 acl_url_path_match
261 )));
262 }
263
264 let valid_match = self.acl.is_valid(
265 scheme,
266 &authority,
267 req.headers()
268 .iter()
269 .filter_map(|(k, v)| Some((k.as_str(), v.to_str().ok()?))),
270 req.body().and_then(|b| b.as_bytes()),
271 );
272 if valid_match.is_denied() {
273 return Err(Error::Middleware(anyhow!(
274 "request is denied - {}",
275 valid_match
276 )));
277 }
278
279 next.run(req, extensions).await
280 } else {
281 return Err(Error::Middleware(anyhow!("missing host")));
282 }
283 }
284}
285
286type BoxError = Box<dyn std::error::Error + Send + Sync>;
287
288struct GaiResolver;
289
290impl Resolve for GaiResolver {
291 fn resolve(&self, name: Name) -> Resolving {
292 Box::pin(async move {
293 let addresses = (name.as_str(), 0)
296 .to_socket_addrs()
297 .map_err(|e| Box::new(e) as BoxError)?;
298 Ok(Box::new(addresses.into_iter()) as Box<dyn Iterator<Item = SocketAddr> + Send>)
299 })
300 }
301}
302
303pub struct HttpAclDnsResolver {
315 dns_resolver: Arc<dyn Resolve>,
316 acl: Arc<HttpAcl>,
317}
318
319impl HttpAclDnsResolver {
320 pub fn new(middleware: &HttpAclMiddleware) -> Self {
322 Self {
323 dns_resolver: Arc::new(GaiResolver),
324 acl: middleware.acl(),
325 }
326 }
327
328 pub fn with_dns_resolver(
331 middleware: &HttpAclMiddleware,
332 dns_resolver: Arc<dyn Resolve>,
333 ) -> Self {
334 Self {
335 dns_resolver,
336 acl: middleware.acl(),
337 }
338 }
339}
340
341impl Resolve for HttpAclDnsResolver {
342 fn resolve(&self, name: Name) -> Resolving {
343 if self.acl.is_host_allowed(name.as_str()).is_denied() {
344 let err: BoxError = Box::new(HttpAclError::HostDenied {
345 host: name.as_str().to_string(),
346 });
347 return Box::pin(future::ready(Err(err)));
348 }
349
350 let acl = self.acl.clone();
351 let resolver = self.dns_resolver.clone();
352
353 Box::pin(async move {
354 if let Some(tcp_address) = acl.resolve_trusted_static_dns_mapping(name.as_str()) {
355 Ok(Box::new(std::iter::once(tcp_address))
359 as Box<dyn Iterator<Item = SocketAddr> + Send>)
360 } else if let Some(tcp_address) = acl.resolve_static_dns_mapping(name.as_str()) {
361 if acl.is_ip_allowed(&tcp_address.ip()).is_allowed()
366 && acl.is_port_allowed(tcp_address.port()).is_allowed()
367 {
368 Ok(Box::new(std::iter::once(tcp_address))
369 as Box<dyn Iterator<Item = SocketAddr> + Send>)
370 } else {
371 let err: BoxError =
372 Box::new(std::io::Error::other("Static DNS mapping denied by ACL"));
373 Err(err)
374 }
375 } else {
376 let resolved = resolver.resolve(name).await;
377 match resolved {
378 Ok(addresses) => {
379 let filtered = addresses
380 .into_iter()
381 .filter(|addr| {
382 acl.is_ip_allowed(&addr.ip()).is_allowed()
383 && acl.is_port_allowed(addr.port()).is_allowed()
384 })
385 .collect::<Vec<_>>();
386 Ok(Box::new(filtered.into_iter())
387 as Box<dyn Iterator<Item = SocketAddr> + Send>)
388 }
389 Err(e) => Err(e),
390 }
391 }
392 })
393 }
394}
395
396#[derive(Error, Debug)]
397pub enum HttpAclError {
403 #[error("Host resolution denied by ACL: {host}")]
405 HostDenied {
406 host: String,
408 },
409}
410
411#[cfg(test)]
412mod tests {
413 use super::*;
414
415 #[tokio::test]
416 async fn test_http_acl_middleware() {
417 let acl = HttpAcl::builder()
418 .add_denied_host("example.com".to_string())
419 .unwrap()
420 .build();
421
422 let middleware = HttpAclMiddleware::new(acl);
423
424 let client = reqwest_middleware::ClientBuilder::new(
425 reqwest::Client::builder()
426 .dns_resolver(middleware.dns_resolver())
427 .build()
428 .unwrap(),
429 )
430 .with(middleware)
431 .build();
432
433 let request = client.get("http://example.com/").send().await;
434
435 assert!(request.is_err());
436 assert_eq!(
437 request.unwrap_err().to_string(),
438 "host example.com is denied - The entity is denied according to the denied ACL."
439 );
440 }
441
442 #[tokio::test]
443 async fn test_middleware_decodes_percent_encoded_path() {
444 let acl = HttpAcl::builder()
445 .add_allowed_host("example.com".to_string())
446 .unwrap()
447 .add_denied_url_path("/secret file".to_string())
448 .unwrap()
449 .build();
450
451 let middleware = HttpAclMiddleware::new(acl);
452
453 let client =
454 reqwest_middleware::ClientBuilder::new(reqwest::Client::builder().build().unwrap())
455 .with(middleware)
456 .build();
457
458 let request = client.get("http://example.com/secret%20file").send().await;
463
464 assert!(request.is_err());
465 assert!(
466 request
467 .unwrap_err()
468 .to_string()
469 .contains("path /secret file is denied")
470 );
471 }
472
473 #[tokio::test]
474 async fn test_dns_resolver_returns_typed_error_for_denied_host() {
475 let acl = HttpAcl::builder()
476 .add_denied_host("denied.example.com".to_string())
477 .unwrap()
478 .build();
479
480 let middleware = HttpAclMiddleware::new(acl);
481 let resolver = middleware.dns_resolver();
482
483 let name: reqwest::dns::Name = "denied.example.com".parse().unwrap();
484 let err = match resolver.resolve(name).await {
485 Ok(_) => panic!("expected resolution to be denied"),
486 Err(e) => e,
487 };
488
489 let acl_err = err
490 .downcast_ref::<HttpAclError>()
491 .expect("expected a HttpAclError");
492 assert!(matches!(
493 acl_err,
494 HttpAclError::HostDenied { host } if host == "denied.example.com"
495 ));
496 }
497
498 #[tokio::test]
499 async fn test_dns_resolver_resolves_hostnames() {
500 use tokio::io::{AsyncReadExt, AsyncWriteExt};
501 use tokio::net::TcpListener;
502
503 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
504 let addr = listener.local_addr().unwrap();
505
506 tokio::spawn(async move {
507 if let Ok((mut socket, _)) = listener.accept().await {
508 let mut buf = [0u8; 1024];
509 let _ = socket.read(&mut buf).await;
510 let response = "HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n";
511 let _ = socket.write_all(response.as_bytes()).await;
512 }
513 });
514
515 let acl = HttpAcl::builder()
516 .non_global_ip_ranges(true)
517 .ip_acl_default(true)
518 .port_acl_default(true)
519 .host_acl_default(true)
520 .build();
521
522 let middleware = HttpAclMiddleware::new(acl);
523
524 let client = reqwest_middleware::ClientBuilder::new(
525 reqwest::Client::builder()
526 .dns_resolver(middleware.dns_resolver())
527 .build()
528 .unwrap(),
529 )
530 .with(middleware)
531 .build();
532
533 let request = client
536 .get(format!("http://localhost:{}/", addr.port()))
537 .send()
538 .await;
539
540 assert!(request.is_ok(), "{:?}", request.err());
541 }
542
543 #[tokio::test]
544 async fn test_trusted_static_dns_mapping_bypasses_ip_port_acl() {
545 use tokio::io::{AsyncReadExt, AsyncWriteExt};
546 use tokio::net::TcpListener;
547
548 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
549 let addr = listener.local_addr().unwrap();
550
551 tokio::spawn(async move {
552 if let Ok((mut socket, _)) = listener.accept().await {
553 let mut buf = [0u8; 1024];
554 let _ = socket.read(&mut buf).await;
555 let response = "HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n";
556 let _ = socket.write_all(response.as_bytes()).await;
557 }
558 });
559
560 let acl = HttpAcl::builder()
563 .host_acl_default(true)
564 .add_trusted_static_dns_mapping("trusted.internal".to_string(), addr)
565 .unwrap()
566 .build();
567
568 assert!(acl.is_ip_allowed(&addr.ip()).is_denied());
569 assert!(acl.is_port_allowed(addr.port()).is_denied());
570
571 let middleware = HttpAclMiddleware::new(acl);
572
573 let client = reqwest_middleware::ClientBuilder::new(
574 reqwest::Client::builder()
575 .dns_resolver(middleware.dns_resolver())
576 .build()
577 .unwrap(),
578 )
579 .with(middleware)
580 .build();
581
582 let request = client.get("http://trusted.internal/").send().await;
585
586 assert!(request.is_ok(), "{:?}", request.err());
587 }
588
589 #[tokio::test]
590 async fn test_redirect_policy_blocks_disallowed_target() {
591 use tokio::io::{AsyncReadExt, AsyncWriteExt};
592 use tokio::net::TcpListener;
593
594 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
595 let addr = listener.local_addr().unwrap();
596
597 tokio::spawn(async move {
598 if let Ok((mut socket, _)) = listener.accept().await {
599 let mut buf = [0u8; 1024];
600 let _ = socket.read(&mut buf).await;
601 let response = "HTTP/1.1 302 Found\r\nLocation: http://192.168.1.1/\r\nContent-Length: 0\r\n\r\n";
602 let _ = socket.write_all(response.as_bytes()).await;
603 }
604 });
605
606 let acl = HttpAcl::builder()
609 .non_global_ip_ranges(true)
610 .ip_acl_default(true)
611 .port_acl_default(true)
612 .add_denied_ip_range((
613 "192.168.1.1".parse::<std::net::IpAddr>().unwrap(),
614 "192.168.1.1".parse::<std::net::IpAddr>().unwrap(),
615 ))
616 .unwrap()
617 .build();
618
619 let middleware = HttpAclMiddleware::new(acl);
620
621 let client = reqwest_middleware::ClientBuilder::new(
622 reqwest::Client::builder()
623 .dns_resolver(middleware.dns_resolver())
624 .redirect(middleware.redirect_policy())
625 .build()
626 .unwrap(),
627 )
628 .with(middleware)
629 .build();
630
631 let request = client
632 .get(format!("http://127.0.0.1:{}/", addr.port()))
633 .send()
634 .await;
635
636 assert!(request.is_err());
637 }
638}