revoke-trace 0.3.0

Distributed tracing with OpenTelemetry for Revoke framework
Documentation
use opentelemetry::trace::{SamplingDecision as OtelSamplingDecision, SamplingResult, TraceId};
use opentelemetry_sdk::trace::ShouldSample;
use std::collections::HashMap;
use std::sync::Arc;

/// 采样决策
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SamplingDecision {
    /// 记录并采样
    RecordAndSample,
    /// 仅记录,不采样
    RecordOnly,
    /// 丢弃
    Drop,
}

impl From<SamplingDecision> for OtelSamplingDecision {
    fn from(decision: SamplingDecision) -> Self {
        match decision {
            SamplingDecision::RecordAndSample => OtelSamplingDecision::RecordAndSample,
            SamplingDecision::RecordOnly => OtelSamplingDecision::RecordOnly,
            SamplingDecision::Drop => OtelSamplingDecision::Drop,
        }
    }
}

/// 采样策略
#[derive(Clone)]
pub enum SamplingStrategy {
    /// 始终采样
    AlwaysOn,
    /// 从不采样
    AlwaysOff,
    /// 基于概率采样
    Probability(f64),
    /// 基于 trace ID 的确定性采样
    TraceIdRatio(f64),
    /// 自定义采样函数
    Custom(Arc<dyn Fn(&TraceId, &str) -> SamplingDecision + Send + Sync>),
}

impl std::fmt::Debug for SamplingStrategy {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        match self {
            Self::AlwaysOn => write!(f, "AlwaysOn"),
            Self::AlwaysOff => write!(f, "AlwaysOff"),
            Self::Probability(rate) => f.debug_tuple("Probability").field(rate).finish(),
            Self::TraceIdRatio(ratio) => f.debug_tuple("TraceIdRatio").field(ratio).finish(),
            Self::Custom(_) => write!(f, "Custom(...)"),
        }
    }
}

/// 追踪采样器
#[derive(Clone, Debug)]
pub struct TraceSampler {
    strategy: SamplingStrategy,
    attributes: HashMap<String, String>,
}

impl TraceSampler {
    /// 创建新的采样器
    pub fn new(strategy: SamplingStrategy) -> Self {
        Self {
            strategy,
            attributes: HashMap::new(),
        }
    }

    /// 创建始终采样的采样器
    pub fn always_on() -> Self {
        Self::new(SamplingStrategy::AlwaysOn)
    }

    /// 创建从不采样的采样器
    pub fn always_off() -> Self {
        Self::new(SamplingStrategy::AlwaysOff)
    }

    /// 创建概率采样器
    pub fn probability(rate: f64) -> Self {
        let rate = rate.clamp(0.0, 1.0);
        Self::new(SamplingStrategy::Probability(rate))
    }

    /// 创建基于 trace ID 的采样器
    pub fn trace_id_ratio(ratio: f64) -> Self {
        let ratio = ratio.clamp(0.0, 1.0);
        Self::new(SamplingStrategy::TraceIdRatio(ratio))
    }

    /// 添加采样属性
    pub fn with_attribute(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
        self.attributes.insert(key.into(), value.into());
        self
    }

    /// 做出采样决策
    pub fn should_sample(
        &self,
        trace_id: &TraceId,
        name: &str,
        parent_sampled: Option<bool>,
    ) -> (SamplingDecision, HashMap<String, String>) {
        let decision = match &self.strategy {
            SamplingStrategy::AlwaysOn => SamplingDecision::RecordAndSample,
            SamplingStrategy::AlwaysOff => SamplingDecision::Drop,
            SamplingStrategy::Probability(rate) => {
                if rand::random::<f64>() < *rate {
                    SamplingDecision::RecordAndSample
                } else {
                    SamplingDecision::Drop
                }
            }
            SamplingStrategy::TraceIdRatio(ratio) => {
                // 基于 trace ID 的确定性采样
                let trace_id_bytes = trace_id.to_bytes();
                let hash = u64::from_be_bytes([
                    trace_id_bytes[8],
                    trace_id_bytes[9],
                    trace_id_bytes[10],
                    trace_id_bytes[11],
                    trace_id_bytes[12],
                    trace_id_bytes[13],
                    trace_id_bytes[14],
                    trace_id_bytes[15],
                ]);
                let threshold = (ratio * u64::MAX as f64) as u64;

                if hash < threshold {
                    SamplingDecision::RecordAndSample
                } else {
                    SamplingDecision::Drop
                }
            }
            SamplingStrategy::Custom(sampler) => sampler(trace_id, name),
        };

        // 如果有父 span 且父 span 被采样,则子 span 也应该被采样
        let final_decision = if let Some(true) = parent_sampled {
            match decision {
                SamplingDecision::Drop => SamplingDecision::RecordOnly,
                _ => decision,
            }
        } else {
            decision
        };

        (final_decision, self.attributes.clone())
    }
}

impl Default for TraceSampler {
    fn default() -> Self {
        Self::always_on()
    }
}

/// OpenTelemetry 采样器适配器
#[derive(Debug, Clone)]
pub struct OtelSamplerAdapter {
    sampler: TraceSampler,
}

impl OtelSamplerAdapter {
    pub fn new(sampler: TraceSampler) -> Self {
        Self { sampler }
    }
}

impl ShouldSample for OtelSamplerAdapter {
    fn should_sample(
        &self,
        parent_context: Option<&opentelemetry::Context>,
        trace_id: TraceId,
        name: &str,
        _span_kind: &opentelemetry::trace::SpanKind,
        _attributes: &[opentelemetry::KeyValue],
        _links: &[opentelemetry::trace::Link],
    ) -> SamplingResult {
        use opentelemetry::trace::TraceContextExt;

        let parent_sampled = parent_context.and_then(|ctx| {
            let span = ctx.span();
            let span_context = span.span_context();
            if span_context.is_valid() {
                Some(span_context.trace_flags().is_sampled())
            } else {
                None
            }
        });

        let (decision, attributes) = self.sampler.should_sample(&trace_id, name, parent_sampled);

        let otel_attributes = attributes
            .into_iter()
            .map(|(k, v)| opentelemetry::KeyValue::new(k, v))
            .collect();

        SamplingResult {
            decision: decision.into(),
            attributes: otel_attributes,
            trace_state: Default::default(),
        }
    }
}

/// 添加 rand 依赖
use rand;

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

    #[test]
    fn test_sampling_strategies() {
        let always_on = TraceSampler::always_on();
        let always_off = TraceSampler::always_off();
        let probability = TraceSampler::probability(0.5);

        let trace_id = TraceId::from_bytes([1; 16]);

        let (on_decision, _) = always_on.should_sample(&trace_id, "test", None);
        assert_eq!(on_decision, SamplingDecision::RecordAndSample);

        let (off_decision, _) = always_off.should_sample(&trace_id, "test", None);
        assert_eq!(off_decision, SamplingDecision::Drop);

        // 概率采样应该有时采样,有时不采样
        let mut sampled_count = 0;
        for _ in 0..1000 {
            let (decision, _) = probability.should_sample(&trace_id, "test", None);
            if decision == SamplingDecision::RecordAndSample {
                sampled_count += 1;
            }
        }

        // 应该大约是 50%
        assert!(sampled_count > 400 && sampled_count < 600);
    }

    #[test]
    fn test_trace_id_ratio_sampling() {
        let sampler = TraceSampler::trace_id_ratio(0.5);

        // 使用不同的 trace ID 测试
        let trace_id1 = TraceId::from_bytes([0; 16]);
        let trace_id2 = TraceId::from_bytes([255; 16]);

        let (decision1, _) = sampler.should_sample(&trace_id1, "test", None);
        let (decision2, _) = sampler.should_sample(&trace_id2, "test", None);

        // 确定性采样:相同的 trace ID 应该总是得到相同的决策
        let (decision1_repeat, _) = sampler.should_sample(&trace_id1, "test", None);
        let (decision2_repeat, _) = sampler.should_sample(&trace_id2, "test", None);
        assert_eq!(decision1, decision1_repeat);
        assert_eq!(decision2, decision2_repeat);
    }

    #[test]
    fn test_parent_sampling() {
        let sampler = TraceSampler::always_off();

        let trace_id = TraceId::from_bytes([1; 16]);

        // 没有父 span 时应该丢弃
        let (decision, _) = sampler.should_sample(&trace_id, "test", None);
        assert_eq!(decision, SamplingDecision::Drop);

        // 父 span 被采样时应该至少记录
        let (decision_with_parent, _) = sampler.should_sample(&trace_id, "test", Some(true));
        assert_eq!(decision_with_parent, SamplingDecision::RecordOnly);
    }
}