Skip to main content

axum_observability/
request_id.rs

1use std::{error::Error, fmt, str::FromStr};
2
3/// A request identifier selected by the configured request-ID policy.
4///
5/// Values created through [`parse`](Self::parse), [`FromStr`], or [`TryFrom`]
6/// satisfy the baseline grammar: 1 to [`MAX_LEN`](Self::MAX_LEN) ASCII
7/// URI-unreserved bytes. Middleware can also construct this type for one
8/// native-safe caller value admitted by a custom validator; such a value is not
9/// guaranteed to satisfy the baseline grammar.
10#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
11pub struct RequestId(Box<str>);
12
13impl RequestId {
14    /// Maximum request-ID length accepted by the baseline parser, in bytes.
15    ///
16    /// A caller value admitted by a custom middleware validator can be longer.
17    pub const MAX_LEN: usize = 128;
18
19    /// Parses and validates a request identifier.
20    ///
21    /// # Errors
22    ///
23    /// Returns the first deterministic baseline validation failure without
24    /// retaining or echoing the rejected value.
25    ///
26    /// # Examples
27    ///
28    /// ```
29    /// use axum_observability::RequestId;
30    ///
31    /// let request_id = RequestId::parse("request-42")?;
32    /// assert_eq!(request_id.as_str(), "request-42");
33    /// # Ok::<(), axum_observability::InvalidRequestId>(())
34    /// ```
35    pub fn parse(value: &str) -> Result<Self, InvalidRequestId> {
36        validate(value)?;
37        Ok(Self(value.into()))
38    }
39
40    pub(crate) fn from_native_header(value: &str) -> Option<Self> {
41        native_field_content(value).then(|| Self(value.into()))
42    }
43
44    /// Returns the validated identifier.
45    #[must_use]
46    pub fn as_str(&self) -> &str {
47        &self.0
48    }
49}
50
51pub(crate) fn native_field_content(value: &str) -> bool {
52    let bytes = value.as_bytes();
53    if bytes.is_empty()
54        || matches!(bytes.first(), Some(b' ' | b'\t'))
55        || matches!(bytes.last(), Some(b' ' | b'\t'))
56    {
57        return false;
58    }
59    bytes
60        .iter()
61        .all(|byte| *byte == b'\t' || *byte >= 0x20 && *byte != 0x7f)
62}
63
64impl AsRef<str> for RequestId {
65    fn as_ref(&self) -> &str {
66        self.as_str()
67    }
68}
69
70impl fmt::Display for RequestId {
71    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
72        formatter.write_str(self.as_str())
73    }
74}
75
76impl FromStr for RequestId {
77    type Err = InvalidRequestId;
78
79    fn from_str(value: &str) -> Result<Self, Self::Err> {
80        Self::parse(value)
81    }
82}
83
84impl TryFrom<&str> for RequestId {
85    type Error = InvalidRequestId;
86
87    fn try_from(value: &str) -> Result<Self, Self::Error> {
88        Self::parse(value)
89    }
90}
91
92impl TryFrom<String> for RequestId {
93    type Error = InvalidRequestId;
94
95    fn try_from(value: String) -> Result<Self, Self::Error> {
96        validate(&value)?;
97        Ok(Self(value.into_boxed_str()))
98    }
99}
100
101/// Reason a request identifier failed baseline validation.
102#[derive(Clone, Copy, Debug, Eq, PartialEq)]
103#[non_exhaustive]
104pub enum InvalidRequestId {
105    /// The identifier was empty.
106    Empty,
107    /// The identifier exceeded [`RequestId::MAX_LEN`] bytes.
108    TooLong {
109        /// Rejected byte length.
110        length: usize,
111    },
112    /// The identifier contained a byte outside the URI-unreserved set.
113    InvalidCharacter {
114        /// Byte index of the first invalid character.
115        index: usize,
116    },
117}
118
119impl fmt::Display for InvalidRequestId {
120    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
121        match self {
122            Self::Empty => formatter.write_str("request ID must not be empty"),
123            Self::TooLong { length } => {
124                write!(formatter, "request ID length {length} exceeds 128 bytes")
125            }
126            Self::InvalidCharacter { index } => write!(
127                formatter,
128                "request ID contains an invalid character at byte index {index}"
129            ),
130        }
131    }
132}
133
134impl Error for InvalidRequestId {}
135
136fn validate(value: &str) -> Result<(), InvalidRequestId> {
137    if value.is_empty() {
138        return Err(InvalidRequestId::Empty);
139    }
140    if value.len() > RequestId::MAX_LEN {
141        return Err(InvalidRequestId::TooLong {
142            length: value.len(),
143        });
144    }
145    if let Some(index) = value.bytes().position(|byte| {
146        !(byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'.' | b'_' | b'~'))
147    }) {
148        return Err(InvalidRequestId::InvalidCharacter { index });
149    }
150    Ok(())
151}
152
153#[cfg(test)]
154mod tests {
155    use std::str::FromStr as _;
156
157    use super::{InvalidRequestId, RequestId, native_field_content};
158
159    #[test]
160    fn native_field_content_admits_internal_htab_but_rejects_controls() {
161        assert!(native_field_content("tenant\trequest"));
162        for value in ["tenant\0request", "tenant\x1frequest", "tenant\x7frequest"] {
163            assert!(!native_field_content(value), "admitted {value:?}");
164        }
165    }
166
167    #[test]
168    fn accepts_exact_length_boundaries_and_conversion_forms() {
169        let one = RequestId::parse("a").expect("one byte");
170        assert_eq!(one.as_str(), "a");
171        assert_eq!(one.as_ref(), "a");
172
173        let maximum = "aZ09-._~".repeat(16);
174        assert_eq!(maximum.len(), RequestId::MAX_LEN);
175        assert_eq!(
176            RequestId::from_str(&maximum).expect("from str").as_str(),
177            maximum
178        );
179        assert_eq!(
180            RequestId::try_from(maximum.as_str())
181                .expect("borrowed")
182                .as_str(),
183            maximum
184        );
185        assert_eq!(
186            RequestId::try_from(maximum.clone())
187                .expect("owned")
188                .as_str(),
189            maximum
190        );
191    }
192
193    #[test]
194    fn reports_deterministic_non_sensitive_failures() {
195        let cases = [
196            ("", InvalidRequestId::Empty),
197            (
198                &"secret".repeat(22),
199                InvalidRequestId::TooLong { length: 132 },
200            ),
201            (
202                "safe/secret",
203                InvalidRequestId::InvalidCharacter { index: 4 },
204            ),
205            (
206                "safeümlaut",
207                InvalidRequestId::InvalidCharacter { index: 4 },
208            ),
209        ];
210
211        for (value, expected) in cases {
212            let error = RequestId::parse(value).expect_err("invalid request ID");
213            assert_eq!(error, expected);
214            if !value.is_empty() {
215                assert!(!error.to_string().contains(value));
216            }
217        }
218    }
219
220    #[test]
221    fn checks_length_before_character_content() {
222        let value = format!("/{}", "a".repeat(128));
223        assert_eq!(
224            RequestId::parse(&value),
225            Err(InvalidRequestId::TooLong { length: 129 })
226        );
227    }
228
229    #[test]
230    fn conversion_forms_enforce_identical_invalid_boundaries() {
231        let oversized = "a".repeat(RequestId::MAX_LEN + 1);
232        let expected = InvalidRequestId::TooLong { length: 129 };
233        assert_eq!(RequestId::from_str(&oversized), Err(expected));
234        assert_eq!(RequestId::try_from(oversized.as_str()), Err(expected));
235        assert_eq!(RequestId::try_from(oversized), Err(expected));
236    }
237
238    #[test]
239    fn validation_errors_have_stable_redacted_messages() {
240        assert_eq!(
241            InvalidRequestId::Empty.to_string(),
242            "request ID must not be empty"
243        );
244        assert_eq!(
245            InvalidRequestId::TooLong { length: 129 }.to_string(),
246            "request ID length 129 exceeds 128 bytes"
247        );
248        assert_eq!(
249            InvalidRequestId::InvalidCharacter { index: 7 }.to_string(),
250            "request ID contains an invalid character at byte index 7"
251        );
252    }
253}