From 7862a2259db8dc1af440604c6c582487a59a2709 Mon Sep 17 00:00:00 2001 From: David Pedersen Date: Wed, 12 May 2021 18:24:09 +0200 Subject: [PATCH] feat(tonic): pass `trace_fn` the request rather than just the headers (#634) --- tonic/src/transport/server/mod.rs | 32 +++++++++++++++++++------------ 1 file changed, 20 insertions(+), 12 deletions(-) diff --git a/tonic/src/transport/server/mod.rs b/tonic/src/transport/server/mod.rs index c2ad4d5..e2007ca 100644 --- a/tonic/src/transport/server/mod.rs +++ b/tonic/src/transport/server/mod.rs @@ -33,7 +33,7 @@ use futures_util::{ future::{self, Either as FutureEither, MapErr}, TryFutureExt, }; -use http::{HeaderMap, Request, Response}; +use http::{Request, Response}; use hyper::{server::accept, Body}; use std::{ fmt, @@ -48,7 +48,7 @@ use tower::{limit::concurrency::ConcurrencyLimitLayer, util::Either, Service, Se use tracing_futures::{Instrument, Instrumented}; type BoxService = tower::util::BoxService, Response, crate::Error>; -type TraceInterceptor = Arc tracing::Span + Send + Sync + 'static>; +type TraceInterceptor = Arc) -> tracing::Span + Send + Sync + 'static>; const DEFAULT_HTTP2_KEEPALIVE_TIMEOUT_SECS: u64 = 20; @@ -290,10 +290,10 @@ impl Server { } } - /// Intercept inbound headers and add a [`tracing::Span`] to each response future. + /// Intercept inbound requests and add a [`tracing::Span`] to each response future. pub fn trace_fn(self, f: F) -> Self where - F: Fn(&HeaderMap) -> tracing::Span + Send + Sync + 'static, + F: Fn(&http::Request<()>) -> tracing::Span + Send + Sync + 'static, { Server { trace_interceptor: Some(Arc::new(f)), @@ -361,7 +361,7 @@ impl Server { IE: Into, F: Future, { - let span = self.trace_interceptor.clone(); + let trace_interceptor = self.trace_interceptor.clone(); let concurrency_limit = self.concurrency_limit; let init_connection_window_size = self.init_connection_window_size; let init_stream_window_size = self.init_stream_window_size; @@ -381,7 +381,7 @@ impl Server { inner: svc, concurrency_limit, timeout, - span, + trace_interceptor, }; let server = hyper::Server::builder(incoming) @@ -582,7 +582,7 @@ impl fmt::Debug for Server { struct Svc { inner: S, - span: Option, + trace_interceptor: Option, conn_info: ConnectionInfo, } @@ -602,8 +602,16 @@ where } fn call(&mut self, mut req: Request) -> Self::Future { - let span = if let Some(trace_interceptor) = &self.span { - trace_interceptor(req.headers()) + let span = if let Some(trace_interceptor) = &self.trace_interceptor { + let (parts, body) = req.into_parts(); + let bodyless_request = Request::from_parts(parts, ()); + + let span = trace_interceptor(&bodyless_request); + + let (parts, _) = bodyless_request.into_parts(); + req = Request::from_parts(parts, body); + + span } else { tracing::Span::none() }; @@ -624,7 +632,7 @@ struct MakeSvc { concurrency_limit: Option, timeout: Option, inner: S, - span: Option, + trace_interceptor: Option, } impl Service<&ServerIo> for MakeSvc @@ -650,7 +658,7 @@ where let svc = self.inner.clone(); let concurrency_limit = self.concurrency_limit; let timeout = self.timeout; - let span = self.span.clone(); + let trace_interceptor = self.trace_interceptor.clone(); Box::pin(async move { let svc = ServiceBuilder::new() @@ -661,7 +669,7 @@ where let svc = BoxService::new(Svc { inner: svc, - span, + trace_interceptor, conn_info, });