diff --git a/crates/geoip/Cargo.toml b/crates/geoip/Cargo.toml index 0c6abba..64b5b51 100644 --- a/crates/geoip/Cargo.toml +++ b/crates/geoip/Cargo.toml @@ -25,3 +25,7 @@ maxminddb = "0.27" [dev-dependencies] tokio = { version = "1", features = ["full"] } axum = { workspace = true } +# `util` carries `ServiceExt` and `ServiceBuilder::service_fn`, which every +# middleware test drives the layer through. The non-dev dependency deliberately +# stays featureless: the middleware itself needs none of it. +tower = { version = "0.4", features = ["util"] } diff --git a/crates/geoip/src/block/middleware.rs b/crates/geoip/src/block/middleware.rs index 9c1a9e6..5846693 100644 --- a/crates/geoip/src/block/middleware.rs +++ b/crates/geoip/src/block/middleware.rs @@ -12,8 +12,10 @@ use { axum_client_ip::InsecureClientIp, futures::future::{self, Either, Ready}, http_body::Body, - hyper::{Request, Response, StatusCode}, + hyper::{http::Extensions, HeaderMap, Request, Response, StatusCode}, std::{ + fmt, + net::IpAddr, sync::Arc, task::{Context, Poll}, }, @@ -24,10 +26,95 @@ use { #[cfg(test)] mod tests; +/// How the middleware decides which address a request came from. +/// +/// This choice decides what the geo-block is actually enforcing, so it is +/// explicit rather than implied. A blocked visitor only stays blocked if the +/// address cannot be chosen by the visitor. +type ExtractClientIp = dyn Fn(&HeaderMap, &Extensions) -> Option + Send + Sync; + +#[derive(Clone)] +pub struct ClientIpExtractor(Arc); + +impl ClientIpExtractor { + /// Resolve the client address with `f`. + /// + /// Return `None` when no address can be established. The configured + /// [`BlockingPolicy`] then decides whether that allows or blocks the + /// request, through [`Error::UnableToExtractIPAddress`]. + pub fn new(f: F) -> Self + where + F: Fn(&HeaderMap, &Extensions) -> Option + Send + Sync + 'static, + { + Self(Arc::new(f)) + } + + /// Read the address from forwarding headers: the leftmost entry of + /// `X-Forwarded-For`, then `Forwarded`, `X-Real-Ip`, `Fly-Client-IP`, + /// `True-Client-IP`, `CF-Connecting-IP` and `CloudFront-Viewer-Address`, + /// and finally the connection peer. + /// + /// **Every header in that chain is set by the caller.** A proxy in front of + /// this service appends to `X-Forwarded-For` rather than replacing it, so + /// the leftmost entry stays whatever the caller sent, and the other headers + /// pass through untouched unless something strips them. A visitor can + /// therefore name any address and choose the geo answer they want. + /// + /// Use this only where no trusted component establishes the caller's + /// address. Where one does — an API gateway or load balancer that stamps a + /// header of its own — pass that source through [`ClientIpExtractor::new`] + /// instead. + pub fn insecure_from_forwarding_headers() -> Self { + Self::new(|headers, extensions| { + InsecureClientIp::from(headers, extensions) + .ok() + .map(|client_ip| client_ip.0) + }) + } + + fn extract(&self, headers: &HeaderMap, extensions: &Extensions) -> Option { + (self.0)(headers, extensions) + } +} + +impl fmt::Debug for ClientIpExtractor { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_tuple("ClientIpExtractor").finish_non_exhaustive() + } +} + #[derive(Debug)] struct Inner { filter: ZoneFilter, ip_resolver: R, + client_ip: ClientIpExtractor, +} + +impl Inner +where + R: Resolver, +{ + fn new( + ip_resolver: R, + blocked_zones: Vec, + blocking_policy: BlockingPolicy, + client_ip: ClientIpExtractor, + ) -> Self { + Self { + filter: ZoneFilter::new(blocked_zones, blocking_policy), + ip_resolver, + client_ip, + } + } + + fn check(&self, request: &Request) -> Result<(), Error> { + let client_ip = self + .client_ip + .extract(request.headers(), request.extensions()) + .ok_or(Error::UnableToExtractIPAddress)?; + + self.filter.check(client_ip, &self.ip_resolver) + } } /// Layer that applies the GeoBlock middleware which blocks requests base on IP @@ -45,16 +132,40 @@ impl GeoBlockLayer where R: Resolver, { + /// Block on an address read from forwarding headers. + /// + /// The caller controls those headers — see + /// [`ClientIpExtractor::insecure_from_forwarding_headers`]. Prefer + /// [`GeoBlockLayer::with_client_ip_extractor`] wherever something in front + /// of this service establishes the address. pub fn new( ip_resolver: R, blocked_countries: Vec, blocking_policy: BlockingPolicy, + ) -> Self { + Self::with_client_ip_extractor( + ip_resolver, + blocked_countries, + blocking_policy, + ClientIpExtractor::insecure_from_forwarding_headers(), + ) + } + + /// Block on the address `client_ip` reports, instead of on forwarding + /// headers. + pub fn with_client_ip_extractor( + ip_resolver: R, + blocked_countries: Vec, + blocking_policy: BlockingPolicy, + client_ip: ClientIpExtractor, ) -> Self { Self { - inner: Arc::new(Inner { - filter: ZoneFilter::new(blocked_countries, blocking_policy), + inner: Arc::new(Inner::new( ip_resolver, - }), + blocked_countries, + blocking_policy, + client_ip, + )), } } } @@ -89,18 +200,44 @@ impl GeoBlockService where R: Resolver, { + /// Block on an address read from forwarding headers. + /// + /// The caller controls those headers — see + /// [`ClientIpExtractor::insecure_from_forwarding_headers`]. Prefer + /// [`GeoBlockService::with_client_ip_extractor`] wherever something in + /// front of this service establishes the address. pub fn new( service: S, ip_resolver: R, blocked_zones: Vec, blocking_policy: BlockingPolicy, + ) -> Self { + Self::with_client_ip_extractor( + service, + ip_resolver, + blocked_zones, + blocking_policy, + ClientIpExtractor::insecure_from_forwarding_headers(), + ) + } + + /// Block on the address `client_ip` reports, instead of on forwarding + /// headers. + pub fn with_client_ip_extractor( + service: S, + ip_resolver: R, + blocked_zones: Vec, + blocking_policy: BlockingPolicy, + client_ip: ClientIpExtractor, ) -> Self { Self { service, - inner: Arc::new(Inner { - filter: ZoneFilter::new(blocked_zones, blocking_policy), + inner: Arc::new(Inner::new( ip_resolver, - }), + blocked_zones, + blocking_policy, + client_ip, + )), } } } @@ -122,11 +259,7 @@ where fn call(&mut self, request: Request) -> Self::Future { let inner = self.inner.as_ref(); - let result = InsecureClientIp::from(request.headers(), request.extensions()) - .map_err(|_| Error::UnableToExtractIPAddress) - .and_then(|client_ip| inner.filter.check(client_ip.0, &inner.ip_resolver)); - - match inner.filter.apply_policy(result) { + match inner.filter.apply_policy(inner.check(&request)) { Ok(_) => Either::Left(self.service.call(request)), Err(err) => { diff --git a/crates/geoip/src/block/middleware/tests.rs b/crates/geoip/src/block/middleware/tests.rs index 93831e6..8eb6c6d 100644 --- a/crates/geoip/src/block/middleware/tests.rs +++ b/crates/geoip/src/block/middleware/tests.rs @@ -1,15 +1,27 @@ use { crate::{ - block::{middleware::GeoBlockLayer, BlockingPolicy}, + block::{ + middleware::{ClientIpExtractor, GeoBlockLayer}, + BlockingPolicy, + }, LocalResolver, }, axum::body::Body, hyper::{Request, Response, StatusCode}, maxminddb::{geoip2, geoip2::City}, - std::{convert::Infallible, net::IpAddr, sync::Arc}, + std::{ + convert::Infallible, + net::{IpAddr, Ipv4Addr}, + sync::Arc, + }, tower::{Service, ServiceBuilder, ServiceExt}, }; +/// Resolves to a blocked country. +const BLOCKED_IP: IpAddr = IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1)); +/// Resolves to an unblocked country, and is what a caller claims to be. +const CLAIMED_IP: &str = "10.0.0.2"; + async fn handle(_request: Request) -> Result, Infallible> { Ok(Response::new(Body::empty())) } @@ -230,3 +242,110 @@ async fn test_arc() { assert_eq!(response.status(), StatusCode::UNAUTHORIZED); } + +/// Resolves one address to a blocked country and everything else to an +/// unblocked one, so a test can tell which address the middleware actually +/// used rather than only whether it blocked. +fn resolve_by_ip(addr: IpAddr) -> City<'static> { + let iso_code = if addr == BLOCKED_IP { "CU" } else { "US" }; + + City { + country: geoip2::city::Country { + iso_code: Some(iso_code), + ..Default::default() + }, + ..Default::default() + } +} + +/// A configured extractor decides the outcome, so forwarding headers naming +/// some other country cannot buy a caller its way out of a block. +#[tokio::test] +async fn test_configured_extractor_beats_spoofed_forwarding_headers() { + let resolver = LocalResolver::new(Some(resolve_by_ip), None); + + let geoblock = GeoBlockLayer::with_client_ip_extractor( + resolver, + vec!["CU".into()], + BlockingPolicy::Block, + ClientIpExtractor::new(|_headers, _extensions| Some(BLOCKED_IP)), + ); + + let mut service = ServiceBuilder::new().layer(geoblock).service_fn(handle); + + // Every header `InsecureClientIp` reads, all naming an unblocked country. + let request = Request::builder() + .header("X-Forwarded-For", CLAIMED_IP) + .header("Forwarded", format!("for={CLAIMED_IP}")) + .header("X-Real-Ip", CLAIMED_IP) + .header("Fly-Client-IP", CLAIMED_IP) + .header("True-Client-IP", CLAIMED_IP) + .header("CF-Connecting-IP", CLAIMED_IP) + .body(Body::empty()) + .unwrap(); + + let response = service.ready().await.unwrap().call(request).await.unwrap(); + + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); +} + +/// The header-reading default does the opposite, which is why it is named for +/// what it trusts. Pinned so a dependency bump cannot quietly change which +/// headers decide a block. +#[tokio::test] +async fn test_default_extractor_still_trusts_forwarding_headers() { + let resolver = LocalResolver::new(Some(resolve_by_ip), None); + + let geoblock = GeoBlockLayer::new(resolver, vec!["CU".into()], BlockingPolicy::Block); + + let mut service = ServiceBuilder::new().layer(geoblock).service_fn(handle); + + let request = Request::builder() + .header("X-Forwarded-For", CLAIMED_IP) + .body(Body::empty()) + .unwrap(); + + let response = service.ready().await.unwrap().call(request).await.unwrap(); + + assert_eq!(response.status(), StatusCode::OK); +} + +/// An extractor that establishes no address is an extraction failure, so the +/// blocking policy decides rather than the request sailing through. +#[tokio::test] +async fn test_extractor_without_an_address_defers_to_the_policy() { + let no_address = || ClientIpExtractor::new(|_headers, _extensions| None); + let request = || Request::builder().body(Body::empty()).unwrap(); + + let blocking = GeoBlockLayer::with_client_ip_extractor( + LocalResolver::new(Some(resolve_by_ip), None), + vec!["CU".into()], + BlockingPolicy::Block, + no_address(), + ); + let mut service = ServiceBuilder::new().layer(blocking).service_fn(handle); + let response = service + .ready() + .await + .unwrap() + .call(request()) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR); + + let allowing = GeoBlockLayer::with_client_ip_extractor( + LocalResolver::new(Some(resolve_by_ip), None), + vec!["CU".into()], + BlockingPolicy::AllowExtractFailure, + no_address(), + ); + let mut service = ServiceBuilder::new().layer(allowing).service_fn(handle); + let response = service + .ready() + .await + .unwrap() + .call(request()) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); +}