aster_forge_actix_middleware/
request_id.rs1use 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#[derive(Clone, Debug)]
18pub struct RequestId(pub String);
19
20pub 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
41pub 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}