revoke-trace 0.3.0

Distributed tracing with OpenTelemetry for Revoke framework
Documentation
use crate::context::TraceContext;
use opentelemetry::propagation::{Extractor, Injector, TextMapPropagator};
use opentelemetry::trace::{SpanContext, SpanId, TraceContextExt, TraceFlags, TraceId, TraceState};
use opentelemetry_sdk::propagation::TraceContextPropagator;
use std::collections::HashMap;

/// HTTP 头部类型别名
pub type HttpHeaders = HashMap<String, String>;

/// 追踪传播器,用于在进程间传递追踪上下文
pub struct TracePropagator {
    propagator: Box<dyn TextMapPropagator + Send + Sync>,
}

impl TracePropagator {
    /// 创建默认的 W3C TraceContext 传播器
    pub fn new() -> Self {
        Self {
            propagator: Box::new(TraceContextPropagator::new()),
        }
    }


    /// 从 HTTP 头部提取追踪上下文
    pub fn extract(&self, headers: &HttpHeaders) -> Option<TraceContext> {
        let extractor = HeaderExtractor(headers);
        let context = self.propagator.extract(&extractor);

        let span = context.span();
        let span_context = span.span_context();

        if span_context.is_valid() {
            Some(TraceContext::from_span_context(&span_context))
        } else {
            None
        }
    }

    /// 将追踪上下文注入到 HTTP 头部
    pub fn inject(&self, context: &TraceContext, headers: &mut HttpHeaders) {
        let span_context = context.to_span_context();
        let otel_context = opentelemetry::Context::current().with_remote_span_context(span_context);

        let mut injector = HeaderInjector(headers);
        self.propagator.inject_context(&otel_context, &mut injector);
    }

    /// 从载体中提取追踪上下文(通用版本)
    pub fn extract_from<T: Extractor>(&self, carrier: &T) -> Option<TraceContext> {
        let context = self.propagator.extract(carrier);
        let span = context.span();
        let span_context = span.span_context();

        if span_context.is_valid() {
            Some(TraceContext::from_span_context(&span_context))
        } else {
            None
        }
    }

    /// 将追踪上下文注入到载体中(通用版本)
    pub fn inject_into<T: Injector>(&self, context: &TraceContext, carrier: &mut T) {
        let span_context = context.to_span_context();
        let otel_context = opentelemetry::Context::current().with_remote_span_context(span_context);

        self.propagator.inject_context(&otel_context, carrier);
    }
}

impl Default for TracePropagator {
    fn default() -> Self {
        Self::new()
    }
}

/// HTTP 头部提取器
struct HeaderExtractor<'a>(&'a HttpHeaders);

impl<'a> Extractor for HeaderExtractor<'a> {
    fn get(&self, key: &str) -> Option<&str> {
        self.0.get(key).map(|v| v.as_str())
    }

    fn keys(&self) -> Vec<&str> {
        self.0.keys().map(|k| k.as_str()).collect()
    }
}

/// HTTP 头部注入器
struct HeaderInjector<'a>(&'a mut HttpHeaders);

impl<'a> Injector for HeaderInjector<'a> {
    fn set(&mut self, key: &str, value: String) {
        self.0.insert(key.to_string(), value);
    }
}

/// 用于手动解析和构建追踪头部的辅助函数
pub mod manual {
    use super::*;

    /// 解析 W3C traceparent 头部
    /// 格式: version-trace_id-span_id-trace_flags
    pub fn parse_traceparent(value: &str) -> Option<SpanContext> {
        let parts: Vec<&str> = value.split('-').collect();
        if parts.len() != 4 {
            return None;
        }

        let version = parts[0];
        if version != "00" {
            return None; // 只支持版本 00
        }

        let trace_id = TraceId::from_hex(parts[1]).ok()?;
        let span_id = SpanId::from_hex(parts[2]).ok()?;
        let trace_flags = u8::from_str_radix(parts[3], 16).ok()?;

        Some(SpanContext::new(
            trace_id,
            span_id,
            TraceFlags::new(trace_flags),
            false,
            TraceState::default(),
        ))
    }

    /// 构建 W3C traceparent 头部
    pub fn build_traceparent(context: &SpanContext) -> String {
        format!(
            "00-{}-{}-{:02x}",
            context.trace_id(),
            context.span_id(),
            context.trace_flags().to_u8()
        )
    }

    /// 解析 B3 单头部格式
    /// 格式: trace_id-span_id-sampling_decision-parent_span_id
    pub fn parse_b3_single(value: &str) -> Option<SpanContext> {
        let parts: Vec<&str> = value.split('-').collect();
        if parts.len() < 2 {
            return None;
        }

        let trace_id = TraceId::from_hex(parts[0]).ok()?;
        let span_id = SpanId::from_hex(parts[1]).ok()?;

        let trace_flags = if parts.len() > 2 && parts[2] == "1" {
            TraceFlags::SAMPLED
        } else {
            TraceFlags::default()
        };

        Some(SpanContext::new(
            trace_id,
            span_id,
            trace_flags,
            false,
            TraceState::default(),
        ))
    }

    /// 构建 B3 单头部格式
    pub fn build_b3_single(context: &SpanContext) -> String {
        let sampled = if context.trace_flags().is_sampled() {
            "1"
        } else {
            "0"
        };
        format!("{}-{}-{}", context.trace_id(), context.span_id(), sampled)
    }
}

// 添加缺失的依赖项到 Cargo.toml
// 需要在 Cargo.toml 中添加:
// rand = "0.8"

#[cfg(test)]
mod tests {
    use super::*;

    #[test]
    fn test_propagator_extract_inject() {
        let propagator = TracePropagator::new();

        // 创建一个追踪上下文
        let context;

        // 创建一个有效的 span context
        let trace_id = TraceId::from_bytes([1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16]);
        let span_id = SpanId::from_bytes([1, 2, 3, 4, 5, 6, 7, 8]);
        let span_context = SpanContext::new(
            trace_id,
            span_id,
            TraceFlags::SAMPLED,
            false,
            TraceState::default(),
        );
        context = TraceContext::from_span_context(&span_context);

        // 注入到头部
        let mut headers = HttpHeaders::new();
        propagator.inject(&context, &mut headers);

        // 应该包含 traceparent 头部
        assert!(headers.contains_key("traceparent"));

        // 从头部提取
        let extracted = propagator.extract(&headers).unwrap();
        assert_eq!(extracted.trace_id(), context.trace_id());
        assert_eq!(extracted.span_id(), context.span_id());
    }

    #[test]
    fn test_manual_traceparent() {
        let trace_id = TraceId::from_bytes([1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16]);
        let span_id = SpanId::from_bytes([1, 2, 3, 4, 5, 6, 7, 8]);
        let span_context = SpanContext::new(
            trace_id,
            span_id,
            TraceFlags::SAMPLED,
            false,
            TraceState::default(),
        );

        let traceparent = manual::build_traceparent(&span_context);
        let parsed = manual::parse_traceparent(&traceparent).unwrap();

        assert_eq!(parsed.trace_id(), trace_id);
        assert_eq!(parsed.span_id(), span_id);
        assert!(parsed.trace_flags().is_sampled());
    }

    #[test]
    fn test_manual_b3() {
        let trace_id = TraceId::from_bytes([1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16]);
        let span_id = SpanId::from_bytes([1, 2, 3, 4, 5, 6, 7, 8]);
        let span_context = SpanContext::new(
            trace_id,
            span_id,
            TraceFlags::SAMPLED,
            false,
            TraceState::default(),
        );

        let b3 = manual::build_b3_single(&span_context);
        let parsed = manual::parse_b3_single(&b3).unwrap();

        assert_eq!(parsed.trace_id(), trace_id);
        assert_eq!(parsed.span_id(), span_id);
        assert!(parsed.trace_flags().is_sampled());
    }
}