otel_bootstrap/
axum_middleware.rs1use axum::{
13 body::Body,
14 http::{HeaderMap, HeaderName, HeaderValue, Request, Response},
15};
16use opentelemetry::{
17 context::FutureExt,
18 global,
19 propagation::{Extractor, Injector},
20 trace::{SpanKind, Status, TraceContextExt, Tracer},
21};
22use opentelemetry_semantic_conventions::attribute::{
23 HTTP_REQUEST_METHOD, HTTP_RESPONSE_STATUS_CODE,
24};
25use std::{
26 future::Future,
27 pin::Pin,
28 task::{self, Poll},
29};
30use tower::{Layer, Service};
31
32#[derive(Clone, Debug)]
44pub struct OtelTraceLayer;
45
46impl<S> Layer<S> for OtelTraceLayer {
47 type Service = OtelTraceService<S>;
48
49 fn layer(&self, inner: S) -> Self::Service {
50 OtelTraceService { inner }
51 }
52}
53
54#[derive(Clone, Debug)]
59pub struct OtelTraceService<S> {
60 inner: S,
61}
62
63impl<S> Service<Request<Body>> for OtelTraceService<S>
64where
65 S: Service<Request<Body>, Response = Response<Body>> + Send + Clone + 'static,
66 S::Future: Send + 'static,
67 S::Error: Into<Box<dyn std::error::Error + Send + Sync>>,
68{
69 type Response = Response<Body>;
70 type Error = S::Error;
71 type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
72
73 fn poll_ready(&mut self, cx: &mut task::Context<'_>) -> Poll<Result<(), Self::Error>> {
74 self.inner.poll_ready(cx)
75 }
76
77 fn call(&mut self, req: Request<Body>) -> Self::Future {
78 let method = req.method().to_string();
79 let route = req.uri().path().to_string();
80
81 let parent_cx = global::get_text_map_propagator(|propagator| {
83 propagator.extract(&HeaderExtractor(req.headers()))
84 });
85
86 let tracer = global::tracer("otel-bootstrap");
88 let span = tracer
89 .span_builder(format!("{method} {route}"))
90 .with_kind(SpanKind::Server)
91 .with_attributes([opentelemetry::KeyValue::new(HTTP_REQUEST_METHOD, method)])
92 .start_with_context(&tracer, &parent_cx);
93
94 let cx = parent_cx.with_span(span);
95
96 let clone = self.inner.clone();
102 let mut inner = std::mem::replace(&mut self.inner, clone);
103
104 Box::pin(async move {
105 let mut response = inner.call(req).with_context(cx.clone()).await?;
114
115 let status_code = response.status().as_u16();
117 cx.span().set_attribute(opentelemetry::KeyValue::new(
118 HTTP_RESPONSE_STATUS_CODE,
119 status_code as i64,
120 ));
121 if response.status().is_server_error() {
122 cx.span().set_status(Status::Error {
123 description: response.status().canonical_reason().unwrap_or("").into(),
124 });
125 }
126
127 let mut injector = HeaderInjector(response.headers_mut());
129 global::get_text_map_propagator(|propagator| {
130 propagator.inject_context(&cx, &mut injector);
131 });
132
133 Ok(response)
134 })
135 }
136}
137
138#[derive(Debug)]
166pub struct SpanEnricherLayer<T>(std::marker::PhantomData<T>);
167
168impl<T> Default for SpanEnricherLayer<T> {
169 fn default() -> Self {
170 Self(std::marker::PhantomData)
171 }
172}
173
174impl<T> Clone for SpanEnricherLayer<T> {
175 fn clone(&self) -> Self {
176 Self(std::marker::PhantomData)
177 }
178}
179
180impl<T, S> Layer<S> for SpanEnricherLayer<T>
181where
182 T: crate::span_enrichment::EnrichSpan + Clone + Send + Sync + 'static,
183{
184 type Service = SpanEnricherService<T, S>;
185
186 fn layer(&self, inner: S) -> Self::Service {
187 SpanEnricherService {
188 inner,
189 _marker: std::marker::PhantomData,
190 }
191 }
192}
193
194#[derive(Clone, Debug)]
196pub struct SpanEnricherService<T, S> {
197 inner: S,
198 _marker: std::marker::PhantomData<T>,
199}
200
201impl<T, S> Service<Request<Body>> for SpanEnricherService<T, S>
202where
203 T: crate::span_enrichment::EnrichSpan + Clone + Send + Sync + 'static,
204 S: Service<Request<Body>, Response = Response<Body>>,
205{
206 type Response = Response<Body>;
207 type Error = S::Error;
208 type Future = S::Future;
209
210 fn poll_ready(&mut self, cx: &mut task::Context<'_>) -> Poll<Result<(), Self::Error>> {
211 self.inner.poll_ready(cx)
212 }
213
214 fn call(&mut self, req: Request<Body>) -> Self::Future {
215 if let Some(ctx) = req.extensions().get::<T>() {
216 ctx.enrich_span(&tracing::Span::current());
217 }
218 self.inner.call(req)
219 }
220}
221
222struct HeaderExtractor<'a>(&'a HeaderMap);
224
225impl Extractor for HeaderExtractor<'_> {
226 fn get(&self, key: &str) -> Option<&str> {
227 self.0.get(key).and_then(|v| v.to_str().ok())
228 }
229
230 fn keys(&self) -> Vec<&str> {
231 self.0.keys().map(HeaderName::as_str).collect()
232 }
233}
234
235struct HeaderInjector<'a>(&'a mut HeaderMap);
237
238impl Injector for HeaderInjector<'_> {
239 fn set(&mut self, key: &str, value: String) {
240 if let (Ok(name), Ok(val)) = (
241 HeaderName::from_bytes(key.as_bytes()),
242 HeaderValue::from_str(&value),
243 ) {
244 self.0.insert(name, val);
245 }
246 }
247}