Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions crates/geoip/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"] }
157 changes: 145 additions & 12 deletions crates/geoip/src/block/middleware.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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},
},
Expand All @@ -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<IpAddr> + Send + Sync;

#[derive(Clone)]
pub struct ClientIpExtractor(Arc<ExtractClientIp>);

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: F) -> Self
where
F: Fn(&HeaderMap, &Extensions) -> Option<IpAddr> + 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<IpAddr> {
(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<R> {
filter: ZoneFilter,
ip_resolver: R,
client_ip: ClientIpExtractor,
}

impl<R> Inner<R>
where
R: Resolver,
{
fn new(
ip_resolver: R,
blocked_zones: Vec<String>,
blocking_policy: BlockingPolicy,
client_ip: ClientIpExtractor,
) -> Self {
Self {
filter: ZoneFilter::new(blocked_zones, blocking_policy),
ip_resolver,
client_ip,
}
}

fn check<ReqBody>(&self, request: &Request<ReqBody>) -> 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
Expand All @@ -45,16 +132,40 @@ impl<R> GeoBlockLayer<R>
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<String>,
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<String>,
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,
)),
}
}
}
Expand Down Expand Up @@ -89,18 +200,44 @@ impl<S, R> GeoBlockService<S, R>
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<String>,
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<String>,
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,
)),
}
}
}
Expand All @@ -122,11 +259,7 @@ where
fn call(&mut self, request: Request<ReqBody>) -> 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) => {
Expand Down
123 changes: 121 additions & 2 deletions crates/geoip/src/block/middleware/tests.rs
Original file line number Diff line number Diff line change
@@ -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<Body>) -> Result<Response<Body>, Infallible> {
Ok(Response::new(Body::empty()))
}
Expand Down Expand Up @@ -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);
}
Loading