use std::task::{Context as TaskContext, Poll};
use futures::{future::BoxFuture, FutureExt};
use opentelemetry::{
global,
propagation::{Extractor, Injector},
Context,
};
use tonic::metadata::{MetadataKey, MetadataMap, MetadataValue};
use tower::{Layer, Service};
use tracing::warn;
pub const TRAFFIC_TYPE_KEY: &str = "traffic_type";
pub const TRAFFIC_TYPE_ORGANIC: &str = "organic";
pub const TRAFFIC_TYPE_SYNTHETIC: &str = "synthetic";
pub const TRAFFIC_TYPE_UNKNOWN: &str = "unknown";
pub const TRAFFIC_TYPE_ENV_VAR: &str = "LINERA_TRAFFIC_TYPE";
#[derive(Clone, Copy, Debug, Default)]
pub struct OtelContextLayer;
#[derive(Clone, Debug)]
pub struct OtelContextService<S> {
inner: S,
}
#[derive(Clone, Debug)]
pub struct ExtractedOtelContext(pub Context);
pub trait HasOtelContext {
fn get_otel_context(&self) -> Option<&ExtractedOtelContext>;
}
impl<B> HasOtelContext for http::Request<B> {
fn get_otel_context(&self) -> Option<&ExtractedOtelContext> {
self.extensions().get::<ExtractedOtelContext>()
}
}
impl<T> HasOtelContext for tonic::Request<T> {
fn get_otel_context(&self) -> Option<&ExtractedOtelContext> {
self.extensions().get::<ExtractedOtelContext>()
}
}
pub fn inject_context(cx: &Context, metadata: &mut MetadataMap) {
global::get_text_map_propagator(|propagator| {
propagator.inject_context(cx, &mut MetadataInjector(metadata));
});
}
pub fn extract_context(metadata: &MetadataMap) -> Context {
global::get_text_map_propagator(|propagator| propagator.extract(&MetadataExtractor(metadata)))
}
pub fn get_context_with_traffic_type() -> Context {
use opentelemetry::{baggage::BaggageExt, Key, KeyValue};
let cx = Context::current();
if std::env::var(TRAFFIC_TYPE_ENV_VAR)
.map(|v| v == TRAFFIC_TYPE_SYNTHETIC)
.unwrap_or(false)
{
cx.with_baggage(vec![KeyValue::new(
Key::new(TRAFFIC_TYPE_KEY),
TRAFFIC_TYPE_SYNTHETIC,
)])
} else {
cx
}
}
pub fn get_traffic_type(cx: &Context) -> &'static str {
use opentelemetry::baggage::BaggageExt;
cx.baggage()
.get(TRAFFIC_TYPE_KEY)
.map(|v| v.as_str())
.and_then(|v| {
if v == TRAFFIC_TYPE_SYNTHETIC {
Some(TRAFFIC_TYPE_SYNTHETIC)
} else {
None
}
})
.unwrap_or(TRAFFIC_TYPE_ORGANIC)
}
pub fn get_traffic_type_from_request<R: HasOtelContext>(request: &R) -> &'static str {
request
.get_otel_context()
.map_or(TRAFFIC_TYPE_ORGANIC, |ext| get_traffic_type(&ext.0))
}
pub fn get_otel_context_from_tonic_request<T>(request: &tonic::Request<T>) -> Option<Context> {
request
.extensions()
.get::<ExtractedOtelContext>()
.map(|ext| ext.0.clone())
}
pub fn create_request_with_context<T>(inner: T, cx: Option<&Context>) -> tonic::Request<T> {
let mut request = tonic::Request::new(inner);
if let Some(cx) = cx {
inject_context(cx, request.metadata_mut());
}
request
}
pub fn create_request_with_current_span_context<T>(inner: T) -> tonic::Request<T> {
use tracing_opentelemetry::OpenTelemetrySpanExt;
let cx = tracing::Span::current().context();
create_request_with_context(inner, Some(&cx))
}
impl<S> Layer<S> for OtelContextLayer {
type Service = OtelContextService<S>;
fn layer(&self, service: S) -> Self::Service {
OtelContextService { inner: service }
}
}
impl<S, B> Service<http::Request<B>> for OtelContextService<S>
where
S: Service<http::Request<B>> + Clone + Send + 'static,
S::Future: Send,
B: Send + 'static,
{
type Response = S::Response;
type Error = S::Error;
type Future = BoxFuture<'static, Result<Self::Response, Self::Error>>;
fn poll_ready(&mut self, cx: &mut TaskContext<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, mut request: http::Request<B>) -> Self::Future {
use tracing::Instrument;
use tracing_opentelemetry::OpenTelemetrySpanExt;
let cx = global::get_text_map_propagator(|propagator| {
propagator.extract(&HttpHeaderExtractor(request.headers()))
});
request
.extensions_mut()
.insert(ExtractedOtelContext(cx.clone()));
let span = tracing::info_span!("grpc_request");
span.set_parent(cx);
let mut inner = self.inner.clone();
async move { inner.call(request).await }
.instrument(span)
.boxed()
}
}
struct MetadataInjector<'a>(&'a mut MetadataMap);
impl Injector for MetadataInjector<'_> {
fn set(&mut self, key: &str, value: String) {
match MetadataKey::from_bytes(key.as_bytes()) {
Ok(key) => match MetadataValue::try_from(&value) {
Ok(value) => {
self.0.insert(key, value);
}
Err(error) => {
warn!(
value,
error = format!("{error:#}"),
"failed to parse metadata value"
);
}
},
Err(error) => {
warn!(
key,
error = format!("{error:#}"),
"failed to parse metadata key"
);
}
}
}
}
struct MetadataExtractor<'a>(&'a MetadataMap);
impl Extractor for MetadataExtractor<'_> {
fn get(&self, key: &str) -> Option<&str> {
self.0.get(key).and_then(|value| value.to_str().ok())
}
fn keys(&self) -> Vec<&str> {
self.0
.keys()
.filter_map(|key| match key {
tonic::metadata::KeyRef::Ascii(key) => Some(key.as_str()),
tonic::metadata::KeyRef::Binary(_) => None,
})
.collect()
}
}
struct HttpHeaderExtractor<'a>(&'a http::HeaderMap);
impl Extractor for HttpHeaderExtractor<'_> {
fn get(&self, key: &str) -> Option<&str> {
self.0.get(key).and_then(|v| v.to_str().ok())
}
fn keys(&self) -> Vec<&str> {
self.0.keys().map(|k| k.as_str()).collect()
}
}
#[cfg(test)]
mod tests {
use opentelemetry::{baggage::BaggageExt, Key, KeyValue};
use super::*;
#[test]
fn test_inject_and_extract_baggage() {
use opentelemetry::propagation::TextMapCompositePropagator;
use opentelemetry_sdk::propagation::{BaggagePropagator, TraceContextPropagator};
let propagator = TextMapCompositePropagator::new(vec![
Box::new(TraceContextPropagator::new()),
Box::new(BaggagePropagator::new()),
]);
global::set_text_map_propagator(propagator);
let cx = Context::current().with_baggage(vec![KeyValue::new(
Key::new(TRAFFIC_TYPE_KEY),
TRAFFIC_TYPE_SYNTHETIC,
)]);
let mut metadata = MetadataMap::new();
inject_context(&cx, &mut metadata);
assert!(
metadata.get("baggage").is_some(),
"baggage header should be present"
);
let extracted_cx = extract_context(&metadata);
let traffic_type = get_traffic_type(&extracted_cx);
assert_eq!(traffic_type, TRAFFIC_TYPE_SYNTHETIC);
}
#[test]
fn test_default_traffic_type_is_organic() {
let cx = Context::current();
assert_eq!(get_traffic_type(&cx), TRAFFIC_TYPE_ORGANIC);
}
}