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#[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 pub fn client(parent: Option<TraceContext>) -> Self {
31 Self {
32 mode: Mode::Client(parent),
33 }
34 }
35
36 pub fn server() -> Self {
38 Self { mode: Mode::Server }
39 }
40
41 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#[cfg(feature = "telemetry")]
81#[derive(Debug, Clone, Copy, PartialEq, Eq)]
82pub enum RpcTelemetryMode {
83 Client,
84 Server,
85}
86
87#[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}