1use std::convert::Infallible;
4
5use axum::extract::{FromRequestParts, Request};
6use axum::http::HeaderValue;
7use axum::http::request::Parts;
8use axum::middleware::Next;
9use axum::response::Response;
10
11pub(crate) const HEADER: &str = "x-request-id";
12
13#[derive(Debug, Clone, PartialEq, Eq)]
27pub struct RequestId(pub String);
28
29impl std::fmt::Display for RequestId {
30 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
31 f.write_str(&self.0)
32 }
33}
34
35impl<S: Send + Sync> FromRequestParts<S> for RequestId {
36 type Rejection = Infallible;
37
38 async fn from_request_parts(parts: &mut Parts, _: &S) -> Result<Self, Infallible> {
39 Ok(parts
40 .extensions
41 .get::<RequestId>()
42 .cloned()
43 .unwrap_or_else(|| RequestId(String::new())))
44 }
45}
46
47fn acceptable(id: &str) -> bool {
48 (8..=64).contains(&id.len())
49 && id
50 .bytes()
51 .all(|b| b.is_ascii_alphanumeric() || matches!(b, b'.' | b'_' | b'-'))
52}
53
54pub(crate) async fn middleware(mut req: Request, next: Next) -> Response {
57 let id = req
58 .headers()
59 .get(HEADER)
60 .and_then(|v| v.to_str().ok())
61 .filter(|id| acceptable(id))
62 .map_or_else(|| crate::random_token()[..20].to_owned(), str::to_owned);
63 if let Ok(value) = HeaderValue::from_str(&id) {
64 req.headers_mut().insert(HEADER, value.clone());
65 req.extensions_mut().insert(RequestId(id));
66 let mut res = next.run(req).await;
67 res.headers_mut().insert(HEADER, value);
68 return res;
69 }
70 next.run(req).await
71}
72
73#[cfg(test)]
74mod tests {
75 #[test]
76 fn accepts_only_safe_ids() {
77 assert!(super::acceptable("abcd-1234_ef.gh"));
78 assert!(!super::acceptable("short"));
79 assert!(!super::acceptable("has space in it"));
80 assert!(!super::acceptable("new\nline-injected"));
81 assert!(!super::acceptable(&"a".repeat(65)));
82 }
83}