1use crate::middleware::{named, MwEntry, Next};
4use crate::request::Request;
5use tracing::Instrument;
6
7tokio::task_local! {
8 static CURRENT_REQUEST_ID: String;
9}
10
11pub fn current_request_id() -> Option<String> {
16 CURRENT_REQUEST_ID.try_with(|s| s.clone()).ok()
17}
18
19#[derive(Debug, Clone, PartialEq, Eq)]
21pub struct RequestId(pub String);
22
23impl std::fmt::Display for RequestId {
24 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
25 self.0.fmt(f)
26 }
27}
28
29impl AsRef<str> for RequestId {
30 fn as_ref(&self) -> &str {
31 &self.0
32 }
33}
34
35pub fn request_id() -> MwEntry {
40 named("request-id", |mut req: Request, next: Next| async move {
41 ensure_request_id(&mut req);
42 let id = req
43 .get::<RequestId>()
44 .map(|r| r.0.clone())
45 .unwrap_or_default();
46 let method = req.method.as_str().to_string();
47 let path = req.path.clone();
48 let quiet = crate::middleware::logger_should_skip(&path);
49 let id_for_span = id.clone();
50 let id_for_header = id.clone();
51
52 let run = CURRENT_REQUEST_ID.scope(id, async move {
53 let mut res = next(req).await;
54 if !id_for_header.is_empty() {
55 res = res.header("x-request-id", &id_for_header);
56 }
57 res
58 });
59
60 if quiet {
61 let span = tracing::debug_span!(
62 "http.server",
63 request_id = %id_for_span,
64 method = %method,
65 path = %path,
66 otel.kind = "server",
67 );
68 run.instrument(span).await
69 } else {
70 let span = tracing::info_span!(
71 "http.server",
72 request_id = %id_for_span,
73 method = %method,
74 path = %path,
75 otel.kind = "server",
76 );
77 run.instrument(span).await
78 }
79 })
80}
81
82pub fn ensure_request_id(req: &mut Request) {
84 if req.get::<RequestId>().is_some() {
85 return;
86 }
87 let id = req
88 .header("x-request-id")
89 .filter(|s| !s.is_empty())
90 .map(str::to_owned)
91 .unwrap_or_else(generate_request_id);
92 req.set(RequestId(id));
93}
94
95fn generate_request_id() -> String {
96 use std::sync::atomic::{AtomicU64, Ordering};
97 static COUNTER: AtomicU64 = AtomicU64::new(1);
98 let n = COUNTER.fetch_add(1, Ordering::Relaxed);
99 let mut bytes = [0u8; 8];
100 let _ = getrandom::getrandom(&mut bytes);
101 let entropy = u64::from_le_bytes(bytes);
102 format!("req-{entropy:016x}-{n:x}")
103}
104
105#[cfg(test)]
106mod tests {
107 use super::*;
108 use crate::middleware::build_chain;
109 use crate::response::Response;
110 use http::Method;
111 use std::sync::Arc;
112
113 #[tokio::test]
114 async fn echoes_and_reuses_inbound() {
115 let leaf: crate::handler::Handler = Arc::new(|req: Request| {
116 Box::pin(async move {
117 let id = req.get::<RequestId>().unwrap().0.clone();
118 Response::text(id)
119 })
120 });
121 let chain = build_chain(&[request_id().mw], leaf);
122 let mut req = Request::new(Method::GET, "/");
123 req.headers
124 .insert("x-request-id", "abc-123".parse().unwrap());
125 let res = chain(req).await;
126 assert_eq!(res.body_bytes(), Some(b"abc-123".as_slice()));
127 assert_eq!(
128 res.headers
129 .get("x-request-id")
130 .and_then(|v| v.to_str().ok()),
131 Some("abc-123")
132 );
133 }
134
135 #[tokio::test]
136 async fn generates_when_missing() {
137 let leaf: crate::handler::Handler = Arc::new(|req: Request| {
138 Box::pin(async move {
139 let id = req.get::<RequestId>().unwrap().0.clone();
140 Response::text(id)
141 })
142 });
143 let chain = build_chain(&[request_id().mw.clone()], leaf);
144 let res = chain(Request::new(Method::GET, "/")).await;
145 let body = String::from_utf8(res.body_bytes().unwrap().to_vec()).unwrap();
146 assert!(body.starts_with("req-"), "{body}");
147 assert_eq!(
148 res.headers
149 .get("x-request-id")
150 .and_then(|v| v.to_str().ok()),
151 Some(body.as_str())
152 );
153 }
154}