use std::convert::Infallible;
use std::fmt;
use axum::extract::Request;
use axum::http::{HeaderMap, HeaderName, HeaderValue};
use axum::response::Response;
use tower::{Layer, Service};
pub const TRACEPARENT: HeaderName = HeaderName::from_static("traceparent");
pub const TRACESTATE: HeaderName = HeaderName::from_static("tracestate");
pub const FLAG_SAMPLED: u8 = 0x01;
pub const MAX_TRACESTATE_MEMBERS: usize = 32;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct TraceParent {
trace_id: [u8; 16],
parent_id: [u8; 8],
flags: u8,
}
impl TraceParent {
#[must_use]
pub fn root() -> Self {
let mut bytes = [0_u8; 24];
fill_unique(&mut bytes);
let mut trace_id = [0_u8; 16];
let mut parent_id = [0_u8; 8];
trace_id.copy_from_slice(&bytes[..16]);
parent_id.copy_from_slice(&bytes[16..]);
if trace_id == [0; 16] {
trace_id[15] = 1;
}
if parent_id == [0; 8] {
parent_id[7] = 1;
}
Self {
trace_id,
parent_id,
flags: FLAG_SAMPLED,
}
}
#[must_use]
pub fn child(&self) -> Self {
let mut bytes = [0_u8; 8];
fill_unique(&mut bytes);
if bytes == [0; 8] {
bytes[7] = 1;
}
Self {
trace_id: self.trace_id,
parent_id: bytes,
flags: self.flags,
}
}
#[must_use]
pub fn parse(value: &str) -> Option<Self> {
let value = value.trim();
if value.len() != 55 {
return None;
}
let bytes = value.as_bytes();
if bytes[2] != b'-' || bytes[35] != b'-' || bytes[52] != b'-' {
return None;
}
let version = hex_byte(&bytes[0..2])?;
if version != 0 {
return None;
}
let mut trace_id = [0_u8; 16];
for (index, slot) in trace_id.iter_mut().enumerate() {
*slot = hex_byte(&bytes[3 + index * 2..5 + index * 2])?;
}
let mut parent_id = [0_u8; 8];
for (index, slot) in parent_id.iter_mut().enumerate() {
*slot = hex_byte(&bytes[36 + index * 2..38 + index * 2])?;
}
if trace_id == [0; 16] || parent_id == [0; 8] {
return None;
}
let flags = hex_byte(&bytes[53..55])?;
Some(Self {
trace_id,
parent_id,
flags,
})
}
#[must_use]
pub fn trace_id(&self) -> String {
hex(&self.trace_id)
}
#[must_use]
pub fn parent_id(&self) -> String {
hex(&self.parent_id)
}
#[must_use]
pub fn sampled(&self) -> bool {
self.flags & FLAG_SAMPLED != 0
}
#[must_use]
pub fn to_header_value(&self) -> HeaderValue {
HeaderValue::from_str(&self.to_string()).unwrap_or(HeaderValue::from_static(""))
}
}
impl fmt::Display for TraceParent {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"00-{}-{}-{:02x}",
hex(&self.trace_id),
hex(&self.parent_id),
self.flags
)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TraceState(String);
impl TraceState {
#[must_use]
pub fn parse(value: &str) -> Option<Self> {
let value = value.trim();
if value.is_empty() {
return None;
}
if value.len() > 512 {
return None;
}
let members = value.split(',').filter(|m| !m.trim().is_empty()).count();
if members == 0 || members > MAX_TRACESTATE_MEMBERS {
return None;
}
if !value.bytes().all(|b| (0x20..=0x7e).contains(&b)) {
return None;
}
Some(Self(value.to_string()))
}
#[must_use]
pub fn as_str(&self) -> &str {
&self.0
}
}
impl fmt::Display for TraceState {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.0)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TraceContext {
parent: TraceParent,
state: Option<TraceState>,
continued: bool,
}
impl TraceContext {
#[must_use]
pub fn from_headers(headers: &HeaderMap) -> Self {
let parent = headers
.get(&TRACEPARENT)
.and_then(|value| value.to_str().ok())
.and_then(TraceParent::parse);
match parent {
Some(parent) => Self {
parent,
state: headers
.get(&TRACESTATE)
.and_then(|value| value.to_str().ok())
.and_then(TraceState::parse),
continued: true,
},
None => Self {
parent: TraceParent::root(),
state: None,
continued: false,
},
}
}
#[must_use]
pub fn parent(&self) -> &TraceParent {
&self.parent
}
#[must_use]
pub fn state(&self) -> Option<&TraceState> {
self.state.as_ref()
}
#[must_use]
pub fn continued(&self) -> bool {
self.continued
}
#[must_use]
pub fn outbound_headers(&self) -> HeaderMap {
let mut headers = HeaderMap::new();
headers.insert(TRACEPARENT, self.parent.child().to_header_value());
if let Some(state) = &self.state
&& let Ok(value) = HeaderValue::from_str(state.as_str())
{
headers.insert(TRACESTATE, value);
}
headers
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct TraceContextLayer;
impl<S> Layer<S> for TraceContextLayer {
type Service = TraceContextService<S>;
fn layer(&self, inner: S) -> Self::Service {
TraceContextService { inner }
}
}
#[derive(Debug, Clone)]
pub struct TraceContextService<S> {
inner: S,
}
impl<S> Service<Request> for TraceContextService<S>
where
S: Service<Request, Response = Response, Error = Infallible> + Clone + Send + 'static,
S::Future: Send + 'static,
{
type Response = Response;
type Error = Infallible;
type Future = std::pin::Pin<
Box<dyn std::future::Future<Output = Result<Self::Response, Self::Error>> + Send>,
>;
fn poll_ready(
&mut self,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, mut request: Request) -> Self::Future {
let context = TraceContext::from_headers(request.headers());
let trace_id = context.parent().trace_id();
let parent_span_id = context.parent().parent_id();
let continued = context.continued();
request.extensions_mut().insert(context);
let clone = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, clone);
Box::pin(async move {
let span = tracing::info_span!(
super::REQUEST,
trace_id = %trace_id,
parent_span_id = %parent_span_id,
continued_trace = continued,
);
let _entered = span.enter();
inner.call(request).await
})
}
}
fn hex(bytes: &[u8]) -> String {
const DIGITS: &[u8; 16] = b"0123456789abcdef";
let mut out = String::with_capacity(bytes.len() * 2);
for byte in bytes {
out.push(DIGITS[usize::from(byte >> 4)] as char);
out.push(DIGITS[usize::from(byte & 0x0f)] as char);
}
out
}
fn hex_byte(pair: &[u8]) -> Option<u8> {
fn nibble(b: u8) -> Option<u8> {
match b {
b'0'..=b'9' => Some(b - b'0'),
b'a'..=b'f' => Some(b - b'a' + 10),
_ => None,
}
}
Some((nibble(*pair.first()?)? << 4) | nibble(*pair.get(1)?)?)
}
fn fill_unique(bytes: &mut [u8]) {
#[cfg(feature = "oauth")]
if getrandom::fill(bytes).is_ok() {
return;
}
use std::sync::atomic::{AtomicU64, Ordering};
static COUNTER: AtomicU64 = AtomicU64::new(0);
let local = 0_u8;
let mut seed = COUNTER.fetch_add(1, Ordering::Relaxed)
^ (std::ptr::from_ref(&local) as u64)
^ std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map_or(0, |d| d.as_nanos() as u64);
for chunk in bytes.chunks_mut(8) {
seed = seed.wrapping_add(0x9e37_79b9_7f4a_7c15);
let mut mixed = seed;
mixed = (mixed ^ (mixed >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
mixed = (mixed ^ (mixed >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
mixed ^= mixed >> 31;
let source = mixed.to_le_bytes();
chunk.copy_from_slice(&source[..chunk.len()]);
}
}
#[cfg(test)]
mod tests {
use super::*;
const SAMPLE: &str = "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01";
#[test]
fn a_well_formed_traceparent_round_trips() {
let parsed = TraceParent::parse(SAMPLE).expect("valid header");
assert_eq!(parsed.trace_id(), "4bf92f3577b34da6a3ce929d0e0e4736");
assert_eq!(parsed.parent_id(), "00f067aa0ba902b7");
assert!(parsed.sampled());
assert_eq!(parsed.to_string(), SAMPLE);
}
#[test]
fn an_unsampled_flag_is_preserved() {
let header = "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-00";
let parsed = TraceParent::parse(header).expect("valid header");
assert!(!parsed.sampled());
assert_eq!(parsed.to_string(), header);
}
#[test]
fn malformed_traceparents_are_all_rejected() {
for header in [
"",
"garbage",
"00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-0",
"00-00000000000000000000000000000000-00f067aa0ba902b7-01",
"00-4bf92f3577b34da6a3ce929d0e0e4736-0000000000000000-01",
"01-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01",
"ff-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01",
"00-4BF92F3577B34DA6A3CE929D0E0E4736-00f067aa0ba902b7-01",
"00:4bf92f3577b34da6a3ce929d0e0e4736:00f067aa0ba902b7:01",
] {
assert!(TraceParent::parse(header).is_none(), "accepted {header:?}");
}
}
#[test]
fn a_child_keeps_the_trace_and_changes_the_span() {
let parent = TraceParent::parse(SAMPLE).expect("valid header");
let child = parent.child();
assert_eq!(child.trace_id(), parent.trace_id());
assert_ne!(child.parent_id(), parent.parent_id());
assert_eq!(child.sampled(), parent.sampled());
}
#[test]
fn two_roots_do_not_collide() {
let one = TraceParent::root();
let two = TraceParent::root();
assert_ne!(one.trace_id(), two.trace_id());
assert_ne!(one.trace_id(), "0".repeat(32));
}
#[test]
fn a_missing_traceparent_starts_a_root() {
let context = TraceContext::from_headers(&HeaderMap::new());
assert!(!context.continued());
assert!(context.state().is_none());
}
#[test]
fn a_tracestate_without_a_traceparent_is_discarded() {
let mut headers = HeaderMap::new();
headers.insert(TRACESTATE, HeaderValue::from_static("vendor=value"));
let context = TraceContext::from_headers(&headers);
assert!(!context.continued());
assert!(context.state().is_none());
}
#[test]
fn a_valid_pair_is_joined_and_carried_forward() {
let mut headers = HeaderMap::new();
headers.insert(TRACEPARENT, HeaderValue::from_static(SAMPLE));
headers.insert(TRACESTATE, HeaderValue::from_static("vendor=value"));
let context = TraceContext::from_headers(&headers);
assert!(context.continued());
assert_eq!(
context.state().map(TraceState::as_str),
Some("vendor=value")
);
let outbound = context.outbound_headers();
let forwarded = outbound
.get(&TRACEPARENT)
.and_then(|v| v.to_str().ok())
.and_then(TraceParent::parse)
.expect("a valid outbound header");
assert_eq!(forwarded.trace_id(), context.parent().trace_id());
assert_ne!(forwarded.parent_id(), context.parent().parent_id());
assert_eq!(outbound.get(&TRACESTATE).unwrap(), "vendor=value");
}
#[test]
fn an_oversized_or_hostile_tracestate_is_refused() {
let too_many = (0..40)
.map(|i| format!("v{i}=x"))
.collect::<Vec<_>>()
.join(",");
assert!(TraceState::parse(&too_many).is_none());
assert!(TraceState::parse(&"a=b,".repeat(400)).is_none());
assert!(TraceState::parse("vendor=value\r\ninjected: yes").is_none());
assert!(TraceState::parse(" ").is_none());
assert!(TraceState::parse("vendor=value").is_some());
}
}