1use crate::middleware::{named, MwEntry, Next};
4use crate::request::Request;
5use tracing::Instrument;
6
7#[derive(Debug, Clone, PartialEq, Eq)]
9pub struct RequestId(pub String);
10
11impl std::fmt::Display for RequestId {
12 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
13 self.0.fmt(f)
14 }
15}
16
17impl AsRef<str> for RequestId {
18 fn as_ref(&self) -> &str {
19 &self.0
20 }
21}
22
23pub fn request_id() -> MwEntry {
25 named("request-id", |mut req: Request, next: Next| async move {
26 ensure_request_id(&mut req);
27 let id = req
28 .get::<RequestId>()
29 .map(|r| r.0.clone())
30 .unwrap_or_default();
31 let method = req.method.as_str().to_string();
32 let path = req.path.clone();
33 let span = tracing::info_span!(
34 "http.server",
35 request_id = %id,
36 method = %method,
37 path = %path,
38 otel.kind = "server",
39 );
40 async move {
41 let mut res = next(req).await;
42 if !id.is_empty() {
43 res = res.header("x-request-id", &id);
44 }
45 res
46 }
47 .instrument(span)
48 .await
49 })
50}
51
52pub fn ensure_request_id(req: &mut Request) {
54 if req.get::<RequestId>().is_some() {
55 return;
56 }
57 let id = req
58 .header("x-request-id")
59 .filter(|s| !s.is_empty())
60 .map(str::to_owned)
61 .unwrap_or_else(generate_request_id);
62 req.set(RequestId(id));
63}
64
65fn generate_request_id() -> String {
66 let nanos = std::time::SystemTime::now()
67 .duration_since(std::time::UNIX_EPOCH)
68 .unwrap_or_default()
69 .as_nanos();
70 format!("req-{nanos:x}")
71}
72
73#[cfg(test)]
74mod tests {
75 use super::*;
76 use crate::middleware::build_chain;
77 use crate::response::Response;
78 use http::Method;
79 use std::sync::Arc;
80
81 #[tokio::test]
82 async fn echoes_and_reuses_inbound() {
83 let leaf: crate::handler::Handler = Arc::new(|req: Request| {
84 Box::pin(async move {
85 let id = req.get::<RequestId>().unwrap().0.clone();
86 Response::text(id)
87 })
88 });
89 let chain = build_chain(&[request_id().mw], leaf);
90 let mut req = Request::new(Method::GET, "/");
91 req.headers
92 .insert("x-request-id", "abc-123".parse().unwrap());
93 let res = chain(req).await;
94 assert_eq!(res.body_bytes(), Some(b"abc-123".as_slice()));
95 assert_eq!(
96 res.headers
97 .get("x-request-id")
98 .and_then(|v| v.to_str().ok()),
99 Some("abc-123")
100 );
101 }
102
103 #[tokio::test]
104 async fn generates_when_missing() {
105 let leaf: crate::handler::Handler = Arc::new(|req: Request| {
106 Box::pin(async move {
107 let id = req.get::<RequestId>().unwrap().0.clone();
108 Response::text(id)
109 })
110 });
111 let chain = build_chain(&[request_id().mw.clone()], leaf);
112 let res = chain(Request::new(Method::GET, "/")).await;
113 let body = String::from_utf8(res.body_bytes().unwrap().to_vec()).unwrap();
114 assert!(body.starts_with("req-"), "{body}");
115 assert_eq!(
116 res.headers
117 .get("x-request-id")
118 .and_then(|v| v.to_str().ok()),
119 Some(body.as_str())
120 );
121 }
122}