acton-service 0.30.0

Production-ready Rust backend framework with type-enforced API versioning
Documentation
//! Per-request context captured once, before any consumer needs it.
//!
//! Client IP, request ID, and user agent were previously re-derived from raw
//! headers at three separate call sites (PASETO auth, JWT auth, audit
//! middleware). Auth middleware runs before the audit layer, so token-failure
//! events were built from headers alone and shipped with a blank IP whenever a
//! proxy stripped `X-Forwarded-For`, and a blank request ID whenever the ID was
//! generated by `SetRequestIdLayer` rather than supplied by the client
//! (issue #17).
//!
//! [`request_context_middleware`] resolves all three exactly once and stores a
//! [`RequestContext`] in the request extensions; every downstream consumer
//! reads that instead of the headers.
//!
//! # Ordering
//!
//! This middleware must run **after** `request_id_layer()` (so a generated
//! `x-request-id` is visible) and **before** the auth, audit, Cedar, and
//! governor layers. `ServiceBuilder::apply_middleware` wires it at that
//! position automatically; services that assemble a `Router` by hand are
//! responsible for the same ordering, and consumers fall back to header-only
//! extraction when the extension is absent.

use axum::{
    extract::{ConnectInfo, Request},
    http::HeaderMap,
    middleware::Next,
    response::Response,
};
use std::net::{IpAddr, SocketAddr};

/// Request metadata resolved once per request and shared via extensions.
#[derive(Clone, Debug)]
pub struct RequestContext {
    /// Client IP, from trusted forwarding headers or the direct TCP peer.
    pub ip: Option<IpAddr>,
    /// Correlation ID, client-supplied or generated by `request_id_layer()`.
    pub request_id: Option<String>,
    /// Raw `User-Agent` header value.
    pub user_agent: Option<String>,
}

impl RequestContext {
    /// Resolve the context from request parts. Pure — the middleware is a thin
    /// wrapper around this so the resolution logic stays unit-testable.
    pub(crate) fn resolve(headers: &HeaderMap, connect_info: Option<&SocketAddr>) -> Self {
        Self {
            // Forwarding headers are trusted here to match the extraction this
            // replaces; the governor applies its own trust policy separately.
            ip: extract_client_ip(headers, connect_info, true),
            request_id: header_string(headers, "x-request-id"),
            user_agent: header_string(headers, "user-agent"),
        }
    }

    /// Project into an [`AuditSource`](crate::audit::event::AuditSource).
    ///
    /// `subject` is left unset: it is only known after token validation, so
    /// callers that have `Claims` fill it in themselves.
    #[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(),
        }
    }
}

/// Read a header as an owned `String`, skipping non-UTF-8 values.
fn header_string(headers: &HeaderMap, name: &str) -> Option<String> {
    headers
        .get(name)
        .and_then(|v| v.to_str().ok())
        .map(String::from)
}

/// Resolve the client IP for a request.
///
/// Order of precedence:
/// 1. If `trust_forwarded_headers` is true:
///    - The first parseable IP in `X-Forwarded-For` (comma-separated, left-most).
///    - The single IP in `X-Real-IP`.
/// 2. The direct TCP peer from `ConnectInfo<SocketAddr>`.
///
/// Returns `None` when no IP can be resolved. TLS serve paths do not install
/// `ConnectInfo`, so the fallback is genuinely optional.
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())
}

/// Build an [`AuditSource`](crate::audit::event::AuditSource) from a request,
/// preferring the [`RequestContext`] extension and falling back to raw headers.
///
/// The fallback keeps hand-assembled routers (auth middleware applied without
/// [`request_context_middleware`]) behaving exactly as they did before.
#[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()))
}

/// Header-only [`AuditSource`](crate::audit::event::AuditSource) extraction.
///
/// Unlike [`extract_client_ip`] this does not parse the IP, so a malformed
/// `X-Forwarded-For` still round-trips verbatim into the audit record.
#[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"),
    }
}

/// Middleware that inserts a [`RequestContext`] into the request extensions.
///
/// Use with [`axum::middleware::from_fn`]. See the module docs for the ordering
/// this requires.
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());
    }

    /// Captures the context observed by a downstream handler.
    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"));
    }

    /// The issue #17 regression: with proxy headers stripped, the peer address
    /// must still reach downstream consumers.
    #[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());
    }

    /// Auth middleware reads the extension, so a stripped-header request still
    /// yields a non-blank IP and request ID on the token-failure path.
    #[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());
    }
}