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#[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#[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}