salvo_extra/
request_id.rs1use 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
28pub const REQUEST_ID_KEY: &str = "::salvo::request_id";
30
31pub trait RequestIdDepotExt {
33 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#[non_exhaustive]
46pub struct RequestId {
47 pub header_name: HeaderName,
49 pub overwrite: bool,
51 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 #[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 #[must_use]
77 pub fn header_name(mut self, name: HeaderName) -> Self {
78 self.header_name = name;
79 self
80 }
81
82 #[must_use]
84 pub fn overwrite(mut self, overwrite: bool) -> Self {
85 self.overwrite = overwrite;
86 self
87 }
88
89 #[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
119pub trait IdGenerator {
121 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#[derive(Default, Debug)]
136pub struct UlidGenerator {}
137impl UlidGenerator {
138 #[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}