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),
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))
}
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) => {
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),
};
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()
}
}
#[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(),
}
}
}
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;
}
}
assert!(sampled_count > 400 && sampled_count < 600);
}
#[test]
fn test_trace_id_ratio_sampling() {
let sampler = TraceSampler::trace_id_ratio(0.5);
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);
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]);
let (decision, _) = sampler.should_sample(&trace_id, "test", None);
assert_eq!(decision, SamplingDecision::Drop);
let (decision_with_parent, _) = sampler.should_sample(&trace_id, "test", Some(true));
assert_eq!(decision_with_parent, SamplingDecision::RecordOnly);
}
}