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.
36pub 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
66/// Set [`RequestId`] from `x-request-id` or generate one (idempotent).
67pub 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}