use opentelemetry::Context as OtelContext;
use opentelemetry::propagation::{Extractor, Injector, TextMapPropagator};
use opentelemetry::trace::{SpanContext, TraceContextExt, TraceFlags, TraceState};
use opentelemetry_sdk::propagation::TraceContextPropagator;
use opentelemetry_sdk::trace::{IdGenerator, RandomIdGenerator};
use tracing::Instrument;
use crate::Headers;
use crate::runtime::{
BlanketLayer, Context, Handler, Layer, Outgoing, PublishContext, PublishTransform, Settle,
};
const TRACEPARENT: &str = "traceparent";
const TRACESTATE: &str = "tracestate";
struct HeaderExtractor<'a>(&'a Headers);
impl Extractor for HeaderExtractor<'_> {
fn get(&self, key: &str) -> Option<&str> {
self.0.get_str(key)
}
fn keys(&self) -> Vec<&str> {
self.0.iter().map(|(name, _)| name).collect()
}
}
struct HeaderInjector<'a>(&'a mut Headers);
impl Injector for HeaderInjector<'_> {
fn set(&mut self, key: &str, value: String) {
if !value.is_empty() {
self.0.insert(key, value);
}
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct OpenTelemetry;
impl OpenTelemetry {
#[must_use]
pub const fn new() -> Self {
Self
}
#[must_use]
pub const fn consume_layer(&self) -> OpenTelemetryLayer {
OpenTelemetryLayer
}
#[must_use]
pub const fn propagation(&self) -> TracePropagation {
TracePropagation
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct OpenTelemetryLayer;
impl<H> Layer<H> for OpenTelemetryLayer {
type Handler = OpenTelemetryHandler<H>;
fn layer(&self, inner: H) -> Self::Handler {
OpenTelemetryHandler { inner }
}
}
impl BlanketLayer for OpenTelemetryLayer {
fn apply<M, C, S, H>(&self, handler: H) -> impl Handler<M, C, S> + 'static
where
M: Send + Sync + 'static,
C: Send + 'static,
S: Send + Sync + 'static,
H: Handler<M, C, S> + 'static,
{
self.layer(handler)
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct OpenTelemetryHandler<H> {
inner: H,
}
impl<M, C, S, H> Handler<M, C, S> for OpenTelemetryHandler<H>
where
M: Sync,
C: Send,
S: Send + Sync,
H: Handler<M, C, S>,
{
fn handle(&self, msg: &M, ctx: &mut Context<'_, C, S>) -> impl Future<Output = Settle> + Send {
let propagator = TraceContextPropagator::new();
let extracted =
propagator.extract_with_context(&OtelContext::new(), &HeaderExtractor(ctx.headers()));
let parent = extracted.span().span_context().clone();
let ids = RandomIdGenerator::default();
let consumer = if parent.is_valid() {
SpanContext::new(
parent.trace_id(),
ids.new_span_id(),
parent.trace_flags(),
false,
parent.trace_state().clone(),
)
} else {
SpanContext::new(
ids.new_trace_id(),
ids.new_span_id(),
TraceFlags::SAMPLED,
false,
TraceState::default(),
)
};
let span = tracing::info_span!(
"ruststream.consume",
otel.kind = "consumer",
subscription = %ctx.name(),
trace_id = %consumer.trace_id(),
span_id = %consumer.span_id(),
);
propagator.inject_context(
&OtelContext::new().with_remote_span_context(consumer),
&mut HeaderInjector(ctx.headers_mut()),
);
self.inner.handle(msg, ctx).instrument(span)
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct TracePropagation;
impl<C> PublishTransform<C> for TracePropagation {
fn apply(&self, out: &mut Outgoing<'_>, cx: &PublishContext<'_, C>) {
if let Some(traceparent) = cx.headers().get_str(TRACEPARENT) {
out.headers_mut()
.insert(TRACEPARENT, traceparent.as_bytes().to_vec());
if let Some(tracestate) = cx.headers().get_str(TRACESTATE) {
out.headers_mut()
.insert(TRACESTATE, tracestate.as_bytes().to_vec());
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
const HEADER: &str = "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01";
#[test]
fn context_round_trips_through_headers() {
let mut headers = Headers::new();
headers.insert(TRACEPARENT, HEADER);
let propagator = TraceContextPropagator::new();
let cx = propagator.extract_with_context(&OtelContext::new(), &HeaderExtractor(&headers));
assert!(cx.span().span_context().is_valid());
let mut out = Headers::new();
propagator.inject_context(&cx, &mut HeaderInjector(&mut out));
assert_eq!(out.get_str(TRACEPARENT), Some(HEADER));
assert!(
!out.contains(TRACESTATE),
"an empty tracestate must not be written"
);
}
#[test]
fn tracestate_rides_extraction_and_injection() {
let mut headers = Headers::new();
headers.insert(TRACEPARENT, HEADER);
headers.insert(TRACESTATE, "vendor=opaque");
let propagator = TraceContextPropagator::new();
let cx = propagator.extract_with_context(&OtelContext::new(), &HeaderExtractor(&headers));
let mut out = Headers::new();
propagator.inject_context(&cx, &mut HeaderInjector(&mut out));
assert_eq!(out.get_str(TRACESTATE), Some("vendor=opaque"));
}
#[test]
fn extractor_lists_the_header_names() {
let mut headers = Headers::new();
headers.insert(TRACEPARENT, HEADER);
assert_eq!(HeaderExtractor(&headers).keys(), vec![TRACEPARENT]);
}
}