diff --git a/docs/filters/http/observability/trace_context.md b/docs/filters/http/observability/trace_context.md index 1e63b5e75..d722b578f 100644 --- a/docs/filters/http/observability/trace_context.md +++ b/docs/filters/http/observability/trace_context.md @@ -3,15 +3,11 @@ # `trace_context` -Propagates W3C Trace Context headers (`traceparent`, `tracestate`). +Propagates W3C Trace Context and `x-request-id` correlation. ## Configuration Notes -On each request: 1. Parses the incoming `traceparent` header (if present and valid) 2. Joins the existing trace or generates a new trace ID 3. Generates a new span ID for the proxy hop 4. Injects the updated `traceparent` into the upstream request 5. Forwards the `tracestate` header (if present and traceparent was valid) 6. Strips the `tracestate` header when traceparent is absent or invalid - -Per W3C Trace Context section 3.3.1.1, `tracestate` MUST NOT be forwarded when `traceparent` is absent or invalid. - -Malformed `traceparent` headers are silently ignored and treated as absent, per the W3C specification. +Per W3C Trace Context section 3.3.1.1, `tracestate` is forwarded only when inbound `traceparent` is valid. Currently accepts no fields; reserved for future options such as trusted-header policies or sampling flag overrides. diff --git a/docs/filters/reference.md b/docs/filters/reference.md index d83c44c28..2d1c9716b 100644 --- a/docs/filters/reference.md +++ b/docs/filters/reference.md @@ -11,7 +11,7 @@ Built-in filters organized by protocol and category. |--------|---------|-------------| | [`access_log`](http/observability/access_log.md) | - | Logs structured access records for each request and response. | | [`request_id`](http/observability/request_id.md) | - | Ensures every request carries a correlation ID. | -| [`trace_context`](http/observability/trace_context.md) | - | Propagates W3C Trace Context headers (`traceparent`, `tracestate`). | +| [`trace_context`](http/observability/trace_context.md) | - | Propagates W3C Trace Context and `x-request-id` correlation. | ## HTTP / Payload Processing diff --git a/filter/src/builtins/http/observability/trace_context.rs b/filter/src/builtins/http/observability/trace_context.rs index f435d3f83..4d30e677f 100644 --- a/filter/src/builtins/http/observability/trace_context.rs +++ b/filter/src/builtins/http/observability/trace_context.rs @@ -3,58 +3,27 @@ //! W3C Trace Context propagation filter. //! -//! Parses incoming `traceparent` and `tracestate` headers per the -//! [W3C Trace Context](https://www.w3.org/TR/trace-context/) specification, -//! joins an existing trace or generates a new trace ID, and injects -//! the updated headers into the upstream request. +//! Joins or starts a trace, injects `traceparent` and `x-request-id` on the +//! forwarded hop, and stores request-scoped [`TraceContext`] for sub-requests. //! //! # Limitations //! -//! This filter is pure header propagation with no `OTel` dependency: the -//! `parent-id` it injects names a proxy hop that is **not exported as a -//! span** anywhere, so tracing backends show the proxy as a missing node -//! between client and upstream spans. Deployments exporting real proxy -//! spans (the `otel` feature) should rely on span-context propagation -//! there instead. New traces are always flagged sampled (`01`) because -//! the filter cannot consult any sampler configuration. +//! Header propagation only: the hop `parent-id` is not exported as a span. +//! New traces are always flagged sampled (`01`). use std::borrow::Cow; use async_trait::async_trait; use serde::Deserialize; -use tracing::debug; +use tracing::{debug, warn}; use crate::{ FilterAction, FilterError, factory::parse_filter_config, filter::{HttpFilter, HttpFilterContext}, + trace_context::{InboundTrace, REQUEST_ID_HEADER, TRACEPARENT_HEADER, TraceContext, parse_traceparent}, }; -// ----------------------------------------------------------------------------- -// Constants -// ----------------------------------------------------------------------------- - -/// Current supported traceparent version. -const TRACEPARENT_VERSION: &str = "00"; - -/// Expected byte length of a well-formed traceparent header value. -/// -/// Format: `{version}-{trace-id}-{parent-id}-{trace-flags}` -/// ` 2 - 32 - 16 - 2 ` = 55 chars -const TRACEPARENT_LEN: usize = 55; // 2 + 1 + 32 + 1 + 16 + 1 + 2 - -/// Byte length of a hex-encoded trace ID (16 bytes = 32 hex chars). -const TRACE_ID_HEX_LEN: usize = 32; - -/// Byte length of a hex-encoded span/parent ID (8 bytes = 16 hex chars). -const SPAN_ID_HEX_LEN: usize = 16; - -/// All-zero trace ID (invalid per spec). -const INVALID_TRACE_ID: &str = "00000000000000000000000000000000"; - -/// All-zero parent ID (invalid per spec). -const INVALID_PARENT_ID: &str = "0000000000000000"; - // ----------------------------------------------------------------------------- // Config // ----------------------------------------------------------------------------- @@ -71,89 +40,14 @@ const INVALID_PARENT_ID: &str = "0000000000000000"; )] struct TraceContextFilterConfig {} -// ----------------------------------------------------------------------------- -// Parsed Traceparent -// ----------------------------------------------------------------------------- - -/// A validated W3C `traceparent` header. -#[derive(Clone, Debug, Eq, PartialEq)] -struct Traceparent { - /// Hex-encoded parent span ID (16 lowercase hex chars). - parent_id: String, - - /// Hex-encoded trace flags (2 lowercase hex chars). - trace_flags: String, - - /// Hex-encoded trace ID (32 lowercase hex chars). - trace_id: String, -} - -impl Traceparent { - /// Parse and validate a `traceparent` header value. - /// - /// Returns `None` for malformed values per the W3C spec: - /// - Too short for minimum traceparent format - /// - Version `ff` (reserved per W3C spec) - /// - Version `00` with length other than 55 characters - /// - Non-lowercase-hex characters in fields - /// - All-zero trace ID or parent ID - /// - /// Per W3C forward-compatibility rules, versions other than `00` - /// (except reserved `ff`) are accepted by parsing the first 55 - /// characters. Future versions may append additional fields. - fn parse(value: &str) -> Option { - let prefix = value.get(..TRACEPARENT_LEN)?; - - let (version, trace_id, parent_id, trace_flags) = split_traceparent(prefix)?; - - // W3C: version ff is reserved and must always be rejected. - if version == "ff" { - return None; - } - - // W3C: version 00 must be exactly 55 chars; future versions may - // produce longer values, so only enforce exact length for v00. - if version == TRACEPARENT_VERSION && value.len() != TRACEPARENT_LEN { - return None; - } - - validate_traceparent_fields(trace_id, parent_id, trace_flags)?; - - Some(Self { - parent_id: parent_id.to_owned(), - trace_flags: trace_flags.to_owned(), - trace_id: trace_id.to_owned(), - }) - } - - /// Format as a W3C traceparent header value with a new parent ID. - fn format(&self, new_parent_id: &str) -> String { - format!( - "{TRACEPARENT_VERSION}-{}-{new_parent_id}-{}", - self.trace_id, self.trace_flags - ) - } -} - // ----------------------------------------------------------------------------- // TraceContextFilter // ----------------------------------------------------------------------------- -/// Propagates W3C Trace Context headers (`traceparent`, `tracestate`). -/// -/// On each request: -/// 1. Parses the incoming `traceparent` header (if present and valid) -/// 2. Joins the existing trace or generates a new trace ID -/// 3. Generates a new span ID for the proxy hop -/// 4. Injects the updated `traceparent` into the upstream request -/// 5. Forwards the `tracestate` header (if present and traceparent was valid) -/// 6. Strips the `tracestate` header when traceparent is absent or invalid -/// -/// Per W3C Trace Context section 3.3.1.1, `tracestate` MUST NOT be -/// forwarded when `traceparent` is absent or invalid. +/// Propagates W3C Trace Context and `x-request-id` correlation. /// -/// Malformed `traceparent` headers are silently ignored and treated -/// as absent, per the W3C specification. +/// Per W3C Trace Context section 3.3.1.1, `tracestate` is forwarded only +/// when inbound `traceparent` is valid. /// /// # YAML configuration /// @@ -193,20 +87,21 @@ impl HttpFilter for TraceContextFilter { } async fn on_request(&self, ctx: &mut HttpFilterContext<'_>) -> Result { + if reuse_existing_trace_context(ctx) { + return Ok(FilterAction::Continue); + } + let incoming = ctx .request .headers - .get("traceparent") + .get(TRACEPARENT_HEADER) .and_then(|v| v.to_str().ok()) - .and_then(Traceparent::parse); + .and_then(parse_traceparent); - let new_span_id = generate_span_id(ctx); - let traceparent = build_traceparent(incoming.as_ref(), &new_span_id, ctx); + let tc = initialize_trace_context(ctx, incoming.as_ref()); + inject_correlation_headers(ctx, &tc); + ctx.extensions.insert(tc); - inject_traceparent(ctx, traceparent); - - // W3C Trace Context section 3.3.1.1: tracestate MUST NOT be - // forwarded when traceparent is absent or invalid. if incoming.is_some() { forward_tracestate(ctx); } else { @@ -217,144 +112,195 @@ impl HttpFilter for TraceContextFilter { } } -// ----------------------------------------------------------------------------- -// Request Processing Helpers -// ----------------------------------------------------------------------------- +/// Reuse an already-initialized request-scoped trace context. +fn reuse_existing_trace_context(ctx: &mut HttpFilterContext<'_>) -> bool { + let Some(tc) = ctx.extensions.get::().cloned() else { + return false; + }; + let request_id = tc.request_id().to_owned(); + let trace_id = tc.trace_id().to_owned(); + + ensure_extra_header(ctx, REQUEST_ID_HEADER, &request_id); + ensure_traceparent_header(ctx, &tc, &trace_id); + warn_competing_request_id(ctx, &request_id); + true +} -/// Build the outgoing `traceparent` value, joining an existing trace -/// or starting a new one. -fn build_traceparent(incoming: Option<&Traceparent>, new_span_id: &str, ctx: &HttpFilterContext<'_>) -> String { - if let Some(tp) = incoming { +/// Build the request-scoped trace context from inbound trace data or a new trace. +fn initialize_trace_context(ctx: &HttpFilterContext<'_>, incoming: Option<&InboundTrace>) -> TraceContext { + let request_id = resolve_request_id(ctx); + if let Some(trace) = incoming { debug!( - trace_id = %tp.trace_id, - parent_id = %tp.parent_id, - trace_flags = %tp.trace_flags, + trace_id = %trace.trace_id, + flags = %trace.flags, "joining existing trace" ); - tp.format(new_span_id) + TraceContext::from_inbound(request_id, trace) } else { - let trace_id = generate_trace_id(ctx); - // New traces are always marked sampled: this filter propagates - // context without an OTel dependency, so it cannot consult the - // configured sampler. Incoming flags are preserved as-is above; - // only proxy-initiated traces get the unconditional 01. - let trace_flags = "01"; // sampled - debug!(trace_id = %trace_id, "starting new trace"); - format!("{TRACEPARENT_VERSION}-{trace_id}-{new_span_id}-{trace_flags}") + let context = TraceContext::new_sampled(request_id, ctx.id_generator, ctx.time_source); + debug!(trace_id = %context.trace_id(), "starting new trace"); + context } } -/// Inject the `traceparent` header into the upstream request, -/// removing any existing value first. -fn inject_traceparent(ctx: &mut HttpFilterContext<'_>, traceparent: String) { +/// Inject correlation headers for the forwarded upstream request. +fn inject_correlation_headers(ctx: &mut HttpFilterContext<'_>, tc: &TraceContext) { + let request_id = tc.request_id().to_owned(); + let [_, (_, traceparent)] = tc.headers_for_hop(ctx.id_generator, ctx.time_source); + ctx.request_headers_to_remove - .push(http::header::HeaderName::from_static("traceparent")); - ctx.extra_request_headers - .push((Cow::Borrowed("traceparent"), traceparent)); + .push(http::header::HeaderName::from_static(TRACEPARENT_HEADER)); + ctx.request_headers_to_remove + .push(http::header::HeaderName::from_static(REQUEST_ID_HEADER)); + + ensure_extra_header(ctx, REQUEST_ID_HEADER, &request_id); + ensure_extra_header(ctx, TRACEPARENT_HEADER, &traceparent); + warn_competing_request_id(ctx, &request_id); } -/// Forward the `tracestate` header if present on the incoming request. -/// -/// Per RFC 7230, multiple `tracestate` headers are valid and MUST be -/// combined into a single comma-separated value. -fn forward_tracestate(ctx: &mut HttpFilterContext<'_>) { - let values: Vec<&str> = ctx + +/// Resolve request id: pending `request_id` filter, inbound header, then generate. +fn resolve_request_id(ctx: &HttpFilterContext<'_>) -> String { + let inbound = ctx .request .headers - .get_all("tracestate") + .get(REQUEST_ID_HEADER) + .and_then(|v| v.to_str().ok()) + .map(str::to_owned); + + let pending = pending_request_id(ctx); + + match (pending, inbound) { + (Some(pending), Some(inbound)) if pending != inbound => { + warn!( + pending = %pending, + inbound = %inbound, + "competing x-request-id values; preferring pending request_id filter value" + ); + pending + }, + (Some(pending), _) => pending, + (None, Some(inbound)) => inbound, + (None, None) => ctx.id_generator.generate(ctx.time_source), + } +} + +/// Return the first pending `x-request-id` value scheduled for upstream injection. +fn pending_request_id(ctx: &HttpFilterContext<'_>) -> Option { + let values: Vec<&str> = ctx + .extra_request_headers .iter() - .filter_map(|v| v.to_str().ok()) + .filter(|(name, _)| name.eq_ignore_ascii_case(REQUEST_ID_HEADER)) + .map(|(_, value)| value.as_str()) .collect(); - - if values.is_empty() { - return; + match values.as_slice() { + [] => None, + [only] => Some((*only).to_owned()), + [first, rest @ ..] => { + if rest.iter().any(|v| *v != *first) { + warn!( + first = %first, + "multiple distinct pending x-request-id values; using the first" + ); + } + Some((*first).to_owned()) + }, } - - let combined = values.join(", "); - debug!(tracestate = %combined, "forwarding tracestate"); - ctx.request_headers_to_remove - .push(http::header::HeaderName::from_static("tracestate")); - ctx.extra_request_headers.push((Cow::Borrowed("tracestate"), combined)); } -/// Strip the `tracestate` header from the upstream request. -/// -/// Called when `traceparent` is absent or invalid per W3C Trace -/// Context section 3.3.1.1. -fn strip_tracestate(ctx: &mut HttpFilterContext<'_>) { - ctx.request_headers_to_remove - .push(http::header::HeaderName::from_static("tracestate")); -} -// ----------------------------------------------------------------------------- -// ID Generation Helpers -// ----------------------------------------------------------------------------- +/// Leave a competing pending extra in place rather than duplicating it. +fn ensure_extra_header(ctx: &mut HttpFilterContext<'_>, name: &'static str, value: &str) { + let existing: Vec = ctx + .extra_request_headers + .iter() + .filter(|(n, _)| n.eq_ignore_ascii_case(name)) + .map(|(_, v)| v.clone()) + .collect(); -/// Generate a 32-hex-char trace ID using the context's ID generator. -fn generate_trace_id(ctx: &HttpFilterContext<'_>) -> String { - let id = ctx.id_generator.generate(ctx.time_source); - debug_assert_eq!( - id.len(), - TRACE_ID_HEX_LEN, - "IdGenerator should produce {TRACE_ID_HEX_LEN} hex chars" - ); - id -} + if existing.is_empty() { + ctx.extra_request_headers.push((Cow::Borrowed(name), value.to_owned())); + return; + } -/// Generate a 16-hex-char span ID using the context's ID generator. -fn generate_span_id(ctx: &HttpFilterContext<'_>) -> String { - let id = ctx.id_generator.generate(ctx.time_source); - // IdGenerator produces 32 hex chars laid out as - // `{timestamp:12}{seed:8}{counter:12}`. Take the LAST 16 (4 seed + 12 - // counter): the first 16 are timestamp-dominated and identical for every - // request in the same microsecond, which would duplicate span IDs. - let id: String = id.chars().skip(id.chars().count() - SPAN_ID_HEX_LEN).collect(); - debug_assert_eq!(id.len(), SPAN_ID_HEX_LEN, "span ID must be 16 hex chars"); - id + for existing_value in &existing { + if existing_value != value { + warn!( + header = name, + existing = %existing_value, + expected = %value, + "competing correlation header pending; leaving existing value in place" + ); + } + } } -// ----------------------------------------------------------------------------- -// Validation Helpers -// ----------------------------------------------------------------------------- +/// Ensure the forwarded hop has a `traceparent` without warning on expected span-id drift. +fn ensure_traceparent_header(ctx: &mut HttpFilterContext<'_>, tc: &TraceContext, expected_trace_id: &str) { + let existing: Vec = ctx + .extra_request_headers + .iter() + .filter(|(name, _)| name.eq_ignore_ascii_case(TRACEPARENT_HEADER)) + .map(|(_, value)| value.clone()) + .collect(); -/// Split a traceparent value into its four dash-separated fields. -/// -/// Returns `None` if the structure is malformed (wrong number of -/// fields or wrong field lengths). -fn split_traceparent(value: &str) -> Option<(&str, &str, &str, &str)> { - let mut parts = value.splitn(4, '-'); - let version = parts.next()?; - let trace_id = parts.next()?; - let parent_id = parts.next()?; - let trace_flags = parts.next()?; - - if version.len() != 2 - || trace_id.len() != TRACE_ID_HEX_LEN - || parent_id.len() != SPAN_ID_HEX_LEN - || trace_flags.len() != 2 - { - return None; + if existing.is_empty() { + let [_, (_, traceparent)] = tc.headers_for_hop(ctx.id_generator, ctx.time_source); + ctx.extra_request_headers + .push((Cow::Borrowed(TRACEPARENT_HEADER), traceparent)); + return; } - Some((version, trace_id, parent_id, trace_flags)) + for value in existing { + match parse_traceparent(&value) { + Some(parsed) if parsed.trace_id == expected_trace_id => {}, + _ => warn!( + existing = %value, + expected_trace_id = %expected_trace_id, + "competing traceparent pending alongside TraceContext" + ), + } + } } -/// Validate hex content and reject all-zero IDs in traceparent fields. -fn validate_traceparent_fields(trace_id: &str, parent_id: &str, trace_flags: &str) -> Option<()> { - if !is_lowercase_hex(trace_id) || !is_lowercase_hex(parent_id) || !is_lowercase_hex(trace_flags) { - return None; +/// Warn when pending forwarded request headers disagree with the shared request ID. +fn warn_competing_request_id(ctx: &HttpFilterContext<'_>, expected: &str) { + for (name, value) in &ctx.extra_request_headers { + if name.eq_ignore_ascii_case(REQUEST_ID_HEADER) && value != expected { + warn!( + existing = %value, + expected = %expected, + "competing x-request-id pending alongside TraceContext" + ); + } } +} - if trace_id == INVALID_TRACE_ID || parent_id == INVALID_PARENT_ID { - return None; +/// Forward inbound `tracestate` when the inbound `traceparent` was valid. +fn forward_tracestate(ctx: &mut HttpFilterContext<'_>) { + let values: Vec<&str> = ctx + .request + .headers + .get_all("tracestate") + .iter() + .filter_map(|v| v.to_str().ok()) + .collect(); + + if values.is_empty() { + return; } - Some(()) + let combined = values.join(", "); + debug!(tracestate = %combined, "forwarding tracestate"); + ctx.request_headers_to_remove + .push(http::header::HeaderName::from_static("tracestate")); + ensure_extra_header(ctx, "tracestate", &combined); } -/// Check that a string contains only lowercase hexadecimal characters. -fn is_lowercase_hex(s: &str) -> bool { - s.bytes().all(|b| b.is_ascii_digit() || (b'a'..=b'f').contains(&b)) +/// Remove inbound `tracestate` when there is no valid inbound `traceparent`. +fn strip_tracestate(ctx: &mut HttpFilterContext<'_>) { + ctx.request_headers_to_remove + .push(http::header::HeaderName::from_static("tracestate")); } // ----------------------------------------------------------------------------- @@ -371,168 +317,30 @@ fn is_lowercase_hex(s: &str) -> bool { reason = "tests" )] mod tests { - use super::*; + use praxis_core::subrequest::FrameworkHeaders; - // ------------------------------------------------------------------------- - // Traceparent Parsing - // ------------------------------------------------------------------------- - - #[test] - fn parse_valid_traceparent() { - let tp = Traceparent::parse("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01").unwrap(); - assert_eq!(tp.trace_id, "4bf92f3577b34da6a3ce929d0e0e4736", "trace_id should match"); - assert_eq!(tp.parent_id, "00f067aa0ba902b7", "parent_id should match"); - assert_eq!(tp.trace_flags, "01", "trace_flags should match"); - } - - #[test] - fn parse_traceparent_unsampled() { - let tp = Traceparent::parse("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-00").unwrap(); - assert_eq!(tp.trace_flags, "00", "trace_flags should indicate unsampled"); - } - - #[test] - fn parse_traceparent_rejects_wrong_length() { - assert!( - Traceparent::parse("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-0").is_none(), - "too-short traceparent should be rejected" - ); - assert!( - Traceparent::parse("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-011").is_none(), - "too-long traceparent should be rejected" - ); - } - - #[test] - fn parse_traceparent_rejects_wrong_delimiters() { - assert!( - Traceparent::parse("00_4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01").is_none(), - "underscore delimiter should be rejected" - ); - } - - #[test] - fn parse_traceparent_accepts_future_version() { - let tp = Traceparent::parse("01-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01") - .expect("future version 01 should be accepted"); - assert_eq!( - tp.trace_id, "4bf92f3577b34da6a3ce929d0e0e4736", - "trace_id should be parsed" - ); - assert_eq!(tp.parent_id, "00f067aa0ba902b7", "parent_id should be parsed"); - assert_eq!(tp.trace_flags, "01", "trace_flags should be parsed"); - } - - #[test] - fn parse_traceparent_accepts_future_version_with_extra_data() { - let tp = Traceparent::parse("02-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01-extra-data") - .expect("future version with extra fields should be accepted"); - assert_eq!( - tp.trace_id, "4bf92f3577b34da6a3ce929d0e0e4736", - "trace_id should be parsed ignoring extra fields" - ); - } - - #[test] - fn parse_traceparent_rejects_reserved_version_ff() { - assert!( - Traceparent::parse("ff-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01").is_none(), - "version ff should be rejected" - ); - } - - #[test] - fn parse_traceparent_rejects_uppercase_hex() { - assert!( - Traceparent::parse("00-4BF92F3577B34DA6A3CE929D0E0E4736-00f067aa0ba902b7-01").is_none(), - "uppercase hex in trace_id should be rejected" - ); - assert!( - Traceparent::parse("00-4bf92f3577b34da6a3ce929d0e0e4736-00F067AA0BA902B7-01").is_none(), - "uppercase hex in parent_id should be rejected" - ); - } - - #[test] - fn parse_traceparent_rejects_all_zero_trace_id() { - assert!( - Traceparent::parse("00-00000000000000000000000000000000-00f067aa0ba902b7-01").is_none(), - "all-zero trace_id should be rejected" - ); - } - - #[test] - fn parse_traceparent_rejects_all_zero_parent_id() { - assert!( - Traceparent::parse("00-4bf92f3577b34da6a3ce929d0e0e4736-0000000000000000-01").is_none(), - "all-zero parent_id should be rejected" - ); - } - - #[test] - fn parse_traceparent_rejects_non_hex() { - assert!( - Traceparent::parse("00-4bf92f3577b34da6a3ce929d0e0e473g-00f067aa0ba902b7-01").is_none(), - "non-hex character in trace_id should be rejected" - ); - } - - #[test] - fn traceparent_format_preserves_trace_id_and_flags() { - let tp = Traceparent { - parent_id: "00f067aa0ba902b7".to_owned(), - trace_flags: "01".to_owned(), - trace_id: "4bf92f3577b34da6a3ce929d0e0e4736".to_owned(), - }; - let formatted = tp.format("abcdef1234567890"); - assert_eq!( - formatted, "00-4bf92f3577b34da6a3ce929d0e0e4736-abcdef1234567890-01", - "format should use new parent_id but preserve trace_id and flags" - ); - } - - // ------------------------------------------------------------------------- - // is_lowercase_hex - // ------------------------------------------------------------------------- - - #[test] - fn is_lowercase_hex_valid() { - assert!(is_lowercase_hex("0123456789abcdef"), "valid lowercase hex should pass"); - } - - #[test] - fn is_lowercase_hex_rejects_uppercase() { - assert!(!is_lowercase_hex("ABCDEF"), "uppercase hex should be rejected"); - } - - #[test] - fn is_lowercase_hex_rejects_non_hex() { - assert!(!is_lowercase_hex("ghijkl"), "non-hex characters should be rejected"); - } - - #[test] - fn is_lowercase_hex_empty_string() { - assert!(is_lowercase_hex(""), "empty string should pass (vacuously true)"); - } - - // ------------------------------------------------------------------------- - // Filter Lifecycle - // ------------------------------------------------------------------------- + use super::*; #[tokio::test] - async fn generates_new_trace_when_no_traceparent() { + async fn generates_new_trace_and_request_id_when_absent() { let filter = make_filter(""); let req = crate::test_utils::make_request(http::Method::GET, "/"); let mut ctx = crate::test_utils::make_filter_context(&req); let action = filter.on_request(&mut ctx).await.unwrap(); + assert!(matches!(action, FilterAction::Continue)); - assert!(matches!(action, FilterAction::Continue), "should continue"); + let tc = ctx.extensions.get::().expect("TraceContext stored"); + assert_eq!(tc.request_id().len(), 32); + assert_eq!(tc.flags(), "01"); - // Should have removed incoming traceparent and added a new one - let traceparent = find_extra_header(&ctx, "traceparent").expect("traceparent should be injected"); - let tp = Traceparent::parse(&traceparent).expect("injected traceparent should be well-formed"); - assert_eq!(tp.trace_flags, "01", "new trace should be sampled"); + let traceparent = find_extra_header(&ctx, "traceparent").expect("traceparent injected"); + let tp = parse_traceparent(&traceparent).expect("well-formed"); + assert_eq!(tp.flags, "01"); + assert_eq!( + find_extra_header(&ctx, "x-request-id").as_deref(), + Some(tc.request_id()) + ); } #[tokio::test] @@ -544,129 +352,208 @@ mod tests { http::header::HeaderValue::from_static("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"), ); let mut ctx = crate::test_utils::make_filter_context(&req); - drop(filter.on_request(&mut ctx).await.unwrap()); - let traceparent = find_extra_header(&ctx, "traceparent").expect("traceparent should be injected"); - let tp = Traceparent::parse(&traceparent).expect("injected traceparent should be well-formed"); - assert_eq!( - tp.trace_id, "4bf92f3577b34da6a3ce929d0e0e4736", - "should preserve incoming trace_id" - ); - assert_ne!( - tp.parent_id, "00f067aa0ba902b7", - "parent_id should be updated to proxy's span" - ); - assert_eq!(tp.trace_flags, "01", "should preserve trace_flags"); + let traceparent = find_extra_header(&ctx, "traceparent").unwrap(); + let tp = parse_traceparent(&traceparent).unwrap(); + assert_eq!(tp.trace_id, "4bf92f3577b34da6a3ce929d0e0e4736"); + let parts: Vec<&str> = traceparent.split('-').collect(); + assert_ne!(parts[2], "00f067aa0ba902b7"); + assert_eq!(tp.flags, "01"); } #[tokio::test] - async fn ignores_malformed_traceparent_and_creates_new_trace() { + async fn malformed_and_all_zero_traceparent_fall_back_to_new_trace() { + for bad in [ + "garbage-value", + "00-00000000000000000000000000000000-00f067aa0ba902b7-01", + "00-4bf92f3577b34da6a3ce929d0e0e4736-0000000000000000-01", + ] { + let filter = make_filter(""); + let mut req = crate::test_utils::make_request(http::Method::GET, "/"); + req.headers.insert( + http::header::HeaderName::from_static("traceparent"), + http::header::HeaderValue::from_str(bad).unwrap(), + ); + let mut ctx = crate::test_utils::make_filter_context(&req); + drop(filter.on_request(&mut ctx).await.unwrap()); + let traceparent = find_extra_header(&ctx, "traceparent").unwrap(); + let tp = parse_traceparent(&traceparent).unwrap(); + assert!(tp.flags == "01", "fallback trace should be sampled for {bad}"); + assert_ne!(tp.trace_id, "00000000000000000000000000000000"); + } + } + + #[tokio::test] + async fn masks_reserved_flags_on_join() { let filter = make_filter(""); let mut req = crate::test_utils::make_request(http::Method::GET, "/"); req.headers.insert( http::header::HeaderName::from_static("traceparent"), - http::header::HeaderValue::from_static("garbage-value"), + http::header::HeaderValue::from_static("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-03"), ); let mut ctx = crate::test_utils::make_filter_context(&req); - drop(filter.on_request(&mut ctx).await.unwrap()); - - let traceparent = find_extra_header(&ctx, "traceparent") - .expect("traceparent should be injected even when incoming is malformed"); - let tp = Traceparent::parse(&traceparent).expect("injected traceparent should be well-formed"); - assert_eq!(tp.trace_flags, "01", "new trace should be sampled"); + let traceparent = find_extra_header(&ctx, "traceparent").unwrap(); + assert!( + traceparent.ends_with("-01"), + "reserved bits must be masked: {traceparent}" + ); } #[tokio::test] - async fn forwards_tracestate_verbatim() { + async fn future_version_accepted_emits_version_00() { let filter = make_filter(""); let mut req = crate::test_utils::make_request(http::Method::GET, "/"); req.headers.insert( http::header::HeaderName::from_static("traceparent"), - http::header::HeaderValue::from_static("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"), - ); - req.headers.insert( - http::header::HeaderName::from_static("tracestate"), - http::header::HeaderValue::from_static("congo=t61rcWkgMzE,rojo=00f067aa0ba902b7"), + http::header::HeaderValue::from_static("02-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01-extra"), ); let mut ctx = crate::test_utils::make_filter_context(&req); - drop(filter.on_request(&mut ctx).await.unwrap()); - - let tracestate = find_extra_header(&ctx, "tracestate").expect("tracestate should be forwarded"); - assert_eq!( - tracestate, "congo=t61rcWkgMzE,rojo=00f067aa0ba902b7", - "tracestate should be preserved verbatim" - ); + let traceparent = find_extra_header(&ctx, "traceparent").unwrap(); + assert!(traceparent.starts_with("00-")); + assert!(traceparent.contains("4bf92f3577b34da6a3ce929d0e0e4736")); } #[tokio::test] - async fn no_tracestate_when_absent() { + async fn reuses_pending_request_id_from_earlier_filter() { let filter = make_filter(""); let req = crate::test_utils::make_request(http::Method::GET, "/"); let mut ctx = crate::test_utils::make_filter_context(&req); - + ctx.extra_request_headers + .push((Cow::Borrowed("x-request-id"), "from-request-id-filter".into())); drop(filter.on_request(&mut ctx).await.unwrap()); - assert!( - find_extra_header(&ctx, "tracestate").is_none(), - "tracestate should not be injected when absent from request" + let tc = ctx.extensions.get::().unwrap(); + assert_eq!(tc.request_id(), "from-request-id-filter"); + assert_eq!( + ctx.extra_request_headers + .iter() + .filter(|(n, _)| n.eq_ignore_ascii_case("x-request-id")) + .count(), + 1, + "must not duplicate pending x-request-id" ); } #[tokio::test] - async fn strips_tracestate_when_traceparent_absent() { + async fn pending_request_id_wins_over_conflicting_inbound_header() { let filter = make_filter(""); let mut req = crate::test_utils::make_request(http::Method::GET, "/"); - req.headers.insert( - http::header::HeaderName::from_static("tracestate"), - http::header::HeaderValue::from_static("congo=t61rcWkgMzE"), - ); + req.headers.insert("x-request-id", "client-request-id".parse().unwrap()); let mut ctx = crate::test_utils::make_filter_context(&req); + ctx.extra_request_headers + .push((Cow::Borrowed("x-request-id"), "from-request-id-filter".into())); drop(filter.on_request(&mut ctx).await.unwrap()); - assert!( - find_extra_header(&ctx, "tracestate").is_none(), - "tracestate should not be forwarded when traceparent is absent" - ); - let removed = ctx.request_headers_to_remove.iter().any(|h| h.as_str() == "tracestate"); - assert!( - removed, - "tracestate header should be removed when traceparent is absent" + let tc = ctx.extensions.get::().unwrap(); + assert_eq!(tc.request_id(), "from-request-id-filter"); + assert_eq!( + find_extra_header(&ctx, "x-request-id").as_deref(), + Some("from-request-id-filter") ); } #[tokio::test] - async fn strips_tracestate_when_traceparent_invalid() { + async fn idempotent_on_request_does_not_duplicate_headers() { let filter = make_filter(""); - let mut req = crate::test_utils::make_request(http::Method::GET, "/"); - req.headers.insert( - http::header::HeaderName::from_static("traceparent"), - http::header::HeaderValue::from_static("garbage-value"), + let req = crate::test_utils::make_request(http::Method::GET, "/"); + let mut ctx = crate::test_utils::make_filter_context(&req); + drop(filter.on_request(&mut ctx).await.unwrap()); + let first_count = ctx.extra_request_headers.len(); + drop(filter.on_request(&mut ctx).await.unwrap()); + assert_eq!( + ctx.extra_request_headers.len(), + first_count, + "second on_request must not duplicate pending headers" ); - req.headers.insert( - http::header::HeaderName::from_static("tracestate"), - http::header::HeaderValue::from_static("congo=t61rcWkgMzE"), + assert_eq!( + ctx.extra_request_headers + .iter() + .filter(|(n, _)| n.eq_ignore_ascii_case("traceparent")) + .count(), + 1 ); + assert_eq!( + ctx.extra_request_headers + .iter() + .filter(|(n, _)| n.eq_ignore_ascii_case("x-request-id")) + .count(), + 1 + ); + } + + #[test] + fn ensure_extra_header_preserves_existing_competing_value() { + let req = crate::test_utils::make_request(http::Method::GET, "/"); let mut ctx = crate::test_utils::make_filter_context(&req); + ctx.extra_request_headers + .push((Cow::Borrowed("traceparent"), "existing-traceparent".into())); - drop(filter.on_request(&mut ctx).await.unwrap()); + ensure_extra_header(&mut ctx, "traceparent", "new-traceparent"); - assert!( - find_extra_header(&ctx, "tracestate").is_none(), - "tracestate should not be forwarded when traceparent is invalid" + let values: Vec<_> = ctx + .extra_request_headers + .iter() + .filter(|(n, _)| n.eq_ignore_ascii_case("traceparent")) + .map(|(_, v)| v.as_str()) + .collect(); + assert_eq!(values, vec!["existing-traceparent"]); + } + + #[tokio::test] + async fn competing_pending_request_id_is_detected() { + let filter = make_filter(""); + let req = crate::test_utils::make_request(http::Method::GET, "/"); + let mut ctx = crate::test_utils::make_filter_context(&req); + drop(filter.on_request(&mut ctx).await.unwrap()); + let expected = ctx.extensions.get::().unwrap().request_id().to_owned(); + ctx.extra_request_headers + .push((Cow::Borrowed("x-request-id"), "later-competing-id".into())); + warn_competing_request_id(&ctx, &expected); + assert_eq!( + ctx.extensions.get::().unwrap().request_id(), + expected, + "competing extras are warn-only; TraceContext stays authoritative" ); - let removed = ctx.request_headers_to_remove.iter().any(|h| h.as_str() == "tracestate"); - assert!( - removed, - "tracestate header should be removed when traceparent is invalid" + } + + #[tokio::test] + async fn apply_trace_propagation_injects_fresh_span_same_trace() { + let filter = make_filter(""); + let req = crate::test_utils::make_request(http::Method::GET, "/"); + let mut ctx = crate::test_utils::make_filter_context(&req); + drop(filter.on_request(&mut ctx).await.unwrap()); + + let primary = find_extra_header(&ctx, "traceparent").unwrap(); + let primary_tp = parse_traceparent(&primary).unwrap(); + + let mut fw = FrameworkHeaders::new(); + ctx.apply_trace_propagation(&mut fw); + let fw_tp = fw + .iter() + .find(|(n, _)| n.as_str() == "traceparent") + .map(|(_, v)| v.to_str().unwrap().to_owned()) + .expect("framework traceparent"); + let fw_parsed = parse_traceparent(&fw_tp).unwrap(); + assert_eq!(fw_parsed.trace_id, primary_tp.trace_id); + assert_ne!( + fw_tp.split("-").nth(2).unwrap(), + primary.split("-").nth(2).unwrap(), + "each outbound hop must mint a fresh span id" ); + let fw_rid = fw + .iter() + .find(|(n, _)| n.as_str() == "x-request-id") + .map(|(_, v)| v.to_str().unwrap().to_owned()) + .unwrap(); + assert_eq!(fw_rid, ctx.extensions.get::().unwrap().request_id()); } #[tokio::test] - async fn combines_multiple_tracestate_headers() { + async fn forwards_tracestate_when_traceparent_valid() { let filter = make_filter(""); let mut req = crate::test_utils::make_request(http::Method::GET, "/"); req.headers.insert( @@ -675,40 +562,32 @@ mod tests { ); req.headers.insert( http::header::HeaderName::from_static("tracestate"), - http::header::HeaderValue::from_static("congo=t61rcWkgMzE"), - ); - req.headers.append( - http::header::HeaderName::from_static("tracestate"), - http::header::HeaderValue::from_static("rojo=00f067aa0ba902b7"), + http::header::HeaderValue::from_static("congo=t61rcWkgMzE,rojo=00f067aa0ba902b7"), ); let mut ctx = crate::test_utils::make_filter_context(&req); - drop(filter.on_request(&mut ctx).await.unwrap()); - - let tracestate = find_extra_header(&ctx, "tracestate").expect("tracestate should be forwarded"); assert_eq!( - tracestate, "congo=t61rcWkgMzE, rojo=00f067aa0ba902b7", - "multiple tracestate headers should be combined with comma separator" + find_extra_header(&ctx, "tracestate").as_deref(), + Some("congo=t61rcWkgMzE,rojo=00f067aa0ba902b7") ); } #[tokio::test] - async fn removes_incoming_traceparent_before_injecting() { + async fn strips_tracestate_when_traceparent_invalid() { let filter = make_filter(""); let mut req = crate::test_utils::make_request(http::Method::GET, "/"); req.headers.insert( http::header::HeaderName::from_static("traceparent"), - http::header::HeaderValue::from_static("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"), + http::header::HeaderValue::from_static("garbage"), + ); + req.headers.insert( + http::header::HeaderName::from_static("tracestate"), + http::header::HeaderValue::from_static("congo=t61rcWkgMzE"), ); let mut ctx = crate::test_utils::make_filter_context(&req); - drop(filter.on_request(&mut ctx).await.unwrap()); - - let removed = ctx - .request_headers_to_remove - .iter() - .any(|h| h.as_str() == "traceparent"); - assert!(removed, "incoming traceparent header should be removed"); + assert!(find_extra_header(&ctx, "tracestate").is_none()); + assert!(ctx.request_headers_to_remove.iter().any(|h| h.as_str() == "tracestate")); } #[tokio::test] @@ -720,48 +599,39 @@ mod tests { http::header::HeaderValue::from_static("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-00"), ); let mut ctx = crate::test_utils::make_filter_context(&req); - drop(filter.on_request(&mut ctx).await.unwrap()); - let traceparent = find_extra_header(&ctx, "traceparent").unwrap(); - let tp = Traceparent::parse(&traceparent).unwrap(); - assert_eq!(tp.trace_flags, "00", "unsampled flag should be preserved"); + assert!(traceparent.ends_with("-00")); + assert_eq!(ctx.extensions.get::().unwrap().flags(), "00"); } #[test] - fn from_config_empty_succeeds() { + fn from_config_empty_and_null_succeed() { let config = serde_yaml::Value::Mapping(serde_yaml::Mapping::new()); - let filter = TraceContextFilter::from_config(&config).unwrap(); - assert_eq!(filter.name(), "trace_context", "filter name should be trace_context"); - } - - #[test] - fn from_config_null_succeeds() { - let filter = TraceContextFilter::from_config(&serde_yaml::Value::Null).unwrap(); - assert_eq!(filter.name(), "trace_context", "filter name should be trace_context"); + assert_eq!( + TraceContextFilter::from_config(&config).unwrap().name(), + "trace_context" + ); + assert_eq!( + TraceContextFilter::from_config(&serde_yaml::Value::Null) + .unwrap() + .name(), + "trace_context" + ); } #[test] fn from_config_rejects_unknown_fields() { let config: serde_yaml::Value = serde_yaml::from_str("bogus: true").unwrap(); - assert!( - TraceContextFilter::from_config(&config).is_err(), - "unknown fields should be rejected" - ); + assert!(TraceContextFilter::from_config(&config).is_err()); } - // ------------------------------------------------------------------------- - // Test Utilities - // ------------------------------------------------------------------------- - - /// Build a [`TraceContextFilter`] from a YAML config string. fn make_filter(yaml: &str) -> TraceContextFilter { let config: serde_yaml::Value = serde_yaml::from_str(yaml).unwrap(); let _cfg: TraceContextFilterConfig = parse_filter_config("trace_context", &config).unwrap(); TraceContextFilter } - /// Find an extra request header by name (case-insensitive). fn find_extra_header(ctx: &HttpFilterContext<'_>, name: &str) -> Option { ctx.extra_request_headers .iter() diff --git a/filter/src/builtins/http/traffic_management/iterative_request_router/runner.rs b/filter/src/builtins/http/traffic_management/iterative_request_router/runner.rs index 6ead5f21f..966028b30 100644 --- a/filter/src/builtins/http/traffic_management/iterative_request_router/runner.rs +++ b/filter/src/builtins/http/traffic_management/iterative_request_router/runner.rs @@ -298,6 +298,7 @@ impl IrrStepRunner { }; let mut framework_headers = FrameworkHeaders::new(); framework_headers.set_depth(self.depth + 1); + filter_ctx.apply_trace_propagation(&mut framework_headers); let transport_budget = step_budget .checked_sub(step_started.elapsed()) .unwrap_or(Duration::ZERO); diff --git a/filter/src/context.rs b/filter/src/context.rs index a65a5c767..c75a0a6be 100644 --- a/filter/src/context.rs +++ b/filter/src/context.rs @@ -559,11 +559,95 @@ impl HttpFilterContext<'_> { self.filter_metadata.get(key).map(String::as_str) } - /// X-Request-ID header value, if present and valid UTF-8. + /// Request-scoped [`TraceContext`] id, else inbound `x-request-id`. + /// + /// [`TraceContext`]: crate::trace_context::TraceContext pub fn request_id(&self) -> Option<&str> { + if let Some(tc) = self.extensions.get::() { + return Some(tc.request_id()); + } self.request.headers.get("x-request-id").and_then(|v| v.to_str().ok()) } + /// Inject hop `x-request-id` and `traceparent` when [`TraceContext`] is present. + /// + /// [`TraceContext`]: crate::trace_context::TraceContext + pub fn apply_trace_propagation(&self, framework_headers: &mut praxis_core::subrequest::FrameworkHeaders) { + let Some(tc) = self.extensions.get::() else { + return; + }; + for (name, value) in &self.extra_request_headers { + if name.eq_ignore_ascii_case("x-request-id") && value != tc.request_id() { + tracing::warn!( + existing = %value, + expected = %tc.request_id(), + "competing x-request-id pending alongside TraceContext during sub-request propagation" + ); + } + } + if let Err(error) = tc.inject_into(framework_headers, self.id_generator, self.time_source) { + tracing::warn!(%error, "failed to inject trace correlation into framework headers"); + } + } + + /// Execute a buffered sub-request, injecting correlation when present. + /// + /// # Errors + /// + /// Returns [`SubRequestError`] if the client is missing or the exchange fails. + /// + /// [`SubRequestError`]: praxis_core::subrequest::SubRequestError + #[expect( + clippy::too_many_arguments, + reason = "mirrors SubRequestClient::execute including framework headers" + )] + pub async fn execute_subrequest( + &self, + peer: &pingora_core::upstreams::peer::HttpPeer, + request: &praxis_core::subrequest::SubRequest, + max_response_bytes: usize, + timeout: std::time::Duration, + mut framework_headers: praxis_core::subrequest::FrameworkHeaders, + ) -> Result { + let client = self.subrequest_client().ok_or_else(|| { + praxis_core::subrequest::SubRequestError::InvalidRequest( + "sub-request client is not available on this filter context".to_owned(), + ) + })?; + self.apply_trace_propagation(&mut framework_headers); + let fw = (!framework_headers.is_empty()).then_some(&framework_headers); + Box::pin(client.execute(peer, request, max_response_bytes, timeout, fw)).await + } + + /// Send a streaming sub-request, injecting correlation when present. + /// + /// # Errors + /// + /// Returns [`SubRequestError`] if the client is missing or the exchange fails. + /// + /// [`SubRequestError`]: praxis_core::subrequest::SubRequestError + #[expect( + clippy::too_many_arguments, + reason = "mirrors SubRequestClient::send_streaming including framework headers" + )] + pub async fn send_streaming_subrequest( + &self, + peer: &pingora_core::upstreams::peer::HttpPeer, + request: &praxis_core::subrequest::SubRequest, + timeout: std::time::Duration, + limits: praxis_core::subrequest::StreamLimits, + mut framework_headers: praxis_core::subrequest::FrameworkHeaders, + ) -> Result { + let client = self.subrequest_client().ok_or_else(|| { + praxis_core::subrequest::SubRequestError::InvalidRequest( + "sub-request client is not available on this filter context".to_owned(), + ) + })?; + self.apply_trace_propagation(&mut framework_headers); + let fw = (!framework_headers.is_empty()).then_some(&framework_headers); + Box::pin(client.send_streaming(peer, request, timeout, limits, fw)).await + } + /// Write a durable metadata value that persists across all phases. /// /// Keys should use dot-prefix namespacing @@ -1054,6 +1138,150 @@ mod tests { ); } + #[tokio::test] + async fn execute_subrequest_returns_error_without_client() { + let req = crate::test_utils::make_request(Method::GET, "/"); + let ctx = crate::test_utils::make_filter_context(&req); + let peer = pingora_core::upstreams::peer::HttpPeer::new("127.0.0.1:9".to_owned(), false, String::new()); + let subrequest = praxis_core::subrequest::SubRequest { + method: Method::GET, + uri: "/sub".parse().unwrap(), + headers: HeaderMap::new(), + body: bytes::Bytes::new(), + }; + + let result = ctx + .execute_subrequest( + &peer, + &subrequest, + 1024, + std::time::Duration::from_secs(1), + praxis_core::subrequest::FrameworkHeaders::new(), + ) + .await; + + assert!( + matches!(result, Err(praxis_core::subrequest::SubRequestError::InvalidRequest(message)) if message.contains("not available")), + "missing subrequest client should return InvalidRequest" + ); + } + + #[tokio::test] + #[allow( + clippy::significant_drop_tightening, + reason = "asserting error result without polling stream" + )] + async fn send_streaming_subrequest_returns_error_without_client() { + let req = crate::test_utils::make_request(Method::GET, "/"); + let ctx = crate::test_utils::make_filter_context(&req); + let peer = pingora_core::upstreams::peer::HttpPeer::new("127.0.0.1:9".to_owned(), false, String::new()); + let subrequest = praxis_core::subrequest::SubRequest { + method: Method::GET, + uri: "/sub".parse().unwrap(), + headers: HeaderMap::new(), + body: bytes::Bytes::new(), + }; + let limits = praxis_core::subrequest::StreamLimits { + idle_timeout: std::time::Duration::from_secs(1), + max_stream_duration: None, + max_total_bytes: None, + }; + + let result = ctx + .send_streaming_subrequest( + &peer, + &subrequest, + std::time::Duration::from_secs(1), + limits, + praxis_core::subrequest::FrameworkHeaders::new(), + ) + .await; + + assert!( + matches!(result, Err(praxis_core::subrequest::SubRequestError::InvalidRequest(message)) if message.contains("not available")), + "missing subrequest client should return InvalidRequest" + ); + } + + async fn capture_subrequest_headers() -> String { + use praxis_core::subrequest::{FrameworkHeaders, SubRequestClient, SubRequestConnector}; + + use crate::trace_context::TraceContext; + + let (addr, server) = start_header_capture_server().await; + let req = crate::test_utils::make_request(Method::GET, "/"); + let mut ctx = crate::test_utils::make_filter_context(&req); + ctx.extensions.insert(TraceContext::new( + "req-from-context".into(), + "4bf92f3577b34da6a3ce929d0e0e4736".into(), + "01".into(), + )); + let client = SubRequestClient::new(SubRequestConnector::new(1, None)); + ctx.subrequest_client = Some(&client); + + let peer = pingora_core::upstreams::peer::HttpPeer::new(addr, false, "localhost".into()); + let response = ctx + .execute_subrequest( + &peer, + &get_subrequest(), + 1024, + std::time::Duration::from_secs(5), + FrameworkHeaders::new(), + ) + .await + .unwrap(); + assert_eq!(response.status, 200); + server.await.unwrap() + } + + async fn start_header_capture_server() -> (std::net::SocketAddr, tokio::task::JoinHandle) { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { read_observed_request(listener).await }); + (addr, server) + } + + async fn read_observed_request(listener: tokio::net::TcpListener) -> String { + use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _}; + + let (mut stream, _) = listener.accept().await.unwrap(); + let mut buf = vec![0; 4096]; + let n = stream.read(&mut buf).await.unwrap(); + let observed = String::from_utf8_lossy(&buf[..n]).to_string(); + stream + .write_all( + b"HTTP/1.1 200 OK +content-length: 0 + +", + ) + .await + .unwrap(); + observed + } + + fn get_subrequest() -> praxis_core::subrequest::SubRequest { + praxis_core::subrequest::SubRequest { + method: Method::GET, + uri: "/sub".parse().unwrap(), + headers: HeaderMap::new(), + body: bytes::Bytes::new(), + } + } + + #[tokio::test] + async fn execute_subrequest_propagates_request_id() { + let observed = capture_subrequest_headers().await; + assert!(observed.contains("x-request-id: req-from-context"), "{observed}"); + } + + #[tokio::test] + async fn execute_subrequest_propagates_traceparent() { + let observed = capture_subrequest_headers().await; + assert!(observed.contains("traceparent: 00-"), "{observed}"); + assert!(observed.contains("4bf92f3577b34da6a3ce929d0e0e4736"), "{observed}"); + } + #[test] fn set_request_body_mode_upgrades_stream_to_stream_buffer() { let req = crate::test_utils::make_request(Method::GET, "/"); diff --git a/filter/src/lib.rs b/filter/src/lib.rs index dbe91d098..4a079d2e0 100644 --- a/filter/src/lib.rs +++ b/filter/src/lib.rs @@ -58,6 +58,7 @@ mod pipeline; mod registry; mod results; mod tcp_filter; +mod trace_context; pub use actions::{FilterAction, Rejection, StreamingResponseBody, StreamingTerminalResponse, TerminalResponse}; pub use any_filter::AnyFilter; @@ -101,6 +102,7 @@ pub use praxis_tls::TlsPeerIdentity; pub use registry::{FilterRegistry, SecurityClass}; pub use results::{FilterResultSet, matches_filter_result}; pub use tcp_filter::{TcpFilter, TcpFilterContext}; +pub use trace_context::TraceContext; // ----------------------------------------------------------------------------- // Custom Filter Registration diff --git a/filter/src/trace_context.rs b/filter/src/trace_context.rs new file mode 100644 index 000000000..bcfdd014c --- /dev/null +++ b/filter/src/trace_context.rs @@ -0,0 +1,360 @@ +// SPDX-License-Identifier: Apache-2.0 +// Copyright (c) 2026 Praxis Contributors + +//! Request-scoped correlation ids for forwarded requests and sub-requests. + +use http::{HeaderName, HeaderValue}; +use praxis_core::{ + id::IdGenerator, + subrequest::{FrameworkHeaders, SubRequestError}, + time::TimeSource, +}; + + +/// Header carrying the request correlation ID. +pub(crate) const REQUEST_ID: HeaderName = HeaderName::from_static("x-request-id"); + +/// Header carrying W3C trace context. +pub(crate) const TRACEPARENT: HeaderName = HeaderName::from_static("traceparent"); + +/// Header name for the request correlation ID. +pub(crate) const REQUEST_ID_HEADER: &str = "x-request-id"; + +/// Header name for W3C trace context. +pub(crate) const TRACEPARENT_HEADER: &str = "traceparent"; + +/// W3C version this proxy emits. +pub(crate) const VERSION: &str = "00"; + +/// Sampled `trace-flags` for a newly started trace. +pub(crate) const SAMPLED: &str = "01"; + +/// Version `00` defines only the sampled bit; other bits are zeroed on emit. +const SAMPLED_BIT: u8 = 0x01; + +/// W3C trace-id hex length. +pub(crate) const TRACE_ID_LEN: usize = 32; + +/// W3C span-id hex length. +pub(crate) const SPAN_ID_LEN: usize = 16; + +/// Number of base fields in a W3C `traceparent` header. +const BASE_FIELDS: usize = 4; + +/// W3C Trace Context section 2.2.2 forbids an all-zero trace-id. +const FALLBACK_TRACE_ID: &str = "00000000000000000000000000000001"; + +/// W3C Trace Context section 2.2.2 forbids an all-zero span-id. +const FALLBACK_SPAN_ID: &str = "0000000000000001"; + + +/// Request-scoped correlation ids. Outbound hops share the trace-id and mint a span-id. +#[derive(Clone, Debug, Eq, PartialEq)] +pub struct TraceContext { + /// W3C `trace-flags` value propagated to outbound hops. + flags: String, + /// Request correlation identifier propagated as `x-request-id`. + request_id: String, + /// W3C trace-id shared by all outbound hops for this request. + trace_id: String, +} + +impl TraceContext { + /// Construct from already-resolved parts. + #[must_use] + pub(crate) fn new(request_id: String, trace_id: String, flags: String) -> Self { + Self { + flags, + request_id, + trace_id, + } + } + + /// Continue a valid inbound trace. + #[must_use] + pub(crate) fn from_inbound(request_id: String, inbound: &InboundTrace) -> Self { + Self::new(request_id, inbound.trace_id.clone(), inbound.flags.clone()) + } + + /// Start a sampled trace. + #[must_use] + pub(crate) fn new_sampled(request_id: String, id_generator: &IdGenerator, time_source: &dyn TimeSource) -> Self { + Self::new( + request_id, + generate_trace_id(id_generator, time_source), + SAMPLED.to_owned(), + ) + } + + /// Correlation headers for one outbound hop. + #[must_use] + pub(crate) fn headers_for_hop( + &self, + id_generator: &IdGenerator, + time_source: &dyn TimeSource, + ) -> [(HeaderName, String); 2] { + [ + (REQUEST_ID, self.request_id.clone()), + (TRACEPARENT, self.traceparent_for_hop(id_generator, time_source)), + ] + } + + /// Resolved `x-request-id`. + #[must_use] + pub fn request_id(&self) -> &str { + &self.request_id + } + + /// Trace-id shared by every outbound hop. + #[must_use] + pub fn trace_id(&self) -> &str { + &self.trace_id + } + + /// W3C `trace-flags`. + #[must_use] + pub fn flags(&self) -> &str { + &self.flags + } + + /// Whether the sampled bit is set. + #[must_use] + pub fn sampled(&self) -> bool { + self.flags == SAMPLED + } + + /// Inject `x-request-id` and a hop `traceparent` into `fw`. + pub(crate) fn inject_into( + &self, + fw: &mut FrameworkHeaders, + id_generator: &IdGenerator, + time_source: &dyn TimeSource, + ) -> Result<(), SubRequestError> { + let [(request_id_name, request_id), (traceparent_name, traceparent)] = + self.headers_for_hop(id_generator, time_source); + fw.insert(request_id_name, header_value(REQUEST_ID_HEADER, &request_id)?)?; + fw.insert(traceparent_name, header_value(TRACEPARENT_HEADER, &traceparent)?) + } + + /// `traceparent` for one hop, with a fresh span-id. + #[must_use] + pub(crate) fn traceparent_for_hop(&self, id_generator: &IdGenerator, time_source: &dyn TimeSource) -> String { + let span_id = generate_span_id(id_generator, time_source); + let Self { flags, trace_id, .. } = self; + format!("{VERSION}-{trace_id}-{span_id}-{flags}") + } +} + + +/// Trace-id and flags continued from a valid inbound `traceparent`. +#[derive(Clone, Debug, Eq, PartialEq)] +pub(crate) struct InboundTrace { + /// Masked `trace-flags`. + pub flags: String, + + /// Shared 32-hex trace-id. + pub trace_id: String, +} + + +/// Generate a 16-hex span-id. +#[must_use] +pub(crate) fn generate_span_id(id_generator: &IdGenerator, time_source: &dyn TimeSource) -> String { + span_id_from(&generate_trace_id(id_generator, time_source)) +} + +/// Generate a 32-hex trace-id. +#[must_use] +pub(crate) fn generate_trace_id(id_generator: &IdGenerator, time_source: &dyn TimeSource) -> String { + sanitize_trace_id(&id_generator.generate(time_source)) +} + +/// Convert a string into a validated HTTP header value. +fn header_value(name: &str, value: &str) -> Result { + HeaderValue::from_str(value).map_err(|e| SubRequestError::InvalidRequest(format!("invalid {name} value: {e}"))) +} + +/// Return true when every character is ASCII zero. +fn is_all_zero(value: &str) -> bool { + value.bytes().all(|b| b == b'0') +} + +/// Return true when every character is lowercase hexadecimal. +fn is_lower_hex(value: &str) -> bool { + value.bytes().all(|b| b.is_ascii_digit() || (b'a'..=b'f').contains(&b)) +} + +/// Keep only supported W3C trace flag bits. +fn mask_flags(flags: &str) -> String { + let bits = u8::from_str_radix(flags, 16).unwrap_or(0); + format!("{:02x}", bits & SAMPLED_BIT) +} + +/// Parse a W3C `traceparent`. `None` if malformed, all-zero, or version `ff`. +#[must_use] +pub(crate) fn parse_traceparent(value: &str) -> Option { + let fields: Vec<&str> = value.split('-').collect(); + let [version, trace_id, span_id, flags] = fields.get(..BASE_FIELDS)? else { + return None; + }; + + if version.len() != 2 || !is_lower_hex(version) || *version == "ff" { + return None; + } + if fields.len() > BASE_FIELDS && *version == VERSION { + return None; + } + if trace_id.len() != TRACE_ID_LEN || !is_lower_hex(trace_id) || is_all_zero(trace_id) { + return None; + } + if span_id.len() != SPAN_ID_LEN || !is_lower_hex(span_id) || is_all_zero(span_id) { + return None; + } + if flags.len() != 2 || !is_lower_hex(flags) { + return None; + } + + Some(InboundTrace { + flags: mask_flags(flags), + trace_id: (*trace_id).to_owned(), + }) +} + +/// Coerce a generated ID into a W3C-valid trace-id (never all-zero). +fn sanitize_trace_id(id: &str) -> String { + if id.len() == TRACE_ID_LEN && is_lower_hex(id) && !is_all_zero(id) { + return id.to_owned(); + } + + let sanitized = format!("{id:0>TRACE_ID_LEN$.TRACE_ID_LEN$}") + .to_ascii_lowercase() + .replace(|c: char| !c.is_ascii_hexdigit(), "0"); + if sanitized.len() != TRACE_ID_LEN || is_all_zero(&sanitized) { + return FALLBACK_TRACE_ID.to_owned(); + } + sanitized +} + +/// Last 16 hex chars; the first 16 collide within the same microsecond. +fn span_id_from(trace_id: &str) -> String { + let span_id = trace_id.get(TRACE_ID_LEN - SPAN_ID_LEN..).unwrap_or(""); + if span_id.len() != SPAN_ID_LEN || is_all_zero(span_id) { + return FALLBACK_SPAN_ID.to_owned(); + } + span_id.to_owned() +} + + +#[cfg(test)] +#[expect(clippy::allow_attributes, reason = "blanket test suppressions")] +#[allow( + clippy::unwrap_used, + clippy::expect_used, + clippy::indexing_slicing, + clippy::panic, + reason = "tests" +)] +mod tests { + use std::time::Duration; + + use praxis_core::time::FixedTimeSource; + + use super::*; + + #[test] + fn parse_valid_traceparent() { + let tp = parse_traceparent("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01").unwrap(); + assert_eq!(tp.trace_id, "4bf92f3577b34da6a3ce929d0e0e4736"); + assert_eq!(tp.flags, "01"); + } + + #[test] + fn parse_traceparent_accepts_unsampled_flags() { + let parsed = parse_traceparent("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-00") + .expect("unsampled but well-formed traceparent should parse"); + assert_eq!(parsed.flags, "00"); + } + + #[test] + fn parse_masks_reserved_flags_to_sampled_bit() { + for (inbound, expected) in [("ff", "01"), ("fe", "00"), ("03", "01"), ("02", "00")] { + let value = format!("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-{inbound}"); + let parsed = parse_traceparent(&value).expect("well-formed flags should parse"); + assert_eq!(parsed.flags, expected); + } + } + + #[test] + fn parse_accepts_future_version_with_extra_fields() { + let tp = parse_traceparent("02-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01-extra-data") + .expect("future version with extra fields should be accepted"); + assert_eq!(tp.trace_id, "4bf92f3577b34da6a3ce929d0e0e4736"); + let ctx = TraceContext::from_inbound("req".into(), &tp); + let generator = IdGenerator::with_seed(1); + let ts = FixedTimeSource::new(Duration::from_micros(42)); + let [_, (_, traceparent)] = ctx.headers_for_hop(&generator, &ts); + assert!(traceparent.starts_with("00-")); + } + + #[test] + fn parse_rejects_all_zero_ids_and_malformed() { + assert!(parse_traceparent("00-00000000000000000000000000000000-00f067aa0ba902b7-01").is_none()); + assert!(parse_traceparent("00-4bf92f3577b34da6a3ce929d0e0e4736-0000000000000000-01").is_none()); + assert!(parse_traceparent("garbage-value").is_none()); + assert!(parse_traceparent("ff-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01").is_none()); + } + + #[test] + fn headers_for_hop_sets_correlation_headers() { + let ctx = TraceContext::new("abc123".into(), "4bf92f3577b34da6a3ce929d0e0e4736".into(), "01".into()); + let generator = IdGenerator::with_seed(1); + let ts = FixedTimeSource::new(Duration::from_micros(42)); + let headers = ctx.headers_for_hop(&generator, &ts); + assert_eq!(headers[0].0, REQUEST_ID); + assert_eq!(headers[0].1, "abc123"); + assert_eq!(headers[1].0, TRACEPARENT); + assert!(headers[1].1.starts_with("00-4bf92f3577b34da6a3ce929d0e0e4736-")); + assert!(headers[1].1.ends_with("-01")); + } + + #[test] + fn inject_into_sets_framework_correlation_headers() { + let ctx = TraceContext::new("abc123".into(), "4bf92f3577b34da6a3ce929d0e0e4736".into(), "01".into()); + let generator = IdGenerator::with_seed(1); + let ts = FixedTimeSource::new(Duration::from_micros(42)); + let mut fw = FrameworkHeaders::new(); + ctx.inject_into(&mut fw, &generator, &ts).unwrap(); + let entries: Vec<_> = fw + .iter() + .map(|(n, v)| (n.as_str().to_owned(), v.to_str().unwrap().to_owned())) + .collect(); + assert!(entries.iter().any(|(n, v)| n == "x-request-id" && v == "abc123")); + assert!( + entries + .iter() + .any(|(n, v)| n == "traceparent" && v.contains("4bf92f3577b34da6a3ce929d0e0e4736")) + ); + } + + #[test] + fn sanitized_ids_are_never_all_zero() { + for id in ["", "0", "00000000000000000000000000000000", "----", "zzzz"] { + let trace_id = sanitize_trace_id(id); + assert_eq!(trace_id.len(), TRACE_ID_LEN); + assert!(is_lower_hex(&trace_id)); + assert!(!is_all_zero(&trace_id)); + + let span_id = span_id_from(&trace_id); + assert_eq!(span_id.len(), SPAN_ID_LEN); + assert!(!is_all_zero(&span_id)); + } + } + + #[test] + fn generate_ids_have_expected_lengths() { + let generator = IdGenerator::with_seed(1); + let ts = FixedTimeSource::new(Duration::from_micros(42)); + assert_eq!(generate_trace_id(&generator, &ts).len(), 32); + assert_eq!(generate_span_id(&generator, &ts).len(), 16); + } +}