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;
pub type HttpHeaders = HashMap<String, String>;
pub struct TracePropagator {
propagator: Box<dyn TextMapPropagator + Send + Sync>,
}
impl TracePropagator {
pub fn new() -> Self {
Self {
propagator: Box::new(TraceContextPropagator::new()),
}
}
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
}
}
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()
}
}
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()
}
}
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::*;
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; }
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(),
))
}
pub fn build_traceparent(context: &SpanContext) -> String {
format!(
"00-{}-{}-{:02x}",
context.trace_id(),
context.span_id(),
context.trace_flags().to_u8()
)
}
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(),
))
}
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)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_propagator_extract_inject() {
let propagator = TracePropagator::new();
let 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);
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());
}
}