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 {
37 named("request-id", |mut req: Request, next: Next| async move {
38 ensure_request_id(&mut req);
39 let id = req
40 .get::<RequestId>()
41 .map(|r| r.0.clone())
42 .unwrap_or_default();
43 let method = req.method.as_str().to_string();
44 let path = req.path.clone();
45 let span = tracing::info_span!(
46 "http.server",
47 request_id = %id,
48 method = %method,
49 path = %path,
50 otel.kind = "server",
51 );
52 let id_for_header = id.clone();
53 CURRENT_REQUEST_ID
54 .scope(id, async move {
55 let mut res = next(req).await;
56 if !id_for_header.is_empty() {
57 res = res.header("x-request-id", &id_for_header);
58 }
59 res
60 })
61 .instrument(span)
62 .await
63 })
64}
65
66pub fn ensure_request_id(req: &mut Request) {
68 if req.get::<RequestId>().is_some() {
69 return;
70 }
71 let id = req
72 .header("x-request-id")
73 .filter(|s| !s.is_empty())
74 .map(str::to_owned)
75 .unwrap_or_else(generate_request_id);
76 req.set(RequestId(id));
77}
78
79fn generate_request_id() -> String {
80 use std::sync::atomic::{AtomicU64, Ordering};
81 static COUNTER: AtomicU64 = AtomicU64::new(1);
82 let n = COUNTER.fetch_add(1, Ordering::Relaxed);
83 let mut bytes = [0u8; 8];
84 let _ = getrandom::getrandom(&mut bytes);
85 let entropy = u64::from_le_bytes(bytes);
86 format!("req-{entropy:016x}-{n:x}")
87}
88
89#[cfg(test)]
90mod tests {
91 use super::*;
92 use crate::middleware::build_chain;
93 use crate::response::Response;
94 use http::Method;
95 use std::sync::Arc;
96
97 #[tokio::test]
98 async fn echoes_and_reuses_inbound() {
99 let leaf: crate::handler::Handler = Arc::new(|req: Request| {
100 Box::pin(async move {
101 let id = req.get::<RequestId>().unwrap().0.clone();
102 Response::text(id)
103 })
104 });
105 let chain = build_chain(&[request_id().mw], leaf);
106 let mut req = Request::new(Method::GET, "/");
107 req.headers
108 .insert("x-request-id", "abc-123".parse().unwrap());
109 let res = chain(req).await;
110 assert_eq!(res.body_bytes(), Some(b"abc-123".as_slice()));
111 assert_eq!(
112 res.headers
113 .get("x-request-id")
114 .and_then(|v| v.to_str().ok()),
115 Some("abc-123")
116 );
117 }
118
119 #[tokio::test]
120 async fn generates_when_missing() {
121 let leaf: crate::handler::Handler = Arc::new(|req: Request| {
122 Box::pin(async move {
123 let id = req.get::<RequestId>().unwrap().0.clone();
124 Response::text(id)
125 })
126 });
127 let chain = build_chain(&[request_id().mw.clone()], leaf);
128 let res = chain(Request::new(Method::GET, "/")).await;
129 let body = String::from_utf8(res.body_bytes().unwrap().to_vec()).unwrap();
130 assert!(body.starts_with("req-"), "{body}");
131 assert_eq!(
132 res.headers
133 .get("x-request-id")
134 .and_then(|v| v.to_str().ok()),
135 Some(body.as_str())
136 );
137 }
138}