Skip to main content

salvo_extra/
request_id.rs

1//! Request id middleware.
2//!
3//! # Example
4//!
5//! ```no_run
6//! use salvo_core::prelude::*;
7//! use salvo_extra::request_id::RequestId;
8//!
9//! #[handler]
10//! async fn hello(req: &mut Request) -> String {
11//!     format!("Request id: {:?}", req.header::<String>("x-request-id"))
12//! }
13//!
14//! #[tokio::main]
15//! async fn main() {
16//!     let acceptor = TcpListener::new("0.0.0.0:8698").bind().await;
17//!     let router = Router::new().hoop(RequestId::new()).get(hello);
18//!     Server::new(acceptor).serve(router).await;
19//! }
20//! ```
21use std::fmt::{self, Debug, Formatter};
22use tracing::Instrument;
23use ulid::Ulid;
24
25use salvo_core::http::{HeaderValue, Request, Response, header::HeaderName};
26use salvo_core::{Depot, FlowCtrl, Handler, async_trait};
27
28/// Key for incoming flash messages in depot.
29pub const REQUEST_ID_KEY: &str = "::salvo::request_id";
30
31/// Extension for Depot.
32pub trait RequestIdDepotExt {
33    /// Get request id reference from depot.
34    fn csrf_token(&self) -> Option<&str>;
35}
36
37impl RequestIdDepotExt for Depot {
38    #[inline]
39    fn csrf_token(&self) -> Option<&str> {
40        self.get::<String>(REQUEST_ID_KEY).map(|v| &**v).ok()
41    }
42}
43
44/// A middleware for generate request id.
45#[non_exhaustive]
46pub struct RequestId {
47    /// The header name for request id.
48    pub header_name: HeaderName,
49    /// Whether overwrite exists request id. Default is `true`
50    pub overwrite: bool,
51    /// The generator for request id.
52    pub generator: Box<dyn IdGenerator + Send + Sync>,
53}
54
55impl Debug for RequestId {
56    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
57        f.debug_struct("RequestId")
58            .field("header_name", &self.header_name)
59            .field("overwrite", &self.overwrite)
60            .finish()
61    }
62}
63
64impl RequestId {
65    /// Create new `CatchPanic` middleware.
66    #[must_use]
67    pub fn new() -> Self {
68        Self {
69            header_name: HeaderName::from_static("x-request-id"),
70            overwrite: true,
71            generator: Box::new(UlidGenerator::new()),
72        }
73    }
74
75    /// Set the header name for request id.
76    #[must_use]
77    pub fn header_name(mut self, name: HeaderName) -> Self {
78        self.header_name = name;
79        self
80    }
81
82    /// Set whether overwrite exists request id. Default is `true`.
83    #[must_use]
84    pub fn overwrite(mut self, overwrite: bool) -> Self {
85        self.overwrite = overwrite;
86        self
87    }
88
89    /// Set the generator for request id.
90    #[must_use]
91    pub fn generator(mut self, generator: impl IdGenerator + Send + Sync + 'static) -> Self {
92        self.generator = Box::new(generator);
93        self
94    }
95
96    fn generate_id(&self, req: &mut Request, depot: &mut Depot) -> HeaderValue {
97        let id = self.generator.generate(req, depot);
98        match HeaderValue::from_str(&id) {
99            Ok(header_value) => header_value,
100            Err(error) => {
101                tracing::warn!(
102                    error = ?error,
103                    generated_id = %id,
104                    "request id generator returned an invalid header value; falling back to ULID"
105                );
106                HeaderValue::from_str(&Ulid::new().to_string())
107                    .expect("ULID should always be a valid header value")
108            }
109        }
110    }
111}
112
113impl Default for RequestId {
114    fn default() -> Self {
115        Self::new()
116    }
117}
118
119/// A trait for generate request id.
120pub trait IdGenerator {
121    /// Generate a new request id.
122    fn generate(&self, req: &mut Request, depot: &mut Depot) -> String;
123}
124
125impl<F> IdGenerator for F
126where
127    F: Fn() -> String + Send + Sync,
128{
129    fn generate(&self, _req: &mut Request, _depot: &mut Depot) -> String {
130        self()
131    }
132}
133
134/// A generator for generate request id with ulid.
135#[derive(Default, Debug)]
136pub struct UlidGenerator {}
137impl UlidGenerator {
138    /// Create new `UlidGenerator`.
139    #[must_use]
140    pub fn new() -> Self {
141        Self {}
142    }
143}
144impl IdGenerator for UlidGenerator {
145    fn generate(&self, _req: &mut Request, _depot: &mut Depot) -> String {
146        Ulid::new().to_string()
147    }
148}
149
150#[async_trait]
151impl Handler for RequestId {
152    async fn handle(
153        &self,
154        req: &mut Request,
155        depot: &mut Depot,
156        res: &mut Response,
157        ctrl: &mut FlowCtrl,
158    ) {
159        let request_id = match req.headers().get(&self.header_name) {
160            None => self.generate_id(req, depot),
161            Some(value) => {
162                if self.overwrite {
163                    self.generate_id(req, depot)
164                } else {
165                    value.clone()
166                }
167            }
168        };
169
170        let _ = req.add_header(self.header_name.clone(), &request_id, false);
171
172        let span = tracing::info_span!("request", ?request_id);
173        res.headers_mut()
174            .insert(self.header_name.clone(), request_id.clone());
175        depot.insert(REQUEST_ID_KEY, request_id);
176
177        async move {
178            ctrl.call_next(req, depot, res).await;
179        }
180        .instrument(span)
181        .await;
182    }
183}
184#[cfg(test)]
185mod tests {
186    use salvo_core::prelude::*;
187    use salvo_core::test::{ResponseExt, TestClient};
188
189    use super::*;
190
191    #[tokio::test]
192    async fn test_request_id_added() {
193        let handler = RequestId::new();
194        let router = Router::new().hoop(handler).get(endpoint);
195        let service = Service::new(router);
196
197        let response = TestClient::get("http://127.0.0.1:8698/")
198            .send(&service)
199            .await;
200        assert_eq!(response.status_code, Some(StatusCode::OK));
201        assert!(response.headers.contains_key("x-request-id"));
202    }
203
204    #[tokio::test]
205    async fn test_request_id_overwrite() {
206        let handler = RequestId::new().overwrite(true);
207        let router = Router::new().hoop(handler).get(endpoint);
208        let service = Service::new(router);
209
210        let response = TestClient::get("http://127.0.0.1:8698/")
211            .add_header("x-request-id", "existing-id", true)
212            .send(&service)
213            .await;
214        assert_eq!(response.status_code, Some(StatusCode::OK));
215        assert_ne!(response.headers.get("x-request-id").unwrap(), "existing-id");
216    }
217
218    #[tokio::test]
219    async fn test_request_id_no_overwrite() {
220        let handler = RequestId::new().overwrite(false);
221        let router = Router::new().hoop(handler).get(endpoint);
222        let service = Service::new(router);
223
224        let response = TestClient::get("http://127.0.0.1:8698/")
225            .add_header("x-request-id", "existing-id", true)
226            .send(&service)
227            .await;
228        assert_eq!(response.status_code, Some(StatusCode::OK));
229        assert_eq!(response.headers.get("x-request-id").unwrap(), "existing-id");
230    }
231
232    #[tokio::test]
233    async fn test_custom_generator() {
234        let handler = RequestId::new().generator(|| "custom-id".to_owned());
235        let router = Router::new().hoop(handler).get(endpoint);
236        let service = Service::new(router);
237
238        let response = TestClient::get("http://127.0.0.1:8698/")
239            .send(&service)
240            .await;
241        assert_eq!(response.status_code, Some(StatusCode::OK));
242        assert_eq!(response.headers.get("x-request-id").unwrap(), "custom-id");
243    }
244
245    #[tokio::test]
246    async fn test_invalid_custom_generator_falls_back() {
247        let handler = RequestId::new().generator(|| "bad\r\nvalue".to_owned());
248        let router = Router::new().hoop(handler).get(endpoint);
249        let service = Service::new(router);
250
251        let response = TestClient::get("http://127.0.0.1:8698/")
252            .send(&service)
253            .await;
254        assert_eq!(response.status_code, Some(StatusCode::OK));
255        let request_id = response
256            .headers
257            .get("x-request-id")
258            .unwrap()
259            .to_str()
260            .unwrap();
261        assert_ne!(request_id, "bad\r\nvalue");
262        assert_eq!(request_id.len(), 26);
263    }
264
265    #[tokio::test]
266    async fn test_depot_storage() {
267        let handler = RequestId::new();
268        #[handler]
269        async fn depot_checker(depot: &mut Depot, res: &mut Response) {
270            let id = depot.get::<HeaderValue>(REQUEST_ID_KEY).unwrap().clone();
271            res.render(Text::Plain(id.to_str().unwrap().to_owned()));
272        }
273        let router = Router::new().hoop(handler).get(depot_checker);
274        let service = Service::new(router);
275
276        let mut response = TestClient::get("http://127.0.0.1:8698/")
277            .send(&service)
278            .await;
279        assert_eq!(response.status_code, Some(StatusCode::OK));
280        let header_id = response
281            .headers
282            .get("x-request-id")
283            .unwrap()
284            .to_str()
285            .unwrap().to_owned();
286        let body = response.take_string().await.unwrap();
287        assert_eq!(header_id, body);
288    }
289
290    #[handler]
291    async fn endpoint() {}
292}