Skip to main content

rest/
trace.rs

1use actix_web::{
2    body::MessageBody,
3    dev::{Service, ServiceRequest, ServiceResponse, Transform},
4    http::header::{HeaderName, HeaderValue},
5    Error, HttpMessage, HttpRequest,
6};
7use futures::future::{ok, LocalBoxFuture, Ready};
8#[cfg(feature = "telemetry")]
9use rust_zero_core::{TelemetrySpan, TelemetrySpanKind};
10use rust_zero_core::{TraceContext, TraceFlags};
11use std::task::{Context, Poll};
12
13const TRACEPARENT: HeaderName = HeaderName::from_static("traceparent");
14
15/// Propagates W3C trace context and creates a server span for every request.
16#[derive(Debug, Clone, Copy, Default)]
17pub struct TraceContextMiddleware;
18
19impl TraceContextMiddleware {
20    pub fn new() -> Self {
21        Self
22    }
23
24    pub fn context(request: &HttpRequest) -> Option<TraceContext> {
25        request.extensions().get::<TraceContext>().cloned()
26    }
27}
28
29impl<S, B> Transform<S, ServiceRequest> for TraceContextMiddleware
30where
31    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
32    S::Future: 'static,
33    B: MessageBody + 'static,
34{
35    type Response = ServiceResponse<B>;
36    type Error = Error;
37    type Transform = TraceContextService<S>;
38    type InitError = ();
39    type Future = Ready<Result<Self::Transform, Self::InitError>>;
40
41    fn new_transform(&self, service: S) -> Self::Future {
42        ok(TraceContextService { service })
43    }
44}
45
46pub struct TraceContextService<S> {
47    service: S,
48}
49
50/// Creates OpenTelemetry server spans and propagates their W3C context.
51///
52/// Install a [`rust_zero_core::Telemetry`] provider before serving requests. This middleware
53/// replaces [`TraceContextMiddleware`] when the `telemetry` feature is enabled.
54#[cfg(feature = "telemetry")]
55#[derive(Debug, Clone, Copy, Default)]
56pub struct OpenTelemetryTracing;
57
58#[cfg(feature = "telemetry")]
59impl OpenTelemetryTracing {
60    pub fn new() -> Self {
61        Self
62    }
63
64    pub fn context(request: &HttpRequest) -> Option<TraceContext> {
65        request.extensions().get::<TraceContext>().cloned()
66    }
67}
68
69#[cfg(feature = "telemetry")]
70impl<S, B> Transform<S, ServiceRequest> for OpenTelemetryTracing
71where
72    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
73    S::Future: 'static,
74    B: MessageBody + 'static,
75{
76    type Response = ServiceResponse<B>;
77    type Error = Error;
78    type Transform = OpenTelemetryTracingService<S>;
79    type InitError = ();
80    type Future = Ready<Result<Self::Transform, Self::InitError>>;
81
82    fn new_transform(&self, service: S) -> Self::Future {
83        ok(OpenTelemetryTracingService { service })
84    }
85}
86
87#[cfg(feature = "telemetry")]
88pub struct OpenTelemetryTracingService<S> {
89    service: S,
90}
91
92#[cfg(feature = "telemetry")]
93impl<S, B> Service<ServiceRequest> for OpenTelemetryTracingService<S>
94where
95    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
96    S::Future: 'static,
97    B: MessageBody + 'static,
98{
99    type Response = ServiceResponse<B>;
100    type Error = Error;
101    type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;
102
103    fn poll_ready(&self, context: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
104        self.service.poll_ready(context)
105    }
106
107    fn call(&self, request: ServiceRequest) -> Self::Future {
108        let method = request.method().to_string();
109        let path = request.path().to_owned();
110        let parent = request
111            .headers()
112            .get(&TRACEPARENT)
113            .and_then(|value| value.to_str().ok())
114            .and_then(|value| TraceContext::parse(value).ok());
115        let span = TelemetrySpan::start(
116            format!("{method} {path}"),
117            TelemetrySpanKind::Server,
118            parent.as_ref(),
119            [
120                ("http.request.method", method),
121                ("url.path", path),
122                (
123                    "server.address",
124                    request.connection_info().host().to_owned(),
125                ),
126            ],
127        );
128        let context = span
129            .trace_context()
130            .cloned()
131            .unwrap_or_else(|| TraceContext::root(TraceFlags::SAMPLED));
132        let traceparent = context.traceparent();
133        request.extensions_mut().insert(context);
134        let future = self.service.call(request);
135
136        Box::pin(async move {
137            match future.await {
138                Ok(mut response) => {
139                    let status = response.status().as_u16();
140                    span.set_attribute("http.response.status_code", status.to_string());
141                    if status >= 500 {
142                        span.set_error(format!("HTTP {status}"));
143                    }
144                    response.headers_mut().insert(
145                        TRACEPARENT,
146                        HeaderValue::from_str(&traceparent)
147                            .expect("generated traceparent values are valid HTTP headers"),
148                    );
149                    span.end();
150                    Ok(response)
151                }
152                Err(error) => {
153                    span.set_error(error.to_string());
154                    span.end();
155                    Err(error)
156                }
157            }
158        })
159    }
160}
161
162impl<S, B> Service<ServiceRequest> for TraceContextService<S>
163where
164    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
165    S::Future: 'static,
166    B: MessageBody + 'static,
167{
168    type Response = ServiceResponse<B>;
169    type Error = Error;
170    type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;
171
172    fn poll_ready(&self, context: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
173        self.service.poll_ready(context)
174    }
175
176    fn call(&self, request: ServiceRequest) -> Self::Future {
177        let context = request
178            .headers()
179            .get(&TRACEPARENT)
180            .and_then(|value| value.to_str().ok())
181            .and_then(|value| TraceContext::parse(value).ok())
182            .map(|parent| parent.child())
183            .unwrap_or_else(|| TraceContext::root(TraceFlags::SAMPLED));
184        let traceparent = context.traceparent();
185        request.extensions_mut().insert(context);
186        let future = self.service.call(request);
187
188        Box::pin(async move {
189            let mut response = future.await?;
190            response.headers_mut().insert(
191                TRACEPARENT,
192                HeaderValue::from_str(&traceparent)
193                    .expect("generated traceparent values are valid HTTP headers"),
194            );
195            Ok(response)
196        })
197    }
198}
199
200#[cfg(test)]
201mod tests {
202    use super::*;
203    use actix_web::{test, web, App, HttpResponse};
204
205    #[actix_rt::test]
206    async fn preserves_trace_id_and_creates_a_server_span() {
207        let inbound = "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01";
208        let app = test::init_service(App::new().wrap(TraceContextMiddleware::new()).route(
209            "/",
210            web::get().to(|request: HttpRequest| async move {
211                HttpResponse::Ok().body(
212                    TraceContextMiddleware::context(&request)
213                        .unwrap()
214                        .parent_span_id()
215                        .unwrap(),
216                )
217            }),
218        ))
219        .await;
220        let response = test::call_service(
221            &app,
222            test::TestRequest::get()
223                .uri("/")
224                .insert_header(("traceparent", inbound))
225                .to_request(),
226        )
227        .await;
228
229        assert_eq!(response.status(), actix_web::http::StatusCode::OK);
230        assert!(response
231            .headers()
232            .get("traceparent")
233            .unwrap()
234            .to_str()
235            .unwrap()
236            .contains("4bf92f3577b34da6a3ce929d0e0e4736"));
237        assert_eq!(test::read_body(response).await, "00f067aa0ba902b7");
238    }
239
240    #[cfg(feature = "telemetry")]
241    #[actix_rt::test]
242    async fn exports_a_server_span_with_matching_request_context() {
243        use rust_zero_core::Telemetry;
244
245        let telemetry = Telemetry::local("test-api", 1.0).unwrap();
246        let inbound = "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01";
247        let app = test::init_service(App::new().wrap(OpenTelemetryTracing::new()).route(
248            "/users",
249            web::get().to(|request: HttpRequest| async move {
250                HttpResponse::Ok().body(OpenTelemetryTracing::context(&request).unwrap().trace_id())
251            }),
252        ))
253        .await;
254        let response = test::call_service(
255            &app,
256            test::TestRequest::get()
257                .uri("/users")
258                .insert_header(("traceparent", inbound))
259                .to_request(),
260        )
261        .await;
262
263        assert_eq!(response.status(), actix_web::http::StatusCode::OK);
264        assert_eq!(
265            test::read_body(response).await,
266            "4bf92f3577b34da6a3ce929d0e0e4736"
267        );
268        telemetry.force_flush().unwrap();
269    }
270}