Skip to main content

sova_core/
request_id.rs

1//! Per-request correlation id (`x-request-id`).
2
3use crate::middleware::{named, MwEntry, Next};
4use crate::request::Request;
5use tracing::Instrument;
6
7tokio::task_local! {
8    static CURRENT_REQUEST_ID: String;
9}
10
11/// Request id for the current async task (set by [`request_id`] middleware).
12///
13/// Plugins (store / redis / tasks) can attach this to tracing events so DevTools
14/// correlates them with the open request bag.
15pub fn current_request_id() -> Option<String> {
16    CURRENT_REQUEST_ID.try_with(|s| s.clone()).ok()
17}
18
19/// Per-request correlation id (inbound `x-request-id` or generated).
20#[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
35/// Ensure [`RequestId`], echo `x-request-id`, wrap in an `http.server` span.
36///
37/// Quiet paths ([`crate::logger_should_skip`]) still get a request id, but use a
38/// `debug` span so `/_devtools/*` / favicon do not flood the console.
39pub 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
82/// Set [`RequestId`] from `x-request-id` or generate one (idempotent).
83pub 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}