use opentelemetry::{
context::FutureExt,
global,
propagation::{Extractor, Injector},
trace::{SpanKind, Status, TraceContextExt, Tracer},
};
use std::{
future::Future,
pin::Pin,
task::{self, Poll},
};
use tonic::body::Body;
use tower::{Layer, Service};
#[derive(Clone, Debug, Default)]
pub struct GrpcClientTraceLayer;
impl<S> Layer<S> for GrpcClientTraceLayer {
type Service = GrpcClientTraceService<S>;
fn layer(&self, inner: S) -> Self::Service {
GrpcClientTraceService { inner }
}
}
#[derive(Clone, Debug)]
pub struct GrpcClientTraceService<S> {
inner: S,
}
impl<S> Service<http::Request<Body>> for GrpcClientTraceService<S>
where
S: Service<http::Request<Body>> + Clone + Send + 'static,
S::Future: Send + 'static,
{
type Response = S::Response;
type Error = S::Error;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, cx: &mut task::Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, mut req: http::Request<Body>) -> Self::Future {
let path = req.uri().path().to_string();
let tracer = global::tracer("otel-bootstrap");
let span = tracer
.span_builder(path)
.with_kind(SpanKind::Client)
.start(&tracer);
let cx = opentelemetry::Context::current_with_span(span);
global::get_text_map_propagator(|propagator| {
propagator.inject_context(&cx, &mut MetadataInjector(req.headers_mut()));
});
let clone = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, clone);
Box::pin(async move { inner.call(req).await })
}
}
#[derive(Clone, Debug, Default)]
pub struct GrpcServerTraceLayer;
impl<S> Layer<S> for GrpcServerTraceLayer {
type Service = GrpcServerTraceService<S>;
fn layer(&self, inner: S) -> Self::Service {
GrpcServerTraceService { inner }
}
}
#[derive(Clone, Debug)]
pub struct GrpcServerTraceService<S> {
inner: S,
}
impl<S> Service<http::Request<Body>> for GrpcServerTraceService<S>
where
S: Service<http::Request<Body>, Response = http::Response<Body>> + Clone + Send + 'static,
S::Future: Send + 'static,
S::Error: Into<Box<dyn std::error::Error + Send + Sync>>,
{
type Response = http::Response<Body>;
type Error = S::Error;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, cx: &mut task::Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, req: http::Request<Body>) -> Self::Future {
let path = req.uri().path().to_string();
let parent_cx = global::get_text_map_propagator(|propagator| {
propagator.extract(&MetadataExtractor(req.headers()))
});
let tracer = global::tracer("otel-bootstrap");
let span = tracer
.span_builder(path)
.with_kind(SpanKind::Server)
.start_with_context(&tracer, &parent_cx);
let cx = parent_cx.with_span(span);
let clone = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, clone);
Box::pin(async move {
let result = inner.call(req).with_context(cx.clone()).await;
match &result {
Ok(resp) => {
if resp.status().is_server_error() {
cx.span().set_status(Status::Error {
description: resp.status().canonical_reason().unwrap_or("").into(),
});
}
}
Err(_) => {
cx.span().set_status(Status::Error {
description: "transport error".into(),
});
}
}
result
})
}
}
struct MetadataExtractor<'a>(&'a http::HeaderMap);
impl Extractor for MetadataExtractor<'_> {
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(http::HeaderName::as_str).collect()
}
}
struct MetadataInjector<'a>(&'a mut http::HeaderMap);
impl Injector for MetadataInjector<'_> {
fn set(&mut self, key: &str, value: String) {
if let (Ok(name), Ok(val)) = (
http::HeaderName::from_bytes(key.as_bytes()),
http::HeaderValue::from_str(&value),
) {
self.0.insert(name, val);
}
}
}