Skip to main content

aster_forge_actix_middleware/
request_id.rs

1//! Request id middleware.
2//!
3//! The middleware assigns a UUID v4 to every request, stores it in request
4//! extensions, adds it to the `X-Request-ID` response header, and instruments
5//! the downstream service call with a tracing span containing request metadata.
6
7use actix_web::{
8    Error, HttpMessage,
9    dev::{Service, ServiceRequest, ServiceResponse, Transform, forward_ready},
10    http::header::{HeaderName, HeaderValue},
11};
12use futures::future::{LocalBoxFuture, Ready, ok};
13use std::rc::Rc;
14use tracing::Instrument;
15
16/// Request id value stored in Actix request extensions.
17#[derive(Clone, Debug)]
18pub struct RequestId(pub String);
19
20/// Actix middleware that creates and propagates request ids.
21pub struct RequestIdMiddleware;
22
23impl<S, B> Transform<S, ServiceRequest> for RequestIdMiddleware
24where
25    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
26    B: 'static,
27{
28    type Response = ServiceResponse<B>;
29    type Error = Error;
30    type InitError = ();
31    type Transform = RequestIdService<S>;
32    type Future = Ready<Result<Self::Transform, Self::InitError>>;
33
34    fn new_transform(&self, service: S) -> Self::Future {
35        ok(RequestIdService {
36            service: Rc::new(service),
37        })
38    }
39}
40
41/// Service wrapper installed by [`RequestIdMiddleware`].
42pub struct RequestIdService<S> {
43    service: Rc<S>,
44}
45
46impl<S, B> Service<ServiceRequest> for RequestIdService<S>
47where
48    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
49    B: 'static,
50{
51    type Response = ServiceResponse<B>;
52    type Error = Error;
53    type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;
54
55    forward_ready!(service);
56
57    fn call(&self, req: ServiceRequest) -> Self::Future {
58        let svc = self.service.clone();
59        let request_id = uuid::Uuid::new_v4().to_string();
60        let method = req.method().to_string();
61        let path = req.path().to_string();
62
63        req.extensions_mut().insert(RequestId(request_id.clone()));
64
65        let span = tracing::info_span!(
66            "request",
67            request_id = %request_id,
68            method = %method,
69            path = %path,
70            user_id = tracing::field::Empty,
71        );
72
73        Box::pin(
74            async move {
75                let mut resp = svc.call(req).await?;
76
77                if let Ok(val) = HeaderValue::from_str(&request_id) {
78                    resp.headers_mut()
79                        .insert(HeaderName::from_static("x-request-id"), val);
80                }
81
82                Ok(resp)
83            }
84            .instrument(span),
85        )
86    }
87}
88
89#[cfg(test)]
90mod tests {
91    use super::{RequestId, RequestIdMiddleware};
92    use actix_web::{HttpMessage, HttpResponse, http::header, test, web};
93
94    #[actix_web::test]
95    async fn request_id_is_stored_and_returned() {
96        let app = test::init_service(actix_web::App::new().wrap(RequestIdMiddleware).route(
97            "/",
98            web::get().to(|req: actix_web::HttpRequest| async move {
99                let request_id = req
100                    .extensions()
101                    .get::<RequestId>()
102                    .map(|value| value.0.clone())
103                    .unwrap_or_default();
104                HttpResponse::Ok().body(request_id)
105            }),
106        ))
107        .await;
108
109        let request = test::TestRequest::get().uri("/").to_request();
110        let response = test::call_service(&app, request).await;
111        let header_value = response
112            .headers()
113            .get(header::HeaderName::from_static("x-request-id"))
114            .expect("request id header should be present")
115            .to_str()
116            .expect("request id should be ASCII")
117            .to_string();
118
119        let body = test::read_body(response).await;
120        assert_eq!(body.as_ref(), header_value.as_bytes());
121    }
122}