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
7/// Per-request correlation id (inbound `x-request-id` or generated).
8#[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
23/// Ensure [`RequestId`], echo `x-request-id`, wrap in an `http.server` span.
24pub 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
52/// Set [`RequestId`] from `x-request-id` or generate one (idempotent).
53pub 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}