use std::collections::HashMap;
use opentelemetry::Context;
use opentelemetry::propagation::{Extractor, Injector, TextMapPropagator};
use opentelemetry::trace::TraceContextExt as _;
#[derive(Debug, Clone)]
pub struct TraceStateEnricher {
additional_entries: Vec<(String, String)>,
}
impl TraceStateEnricher {
pub fn new(tracestate_str: &str) -> Self {
Self {
additional_entries: parse_pairs(tracestate_str).into_iter().collect(),
}
}
}
impl TextMapPropagator for TraceStateEnricher {
fn inject_context(&self, cx: &Context, injector: &mut dyn Injector) {
if self.additional_entries.is_empty() {
return;
}
let span = cx.span();
let span_context = span.span_context();
if !span_context.is_valid() {
return;
}
let mut trace_state = span_context.trace_state().clone();
for (key, value) in &self.additional_entries {
trace_state = trace_state
.insert(key.clone(), value.clone())
.unwrap_or(trace_state);
}
let header = trace_state.header();
if !header.is_empty() {
injector.set("tracestate", header);
}
}
fn extract_with_context(&self, cx: &Context, _extractor: &dyn Extractor) -> Context {
cx.clone()
}
fn fields(&self) -> opentelemetry::propagation::text_map_propagator::FieldIter<'_> {
static FIELDS: &[String] = &[];
opentelemetry::propagation::text_map_propagator::FieldIter::new(FIELDS)
}
}
pub fn parse_pairs(input: &str) -> HashMap<String, String> {
if input.is_empty() {
return HashMap::default();
}
input
.split(',')
.filter_map(|pair| {
let pair = pair.trim();
if pair.is_empty() {
return None;
}
let mut parts = pair.splitn(2, '=');
match (parts.next(), parts.next()) {
(Some(k), Some(v)) if !k.is_empty() && !v.is_empty() => {
Some((k.trim().to_owned(), v.trim().to_owned()))
}
_ => {
tracing::debug!("Ignoring malformed tracestate pair: '{}'", pair);
None
}
}
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_pairs_resilient() {
let result = parse_pairs("my_id=id123,env=prod");
assert_eq!(result.len(), 2);
assert_eq!(result.get("my_id"), Some(&"id123".to_owned()));
assert_eq!(result.get("env"), Some(&"prod".to_owned()));
let result = parse_pairs("valid=ok,invalid,key=,=value,also_valid=good");
assert_eq!(result.len(), 2);
assert_eq!(result.get("valid"), Some(&"ok".to_owned()));
assert_eq!(result.get("also_valid"), Some(&"good".to_owned()));
let result = parse_pairs("");
assert!(result.is_empty());
let result = parse_pairs("invalid,key=,=value");
assert!(result.is_empty());
}
}