use axum::{
extract::{ConnectInfo, Request},
http::HeaderMap,
middleware::Next,
response::Response,
};
use std::net::{IpAddr, SocketAddr};
#[derive(Clone, Debug)]
pub struct RequestContext {
pub ip: Option<IpAddr>,
pub request_id: Option<String>,
pub user_agent: Option<String>,
}
impl RequestContext {
pub(crate) fn resolve(headers: &HeaderMap, connect_info: Option<&SocketAddr>) -> Self {
Self {
ip: extract_client_ip(headers, connect_info, true),
request_id: header_string(headers, "x-request-id"),
user_agent: header_string(headers, "user-agent"),
}
}
#[cfg(feature = "audit")]
pub fn audit_source(&self) -> crate::audit::event::AuditSource {
crate::audit::event::AuditSource {
ip: self.ip.map(|ip| ip.to_string()),
user_agent: self.user_agent.clone(),
subject: None,
request_id: self.request_id.clone(),
}
}
}
fn header_string(headers: &HeaderMap, name: &str) -> Option<String> {
headers
.get(name)
.and_then(|v| v.to_str().ok())
.map(String::from)
}
pub(crate) fn extract_client_ip(
headers: &HeaderMap,
connect_info: Option<&SocketAddr>,
trust_forwarded_headers: bool,
) -> Option<IpAddr> {
if trust_forwarded_headers {
if let Some(value) = headers.get("x-forwarded-for") {
if let Ok(s) = value.to_str() {
if let Some(first) = s.split(',').next() {
if let Ok(ip) = first.trim().parse::<IpAddr>() {
return Some(ip);
}
}
}
}
if let Some(value) = headers.get("x-real-ip") {
if let Ok(s) = value.to_str() {
if let Ok(ip) = s.trim().parse::<IpAddr>() {
return Some(ip);
}
}
}
}
connect_info.map(|sa| sa.ip())
}
#[cfg(feature = "audit")]
pub(crate) fn audit_source_for_request<B>(
request: &http::Request<B>,
) -> crate::audit::event::AuditSource {
request
.extensions()
.get::<RequestContext>()
.map(RequestContext::audit_source)
.unwrap_or_else(|| audit_source_from_headers(request.headers()))
}
#[cfg(feature = "audit")]
pub(crate) fn audit_source_from_headers(headers: &HeaderMap) -> crate::audit::event::AuditSource {
crate::audit::event::AuditSource {
ip: headers
.get("x-forwarded-for")
.or_else(|| headers.get("x-real-ip"))
.and_then(|v| v.to_str().ok())
.map(|s| s.split(',').next().unwrap_or(s).trim().to_string()),
user_agent: header_string(headers, "user-agent"),
subject: None,
request_id: header_string(headers, "x-request-id"),
}
}
pub async fn request_context_middleware(mut request: Request, next: Next) -> Response {
let connect_info = request
.extensions()
.get::<ConnectInfo<SocketAddr>>()
.map(|ci| ci.0);
let context = RequestContext::resolve(request.headers(), connect_info.as_ref());
request.extensions_mut().insert(context);
next.run(request).await
}
#[cfg(test)]
mod tests {
use super::*;
use axum::{body::Body, routing::get, Router};
use http::StatusCode;
use std::net::Ipv4Addr;
use std::sync::{Arc, Mutex};
use tower::ServiceExt;
fn peer(octets: [u8; 4], port: u16) -> SocketAddr {
SocketAddr::new(IpAddr::V4(Ipv4Addr::from(octets)), port)
}
#[test]
fn extract_client_ip_prefers_first_xff_hop() {
let mut headers = HeaderMap::new();
headers.insert("x-forwarded-for", "203.0.113.5, 10.0.0.1".parse().unwrap());
let ip = extract_client_ip(&headers, None, true).expect("IP from XFF");
assert_eq!(ip.to_string(), "203.0.113.5");
}
#[test]
fn extract_client_ip_falls_back_to_real_ip() {
let mut headers = HeaderMap::new();
headers.insert("x-real-ip", "198.51.100.7".parse().unwrap());
let ip = extract_client_ip(&headers, None, true).expect("IP from X-Real-IP");
assert_eq!(ip.to_string(), "198.51.100.7");
}
#[test]
fn extract_client_ip_falls_back_to_connect_info() {
let headers = HeaderMap::new();
let sa = peer([192, 0, 2, 9], 12345);
let ip = extract_client_ip(&headers, Some(&sa), true).expect("IP from connect info");
assert_eq!(ip.to_string(), "192.0.2.9");
}
#[test]
fn extract_client_ip_ignores_headers_when_untrusted() {
let mut headers = HeaderMap::new();
headers.insert("x-forwarded-for", "203.0.113.5".parse().unwrap());
let sa = peer([192, 0, 2, 9], 12345);
let ip = extract_client_ip(&headers, Some(&sa), false)
.expect("IP from connect info despite XFF");
assert_eq!(ip.to_string(), "192.0.2.9");
}
#[test]
fn extract_client_ip_none_without_sources() {
let headers = HeaderMap::new();
assert!(extract_client_ip(&headers, None, true).is_none());
assert!(extract_client_ip(&headers, None, false).is_none());
}
#[test]
fn resolve_reads_request_id_and_user_agent() {
let mut headers = HeaderMap::new();
headers.insert("x-request-id", "req_abc123".parse().unwrap());
headers.insert("user-agent", "acton-test/1.0".parse().unwrap());
let ctx = RequestContext::resolve(&headers, None);
assert_eq!(ctx.request_id.as_deref(), Some("req_abc123"));
assert_eq!(ctx.user_agent.as_deref(), Some("acton-test/1.0"));
assert!(ctx.ip.is_none());
}
#[test]
fn resolve_yields_all_none_for_bare_request() {
let ctx = RequestContext::resolve(&HeaderMap::new(), None);
assert!(ctx.ip.is_none());
assert!(ctx.request_id.is_none());
assert!(ctx.user_agent.is_none());
}
fn capture_router(sink: Arc<Mutex<Option<RequestContext>>>) -> Router {
Router::new()
.route(
"/",
get(move |request: Request| {
let sink = Arc::clone(&sink);
async move {
*sink.lock().expect("sink poisoned") =
request.extensions().get::<RequestContext>().cloned();
StatusCode::OK
}
}),
)
.layer(axum::middleware::from_fn(request_context_middleware))
}
#[tokio::test]
async fn middleware_inserts_context_from_headers() {
let sink = Arc::new(Mutex::new(None));
let router = capture_router(Arc::clone(&sink));
let request = http::Request::builder()
.uri("/")
.header("x-forwarded-for", "203.0.113.5, 10.0.0.1")
.header("x-request-id", "req_generated")
.header("user-agent", "curl/8.0")
.body(Body::empty())
.unwrap();
let response = router.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let ctx = sink.lock().unwrap().clone().expect("context inserted");
assert_eq!(
ctx.ip.map(|ip| ip.to_string()).as_deref(),
Some("203.0.113.5")
);
assert_eq!(ctx.request_id.as_deref(), Some("req_generated"));
assert_eq!(ctx.user_agent.as_deref(), Some("curl/8.0"));
}
#[tokio::test]
async fn middleware_uses_connect_info_when_headers_absent() {
let sink = Arc::new(Mutex::new(None));
let router = capture_router(Arc::clone(&sink));
let mut request = http::Request::builder()
.uri("/")
.body(Body::empty())
.unwrap();
request
.extensions_mut()
.insert(ConnectInfo(peer([198, 51, 100, 42], 51234)));
let response = router.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let ctx = sink.lock().unwrap().clone().expect("context inserted");
assert_eq!(
ctx.ip.map(|ip| ip.to_string()).as_deref(),
Some("198.51.100.42")
);
}
#[cfg(feature = "audit")]
#[test]
fn audit_source_maps_context_fields() {
let ctx = RequestContext {
ip: Some(IpAddr::V4(Ipv4Addr::new(198, 51, 100, 42))),
request_id: Some("req_abc".to_string()),
user_agent: Some("curl/8.0".to_string()),
};
let source = ctx.audit_source();
assert_eq!(source.ip.as_deref(), Some("198.51.100.42"));
assert_eq!(source.request_id.as_deref(), Some("req_abc"));
assert_eq!(source.user_agent.as_deref(), Some("curl/8.0"));
assert!(source.subject.is_none());
}
#[cfg(feature = "audit")]
#[test]
fn audit_source_for_request_prefers_extension_over_headers() {
let mut request = http::Request::builder()
.uri("/")
.header("x-forwarded-for", "10.0.0.1")
.body(Body::empty())
.unwrap();
request.extensions_mut().insert(RequestContext {
ip: Some(IpAddr::V4(Ipv4Addr::new(198, 51, 100, 42))),
request_id: Some("req_generated".to_string()),
user_agent: None,
});
let source = audit_source_for_request(&request);
assert_eq!(source.ip.as_deref(), Some("198.51.100.42"));
assert_eq!(source.request_id.as_deref(), Some("req_generated"));
}
#[cfg(feature = "audit")]
#[test]
fn audit_source_for_request_falls_back_to_headers() {
let request = http::Request::builder()
.uri("/")
.header("x-forwarded-for", "203.0.113.5, 10.0.0.1")
.header("x-request-id", "req_client")
.header("user-agent", "curl/8.0")
.body(Body::empty())
.unwrap();
let source = audit_source_for_request(&request);
assert_eq!(source.ip.as_deref(), Some("203.0.113.5"));
assert_eq!(source.request_id.as_deref(), Some("req_client"));
assert_eq!(source.user_agent.as_deref(), Some("curl/8.0"));
assert!(source.subject.is_none());
}
}