Skip to main content

http_acl_reqwest/
lib.rs

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)]
22/// A reqwest middleware that enforces an [`HttpAcl`].
23///
24/// On each request, checks (in order) the scheme, method, host or IP, port, headers,
25/// URL path, and finally any custom `ValidateFn`, returning
26/// [`reqwest_middleware::Error::Middleware`] on the first denial. This alone only
27/// covers the request as originally built: attach [`Self::dns_resolver`] to the
28/// `Client` as well, so domains are checked against the ACL as they resolve, and
29/// [`Self::redirect_policy`], so redirect targets are checked too. See the crate-level
30/// documentation for a full example wiring all three together.
31pub struct HttpAclMiddleware {
32    acl: Arc<HttpAcl>,
33}
34
35impl HttpAclMiddleware {
36    /// Create a new HTTP ACL middleware from an already-built [`HttpAcl`].
37    pub fn new(acl: HttpAcl) -> Self {
38        Self { acl: Arc::new(acl) }
39    }
40
41    /// Get the [`HttpAcl`] this middleware enforces.
42    pub fn acl(&self) -> Arc<HttpAcl> {
43        self.acl.clone()
44    }
45
46    /// Create a DNS resolver that enforces the ACL, using `getaddrinfo` to actually
47    /// resolve hostnames.
48    ///
49    /// Set via `Client::builder().dns_resolver(...)`. Without this, a domain that
50    /// resolves to a denied or non-global IP (the classic SSRF vector) is never
51    /// checked, since [`HttpAclMiddleware`] only ever sees the request as built, not
52    /// the address it eventually connects to.
53    pub fn dns_resolver(&self) -> Arc<HttpAclDnsResolver> {
54        Arc::new(HttpAclDnsResolver::new(self))
55    }
56
57    /// Same as [`Self::dns_resolver`], but delegating actual resolution to a custom
58    /// [`Resolve`] implementation instead of `getaddrinfo`.
59    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    /// Create a [`redirect::Policy`] that enforces the ACL on every redirect hop.
64    ///
65    /// # Why this is necessary
66    ///
67    /// `HttpAclMiddleware` only validates the request it is given. By default `reqwest`
68    /// follows HTTP redirects internally (up to 10 hops) before control ever returns to
69    /// the middleware chain, so a server an allowed host redirects to - e.g. a `302` to
70    /// `http://169.254.169.254/` - is never re-checked against the ACL. Set this policy
71    /// on the `Client` (in addition to [`Self::dns_resolver`]) to close that gap:
72    ///
73    /// ```no_run
74    /// # use http_acl_reqwest::HttpAclMiddleware;
75    /// # use http_acl::HttpAcl;
76    /// # let middleware = HttpAclMiddleware::new(HttpAcl::builder().build());
77    /// let client = reqwest::Client::builder()
78    ///     .dns_resolver(middleware.dns_resolver())
79    ///     .redirect(middleware.redirect_policy())
80    ///     .build()
81    ///     .unwrap();
82    /// ```
83    ///
84    /// Uses a maximum of 10 redirects, matching `reqwest`'s own default. Use
85    /// [`Self::redirect_policy_with_max`] to customise this.
86    ///
87    /// # Limitations
88    ///
89    /// Only the scheme, host/IP, port, and URL path of each redirect target can be
90    /// checked this way - `reqwest`'s redirect policy does not expose the headers or
91    /// body of the redirected request, so denied headers, denied bodies, and any custom
92    /// `validate_fn` are not re-evaluated per hop.
93    pub fn redirect_policy(&self) -> redirect::Policy {
94        self.redirect_policy_with_max(10)
95    }
96
97    /// Same as [`Self::redirect_policy`], but with a custom maximum number of redirects.
98    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            // `Attempt::error`/`follow`/`stop` consume `attempt` by value, so the denial
102            // reason (if any) is computed into an owned `String` first, in its own scope,
103            // to release the borrow of `attempt` held by `attempt.url()`/`attempt.previous()`.
104            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                // `Url::path()` is percent-encoded; `is_url_path_allowed` expects a
147                // decoded path.
148                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            // `Url::path()` is percent-encoded; `is_url_path_allowed` expects a
251            // decoded path.
252            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            // `Name` is a bare hostname with no port, so `ToSocketAddrs` must be given one
294            // explicitly (e.g. via a tuple) - calling it on the string directly always fails.
295            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
303/// A [`Resolve`]r that checks each resolved address against an [`HttpAcl`] before
304/// handing it back to `reqwest`.
305///
306/// Denies the hostname itself first via [`HttpAcl::is_host_allowed`]. For the
307/// addresses it resolves to, a trusted static DNS mapping (see
308/// [`HttpAclBuilder::add_trusted_static_dns_mapping`]) is returned as-is; a regular
309/// static mapping or a genuinely resolved address is
310/// filtered through [`HttpAcl::is_ip_allowed`] and [`HttpAcl::is_port_allowed`], so
311/// only addresses the ACL permits are ever handed to `reqwest`. Constructed via
312/// [`HttpAclMiddleware::dns_resolver`] or [`HttpAclMiddleware::with_dns_resolver`],
313/// not directly.
314pub struct HttpAclDnsResolver {
315    dns_resolver: Arc<dyn Resolve>,
316    acl: Arc<HttpAcl>,
317}
318
319impl HttpAclDnsResolver {
320    /// Create a new ACL resolver that resolves hostnames via `getaddrinfo`.
321    pub fn new(middleware: &HttpAclMiddleware) -> Self {
322        Self {
323            dns_resolver: Arc::new(GaiResolver),
324            acl: middleware.acl(),
325        }
326    }
327
328    /// Create a new ACL resolver that delegates actual resolution to a custom
329    /// [`Resolve`] implementation.
330    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                // Trusted mappings intentionally bypass the IP/port ACL - the caller
356                // vouches for this destination (e.g. pinning a hostname to an internal
357                // address on purpose).
358                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                // Regular static mappings must still pass the IP/port ACL, just like
362                // resolved addresses do below - otherwise they'd be a way to bypass it
363                // entirely (e.g. mapping a host to a private IP while non-global IPs
364                // are denied).
365                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)]
397/// An error that can occur when resolving a host.
398///
399/// Returned by [`HttpAclDnsResolver`] when a hostname itself is denied by the ACL.
400/// Downcast the boxed error from a failed resolution to check for this specifically,
401/// as opposed to a lower-level resolution failure.
402pub enum HttpAclError {
403    /// Host resolution denied by ACL.
404    #[error("Host resolution denied by ACL: {host}")]
405    HostDenied {
406        /// The host that was denied.
407        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        // Regression test: `Url::path()` is percent-encoded ("%20" for the space
459        // here), but `is_url_path_allowed` matches against the decoded path, so the
460        // middleware must decode before checking - otherwise this denied path would
461        // never match and the request would go through.
462        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        // Regression test: the default `GaiResolver` used to call `to_socket_addrs()` on a
534        // bare hostname (no port), which always errors, so *no* hostname could ever resolve.
535        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        // Deny everything at the IP/port level (the default), but pin "trusted.internal"
561        // to our mock server via a *trusted* static mapping, which should bypass that.
562        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        // No explicit port in the URL - the connector must pick up the trusted
583        // mapping's port, proving both the IP and port ACL were bypassed for it.
584        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        // Allow everything except one specific (non-global) IP, so the initial request to
607        // our local mock server succeeds but the redirect target is denied.
608        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}