use aether_core::events::{TRACEPARENT_KEY, TRACESTATE_KEY, TraceContext};
use opentelemetry::Context;
use opentelemetry::propagation::{Extractor, Injector, TextMapPropagator};
use opentelemetry::trace::TraceContextExt;
use opentelemetry_sdk::propagation::TraceContextPropagator;
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(untagged, rename_all = "camelCase", deny_unknown_fields)]
pub enum AgentTraceContext {
Parent(TraceContext),
Root {
#[serde(rename = "traceId")]
#[schemars(rename = "traceId")]
trace_id: String,
},
}
pub(crate) fn inject_trace_context(context: &Context) -> Option<TraceContext> {
if !context.span().span_context().is_valid() {
return None;
}
let mut carrier = W3cCarrier::default();
TraceContextPropagator::new().inject_context(context, &mut carrier);
Some(TraceContext { traceparent: carrier.traceparent?, tracestate: carrier.tracestate })
}
pub(crate) fn extract_trace_context(trace_context: &TraceContext) -> Option<Context> {
let context = TraceContextPropagator::new().extract(&W3cCarrier::from(trace_context));
context.span().span_context().is_valid().then_some(context)
}
#[derive(Default)]
struct W3cCarrier {
traceparent: Option<String>,
tracestate: Option<String>,
}
impl From<&TraceContext> for W3cCarrier {
fn from(trace_context: &TraceContext) -> Self {
Self { traceparent: Some(trace_context.traceparent.clone()), tracestate: trace_context.tracestate.clone() }
}
}
impl Injector for W3cCarrier {
fn set(&mut self, key: &str, value: String) {
match key {
TRACEPARENT_KEY => self.traceparent = Some(value),
TRACESTATE_KEY => self.tracestate = Some(value),
_ => {}
}
}
}
impl Extractor for W3cCarrier {
fn get(&self, key: &str) -> Option<&str> {
match key {
TRACEPARENT_KEY => self.traceparent.as_deref(),
TRACESTATE_KEY => self.tracestate.as_deref(),
_ => None,
}
}
fn keys(&self) -> Vec<&str> {
[(TRACEPARENT_KEY, &self.traceparent), (TRACESTATE_KEY, &self.tracestate)]
.into_iter()
.filter_map(|(key, value)| value.is_some().then_some(key))
.collect()
}
}