Skip to main content

rpc/
trace.rs

1#[cfg(feature = "telemetry")]
2use rust_zero_core::{TelemetrySpan, TelemetrySpanKind};
3use rust_zero_core::{TraceContext, TraceFlags};
4use tonic::{service::Interceptor, Request, Status};
5
6#[cfg(feature = "telemetry")]
7use std::{
8    fmt,
9    future::Future,
10    pin::Pin,
11    task::{Context, Poll},
12};
13#[cfg(feature = "telemetry")]
14use tower::{Layer, Service};
15
16/// A W3C trace-context interceptor for Tonic clients and servers.
17#[derive(Debug, Clone)]
18pub struct RpcTrace {
19    mode: Mode,
20}
21
22#[derive(Debug, Clone)]
23enum Mode {
24    Client(Option<TraceContext>),
25    Server,
26}
27
28impl RpcTrace {
29    /// Creates a client interceptor. A configured parent is used to create a child span per call.
30    pub fn client(parent: Option<TraceContext>) -> Self {
31        Self {
32            mode: Mode::Client(parent),
33        }
34    }
35
36    /// Creates a server interceptor that accepts `traceparent` metadata and creates a server span.
37    pub fn server() -> Self {
38        Self { mode: Mode::Server }
39    }
40
41    /// Retrieves the server span installed in a request's extensions.
42    pub fn context<T>(request: &Request<T>) -> Option<TraceContext> {
43        request.extensions().get::<TraceContext>().cloned()
44    }
45}
46
47impl Interceptor for RpcTrace {
48    fn call(&mut self, mut request: Request<()>) -> Result<Request<()>, Status> {
49        match &self.mode {
50            Mode::Client(parent) => {
51                let context = parent
52                    .as_ref()
53                    .map(TraceContext::child)
54                    .unwrap_or_else(|| TraceContext::root(TraceFlags::SAMPLED));
55                request.metadata_mut().insert(
56                    "traceparent",
57                    context
58                        .traceparent()
59                        .parse()
60                        .expect("generated traceparent values are valid ASCII metadata"),
61                );
62                request.extensions_mut().insert(context);
63            }
64            Mode::Server => {
65                let context = request
66                    .metadata()
67                    .get("traceparent")
68                    .and_then(|value| value.to_str().ok())
69                    .and_then(|value| TraceContext::parse(value).ok())
70                    .map(|parent| parent.child())
71                    .unwrap_or_else(|| TraceContext::root(TraceFlags::SAMPLED));
72                request.extensions_mut().insert(context);
73            }
74        }
75        Ok(request)
76    }
77}
78
79/// Whether a gRPC telemetry layer instruments outbound or inbound requests.
80#[cfg(feature = "telemetry")]
81#[derive(Debug, Clone, Copy, PartialEq, Eq)]
82pub enum RpcTelemetryMode {
83    Client,
84    Server,
85}
86
87/// A Tower layer that creates complete OpenTelemetry spans around gRPC calls.
88///
89/// Apply [`RpcTelemetryLayer::server`] with `tonic::transport::Server::layer`, or wrap a client
90/// channel with [`RpcTelemetryLayer::client`] before constructing a generated Tonic client.
91#[cfg(feature = "telemetry")]
92#[derive(Debug, Clone, Copy)]
93pub struct RpcTelemetryLayer {
94    mode: RpcTelemetryMode,
95}
96
97#[cfg(feature = "telemetry")]
98impl RpcTelemetryLayer {
99    pub fn client() -> Self {
100        Self {
101            mode: RpcTelemetryMode::Client,
102        }
103    }
104
105    pub fn server() -> Self {
106        Self {
107            mode: RpcTelemetryMode::Server,
108        }
109    }
110
111    pub fn wrap<S>(&self, inner: S) -> RpcTelemetryService<S> {
112        self.layer(inner)
113    }
114}
115
116#[cfg(feature = "telemetry")]
117impl<S> Layer<S> for RpcTelemetryLayer {
118    type Service = RpcTelemetryService<S>;
119
120    fn layer(&self, inner: S) -> Self::Service {
121        RpcTelemetryService {
122            inner,
123            mode: self.mode,
124        }
125    }
126}
127
128#[cfg(feature = "telemetry")]
129#[derive(Debug, Clone)]
130pub struct RpcTelemetryService<S> {
131    inner: S,
132    mode: RpcTelemetryMode,
133}
134
135#[cfg(feature = "telemetry")]
136impl<S, RequestBody, ResponseBody> Service<http::Request<RequestBody>> for RpcTelemetryService<S>
137where
138    S: Service<http::Request<RequestBody>, Response = http::Response<ResponseBody>>,
139    S::Future: Send + 'static,
140    S::Error: fmt::Display + Send + 'static,
141    ResponseBody: Send + 'static,
142{
143    type Response = http::Response<ResponseBody>;
144    type Error = S::Error;
145    type Future =
146        Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send + 'static>>;
147
148    fn poll_ready(&mut self, context: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
149        self.inner.poll_ready(context)
150    }
151
152    fn call(&mut self, mut request: http::Request<RequestBody>) -> Self::Future {
153        let path = request.uri().path().to_owned();
154        let parent = request
155            .extensions()
156            .get::<TraceContext>()
157            .cloned()
158            .or_else(|| {
159                request
160                    .headers()
161                    .get("traceparent")
162                    .and_then(|value| value.to_str().ok())
163                    .and_then(|value| TraceContext::parse(value).ok())
164            });
165        let span = TelemetrySpan::start(
166            path.clone(),
167            match self.mode {
168                RpcTelemetryMode::Client => TelemetrySpanKind::Client,
169                RpcTelemetryMode::Server => TelemetrySpanKind::Server,
170            },
171            parent.as_ref(),
172            [("rpc.system", "grpc".to_owned()), ("rpc.method", path)],
173        );
174
175        if let Some(context) = span.trace_context().cloned() {
176            if self.mode == RpcTelemetryMode::Client {
177                if let Ok(value) = context.traceparent().parse() {
178                    request.headers_mut().insert("traceparent", value);
179                }
180            }
181            request.extensions_mut().insert(context);
182        }
183        let future = self.inner.call(request);
184
185        Box::pin(async move {
186            match future.await {
187                Ok(response) => {
188                    span.set_attribute(
189                        "http.response.status_code",
190                        response.status().as_u16().to_string(),
191                    );
192                    if let Some(status) = response
193                        .headers()
194                        .get("grpc-status")
195                        .and_then(|value| value.to_str().ok())
196                    {
197                        span.set_attribute("rpc.grpc.status_code", status.to_owned());
198                        if status != "0" {
199                            span.set_error(format!("gRPC status {status}"));
200                        }
201                    }
202                    span.end();
203                    Ok(response)
204                }
205                Err(error) => {
206                    span.set_error(error.to_string());
207                    span.end();
208                    Err(error)
209                }
210            }
211        })
212    }
213}
214
215#[cfg(test)]
216mod tests {
217    use super::*;
218
219    #[test]
220    fn propagates_a_client_trace_to_a_server_span() {
221        let parent =
222            TraceContext::parse("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01").unwrap();
223        let mut client = RpcTrace::client(Some(parent));
224        let outgoing = client.call(Request::new(())).unwrap();
225        let mut server = RpcTrace::server();
226        let incoming = server.call(outgoing).unwrap();
227        let context = RpcTrace::context(&incoming).unwrap();
228
229        assert_eq!(context.trace_id(), "4bf92f3577b34da6a3ce929d0e0e4736");
230        assert!(context.parent_span_id().is_some());
231    }
232
233    #[cfg(feature = "telemetry")]
234    #[tokio::test]
235    async fn telemetry_layer_injects_a_child_context() {
236        use rust_zero_core::Telemetry;
237        use std::{convert::Infallible, future::Ready};
238
239        #[derive(Clone)]
240        struct Capture;
241
242        impl Service<http::Request<()>> for Capture {
243            type Response = http::Response<String>;
244            type Error = Infallible;
245            type Future = Ready<Result<Self::Response, Self::Error>>;
246
247            fn poll_ready(&mut self, _context: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
248                Poll::Ready(Ok(()))
249            }
250
251            fn call(&mut self, request: http::Request<()>) -> Self::Future {
252                std::future::ready(Ok(http::Response::new(
253                    request
254                        .headers()
255                        .get("traceparent")
256                        .unwrap()
257                        .to_str()
258                        .unwrap()
259                        .to_owned(),
260                )))
261            }
262        }
263
264        let telemetry = Telemetry::local("rpc-client", 1.0).unwrap();
265        let parent =
266            TraceContext::parse("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01").unwrap();
267        let mut request = http::Request::builder()
268            .uri("/rust_zero.echo.Echo/Echo")
269            .body(())
270            .unwrap();
271        request.extensions_mut().insert(parent);
272        let mut service = RpcTelemetryLayer::client().layer(Capture);
273
274        let traceparent = service.call(request).await.unwrap().into_body();
275        assert!(traceparent.contains("4bf92f3577b34da6a3ce929d0e0e4736"));
276        telemetry.force_flush().unwrap();
277    }
278}