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::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)]
28/// A reqwest middleware that enforces an [`HttpAcl`].
29///
30/// On each request, checks (in order) the scheme, method, host or IP, port, headers,
31/// URL path, and finally any custom `ValidateFn`, returning
32/// [`reqwest_middleware::Error::Middleware`] on the first denial. This alone only
33/// covers the request as originally built: attach [`Self::dns_resolver`] to the
34/// `Client` as well, so domains are checked against the ACL as they resolve, and
35/// [`Self::redirect_policy`], so redirect targets are checked too. See the crate-level
36/// documentation for a full example wiring all three together.
37///
38/// If the `HttpAcl` has a `ModifyRequestFn` and/or `ModifyResponseFn` attached (see
39/// `HttpAclHooks`), those run too - a `ModifyRequestFn` once all the checks above
40/// pass and just before the request is sent, and a `ModifyResponseFn` on the way
41/// back, before the caller ever sees the `Response`. Neither is applied when not
42/// configured - see `HttpAcl::has_modify_request`/`has_modify_response`. A
43/// configured `ModifyResponseFn` forces the whole response body to be buffered and
44/// the response rebuilt from scratch; see the crate README for the performance
45/// trade-off.
46pub struct HttpAclMiddleware {
47    acl: Arc<HttpAcl>,
48}
49
50impl HttpAclMiddleware {
51    /// Create a new HTTP ACL middleware from an already-built [`HttpAcl`].
52    pub fn new(acl: HttpAcl) -> Self {
53        Self { acl: Arc::new(acl) }
54    }
55
56    /// Get the [`HttpAcl`] this middleware enforces.
57    pub fn acl(&self) -> Arc<HttpAcl> {
58        self.acl.clone()
59    }
60
61    /// Create a DNS resolver that enforces the ACL, using `getaddrinfo` to actually
62    /// resolve hostnames.
63    ///
64    /// Set via `Client::builder().dns_resolver(...)`. Without this, a domain that
65    /// resolves to a denied or non-global IP (the classic SSRF vector) is never
66    /// checked, since [`HttpAclMiddleware`] only ever sees the request as built, not
67    /// the address it eventually connects to.
68    pub fn dns_resolver(&self) -> Arc<HttpAclDnsResolver> {
69        Arc::new(HttpAclDnsResolver::new(self))
70    }
71
72    /// Same as [`Self::dns_resolver`], but delegating actual resolution to a custom
73    /// [`Resolve`] implementation instead of `getaddrinfo`.
74    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    /// Create a [`redirect::Policy`] that enforces the ACL on every redirect hop.
79    ///
80    /// # Why this is necessary
81    ///
82    /// `HttpAclMiddleware` only validates the request it is given. By default `reqwest`
83    /// follows HTTP redirects internally (up to 10 hops) before control ever returns to
84    /// the middleware chain, so a server an allowed host redirects to - e.g. a `302` to
85    /// `http://169.254.169.254/` - is never re-checked against the ACL. Set this policy
86    /// on the `Client` (in addition to [`Self::dns_resolver`]) to close that gap:
87    ///
88    /// ```no_run
89    /// # use http_acl_reqwest::HttpAclMiddleware;
90    /// # use http_acl::HttpAcl;
91    /// # let middleware = HttpAclMiddleware::new(HttpAcl::builder().build());
92    /// let client = reqwest::Client::builder()
93    ///     .dns_resolver(middleware.dns_resolver())
94    ///     .redirect(middleware.redirect_policy())
95    ///     .build()
96    ///     .unwrap();
97    /// ```
98    ///
99    /// Uses a maximum of 10 redirects, matching `reqwest`'s own default. Use
100    /// [`Self::redirect_policy_with_max`] to customise this.
101    ///
102    /// # Limitations
103    ///
104    /// Only the scheme, host/IP, port, and URL path of each redirect target can be
105    /// checked this way - `reqwest`'s redirect policy does not expose the headers or
106    /// body of the redirected request, so denied headers, denied bodies, and any custom
107    /// `validate_fn` are not re-evaluated per hop. `ModifyResponseFn` only ever sees
108    /// the final response of a redirect chain, never an intermediate `3xx`, for the
109    /// same reason.
110    ///
111    /// `ModifyRequestFn` is different: it's called once, against the original
112    /// outgoing request, before `Middleware::handle` hands it to `reqwest`, but a
113    /// header it injects is - like any header set before `send()` - carried forward
114    /// by `reqwest`'s own redirect handling to every subsequent hop (`reqwest` may
115    /// still strip specific headers, e.g. `Authorization`, when a redirect crosses
116    /// origins). So an injected header reaches the whole chain even though the
117    /// closure itself does not run again.
118    pub fn redirect_policy(&self) -> redirect::Policy {
119        self.redirect_policy_with_max(10)
120    }
121
122    /// Same as [`Self::redirect_policy`], but with a custom maximum number of redirects.
123    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            // `Attempt::error`/`follow`/`stop` consume `attempt` by value, so the denial
127            // reason (if any) is computed into an owned `String` first, in its own scope,
128            // to release the borrow of `attempt` held by `attempt.url()`/`attempt.previous()`.
129            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                // `Url::path()` is percent-encoded; `is_url_path_allowed` expects a
172                // decoded path.
173                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        // Owned, rather than borrowed from `req.url()`, since it's still needed
202        // after `req` is moved into `next.run(...)` below (to check the response
203        // against `ModifyResponseFn`).
204        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            // `Url::path()` is percent-encoded; `is_url_path_allowed` expects a
279            // decoded path.
280            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                // The point where the whole response body is buffered into memory -
356                // only reached when a `ModifyResponseFn` is actually configured.
357                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                // `Response`'s own `Extensions` (distinct from the `extensions`
373                // parameter of this function) carries things like `HttpInfo`
374                // (backing `Response::remote_addr()`) - restore it here, before
375                // `.url(...)`, not after: `.url()` inserts its own extension entry,
376                // and a wholesale `*extensions_mut() = ...` afterwards would
377                // clobber that again.
378                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            // `Name` is a bare hostname with no port, so `ToSocketAddrs` must be given one
403            // explicitly (e.g. via a tuple) - calling it on the string directly always fails.
404            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
412/// A [`Resolve`]r that checks each resolved address against an [`HttpAcl`] before
413/// handing it back to `reqwest`.
414///
415/// Denies the hostname itself first via [`HttpAcl::is_host_allowed`]. For the
416/// addresses it resolves to, a trusted static DNS mapping (see
417/// [`HttpAclBuilder::add_trusted_static_dns_mapping`]) is returned as-is; a regular
418/// static mapping or a genuinely resolved address is
419/// filtered through [`HttpAcl::is_ip_allowed`] and [`HttpAcl::is_port_allowed`], so
420/// only addresses the ACL permits are ever handed to `reqwest`. Constructed via
421/// [`HttpAclMiddleware::dns_resolver`] or [`HttpAclMiddleware::with_dns_resolver`],
422/// not directly.
423pub struct HttpAclDnsResolver {
424    dns_resolver: Arc<dyn Resolve>,
425    acl: Arc<HttpAcl>,
426}
427
428impl HttpAclDnsResolver {
429    /// Create a new ACL resolver that resolves hostnames via `getaddrinfo`.
430    pub fn new(middleware: &HttpAclMiddleware) -> Self {
431        Self {
432            dns_resolver: Arc::new(GaiResolver),
433            acl: middleware.acl(),
434        }
435    }
436
437    /// Create a new ACL resolver that delegates actual resolution to a custom
438    /// [`Resolve`] implementation.
439    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                // Trusted mappings intentionally bypass the IP/port ACL - the caller
465                // vouches for this destination (e.g. pinning a hostname to an internal
466                // address on purpose).
467                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                // Regular static mappings must still pass the IP/port ACL, just like
471                // resolved addresses do below - otherwise they'd be a way to bypass it
472                // entirely (e.g. mapping a host to a private IP while non-global IPs
473                // are denied).
474                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)]
506/// An error that can occur when resolving a host.
507///
508/// Returned by [`HttpAclDnsResolver`] when a hostname itself is denied by the ACL.
509/// Downcast the boxed error from a failed resolution to check for this specifically,
510/// as opposed to a lower-level resolution failure.
511pub enum HttpAclError {
512    /// Host resolution denied by ACL.
513    #[error("Host resolution denied by ACL: {host}")]
514    HostDenied {
515        /// The host that was denied.
516        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        // Regression test: `Url::path()` is percent-encoded ("%20" for the space
568        // here), but `is_url_path_allowed` matches against the decoded path, so the
569        // middleware must decode before checking - otherwise this denied path would
570        // never match and the request would go through.
571        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        // Regression test: the default `GaiResolver` used to call `to_socket_addrs()` on a
643        // bare hostname (no port), which always errors, so *no* hostname could ever resolve.
644        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        // Deny everything at the IP/port level (the default), but pin "trusted.internal"
670        // to our mock server via a *trusted* static mapping, which should bypass that.
671        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        // No explicit port in the URL - the connector must pick up the trusted
692        // mapping's port, proving both the IP and port ACL were bypassed for it.
693        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        // Allow everything except one specific (non-global) IP, so the initial request to
716        // our local mock server succeeds but the redirect target is denied.
717        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        // `ModifyRequestFn` only ever runs once, against the original request, but
1055        // `reqwest`'s own redirect handling carries a header set before `send()`
1056        // (which is indistinguishable from one injected by the closure) forward to
1057        // every hop - the header shows up here even though the closure itself
1058        // never ran again. See the `redirect_policy` doc comment.
1059        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}