Skip to main content

renox_core/
request_id.rs

1//! Request ids: one per request, in the logs, the response and error reports.
2
3use std::convert::Infallible;
4
5use axum::extract::{FromRequestParts, Request};
6use axum::http::HeaderValue;
7use axum::http::request::Parts;
8use axum::middleware::Next;
9use axum::response::Response;
10
11pub(crate) const HEADER: &str = "x-request-id";
12
13/// The request's id: the `X-Request-Id` a proxy sent (when it looks safe:
14/// 8–64 letters, digits, `.`, `_` or `-`), otherwise a new random one. It's
15/// in every log line of the request, in the response's `X-Request-Id`, and
16/// in error reports, so a visitor's "request id" leads to the logs.
17///
18/// ```
19/// # use renox::prelude::*;
20/// use renox::RequestId;
21///
22/// async fn show(id: RequestId) -> String {
23///     format!("Quote this if something went wrong: {id}")
24/// }
25/// ```
26#[derive(Debug, Clone, PartialEq, Eq)]
27pub struct RequestId(pub String);
28
29impl std::fmt::Display for RequestId {
30    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
31        f.write_str(&self.0)
32    }
33}
34
35impl<S: Send + Sync> FromRequestParts<S> for RequestId {
36    type Rejection = Infallible;
37
38    async fn from_request_parts(parts: &mut Parts, _: &S) -> Result<Self, Infallible> {
39        Ok(parts
40            .extensions
41            .get::<RequestId>()
42            .cloned()
43            .unwrap_or_else(|| RequestId(String::new())))
44    }
45}
46
47fn acceptable(id: &str) -> bool {
48    (8..=64).contains(&id.len())
49        && id
50            .bytes()
51            .all(|b| b.is_ascii_alphanumeric() || matches!(b, b'.' | b'_' | b'-'))
52}
53
54/// Outermost: picks the id and puts it on the request (for the log span)
55/// and on the response.
56pub(crate) async fn middleware(mut req: Request, next: Next) -> Response {
57    let id = req
58        .headers()
59        .get(HEADER)
60        .and_then(|v| v.to_str().ok())
61        .filter(|id| acceptable(id))
62        .map_or_else(|| crate::random_token()[..20].to_owned(), str::to_owned);
63    if let Ok(value) = HeaderValue::from_str(&id) {
64        req.headers_mut().insert(HEADER, value.clone());
65        req.extensions_mut().insert(RequestId(id));
66        let mut res = next.run(req).await;
67        res.headers_mut().insert(HEADER, value);
68        return res;
69    }
70    next.run(req).await
71}
72
73#[cfg(test)]
74mod tests {
75    #[test]
76    fn accepts_only_safe_ids() {
77        assert!(super::acceptable("abcd-1234_ef.gh"));
78        assert!(!super::acceptable("short"));
79        assert!(!super::acceptable("has space in it"));
80        assert!(!super::acceptable("new\nline-injected"));
81        assert!(!super::acceptable(&"a".repeat(65)));
82    }
83}