use opentelemetry::Context;
use opentelemetry::propagation::{Extractor, Injector, TextMapCompositePropagator};
use opentelemetry_sdk::propagation::{BaggagePropagator, TraceContextPropagator};
pub fn install_w3c_propagators() {
let composite = TextMapCompositePropagator::new(vec![
Box::new(TraceContextPropagator::new()),
Box::new(BaggagePropagator::new()),
]);
opentelemetry::global::set_text_map_propagator(composite);
}
pub fn install_w3c_trace_context_propagator() {
install_w3c_propagators();
}
struct VecHeadersInjector<'a>(&'a mut Vec<(String, String)>);
impl Injector for VecHeadersInjector<'_> {
fn set(&mut self, key: &str, value: String) {
self.0.retain(|(k, _)| !k.eq_ignore_ascii_case(key));
self.0.push((key.to_string(), value));
}
}
pub struct VecHeadersExtractor<'a> {
headers: &'a [(String, String)],
}
impl<'a> VecHeadersExtractor<'a> {
pub fn new(headers: &'a [(String, String)]) -> Self {
Self { headers }
}
}
impl Extractor for VecHeadersExtractor<'_> {
fn get(&self, key: &str) -> Option<&str> {
self.headers.iter().find_map(|(k, v)| {
if k.eq_ignore_ascii_case(key) {
Some(v.as_str())
} else {
None
}
})
}
fn keys(&self) -> Vec<&str> {
self.headers.iter().map(|(k, _)| k.as_str()).collect()
}
}
pub fn inject_trace_context_into_headers(cx: &Context, headers: &mut Vec<(String, String)>) {
let mut inj = VecHeadersInjector(headers);
opentelemetry::global::get_text_map_propagator(|prop| prop.inject_context(cx, &mut inj));
}
#[inline]
pub fn inject_current_trace_context(headers: &mut Vec<(String, String)>) {
inject_trace_context_into_headers(&Context::current(), headers);
}
pub fn extract_trace_context_from_headers(base: &Context, headers: &[(String, String)]) -> Context {
let ext = VecHeadersExtractor::new(headers);
opentelemetry::global::get_text_map_propagator(|prop| prop.extract_with_context(base, &ext))
}
#[cfg(feature = "platform")]
pub fn inject_into_http_request(req: &mut id_effect_platform::http::HttpRequest) {
inject_current_trace_context(&mut req.headers);
}
#[cfg(feature = "platform")]
pub fn extract_from_http_request(req: &id_effect_platform::http::HttpRequest) -> Context {
extract_trace_context_from_headers(&Context::new(), &req.headers)
}
#[cfg(test)]
mod tests {
use super::*;
use opentelemetry::trace::TracerProvider as _;
use opentelemetry::trace::{TraceContextExt, Tracer};
use opentelemetry_sdk::trace::{InMemorySpanExporter, SdkTracerProvider};
#[test]
fn inject_then_extract_preserves_remote_span_id() {
install_w3c_propagators();
let exporter = InMemorySpanExporter::default();
let provider = SdkTracerProvider::builder()
.with_simple_exporter(exporter.clone())
.build();
let tracer = provider.tracer("propagation_test");
let span = tracer.start("remote");
let cx = opentelemetry::Context::current_with_span(span);
let mut headers = Vec::new();
inject_trace_context_into_headers(&cx, &mut headers);
assert!(
headers
.iter()
.any(|(k, _)| k.eq_ignore_ascii_case("traceparent")),
"expected traceparent header, got {headers:?}"
);
let extracted = extract_trace_context_from_headers(&Context::default(), &headers);
assert!(extracted.span().span_context().is_valid());
let _ = provider.shutdown();
}
#[test]
fn get_is_case_insensitive_for_header_name() {
let headers = vec![(
"TraceParent".to_string(),
"00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01".to_string(),
)];
let ext = VecHeadersExtractor::new(&headers);
assert!(ext.get("traceparent").is_some());
}
}