Skip to main content

axum_observability/
context.rs

1use std::{error::Error, fmt};
2
3use axum::{
4    http::StatusCode,
5    response::{IntoResponse, Response},
6};
7
8use crate::{RequestId, TraceContextLevel};
9
10/// Validated inbound W3C trace context.
11#[derive(Clone, Debug, Eq, PartialEq)]
12pub struct TraceContext {
13    version: u8,
14    trace_id: String,
15    parent_id: String,
16    flags: u8,
17    level: TraceContextLevel,
18    traceparent: Box<[u8]>,
19    tracestate: Option<String>,
20}
21
22impl TraceContext {
23    pub(crate) fn new(
24        version: u8,
25        trace_id: String,
26        parent_id: String,
27        flags: u8,
28        level: TraceContextLevel,
29        traceparent: Box<[u8]>,
30    ) -> Self {
31        Self {
32            version,
33            trace_id,
34            parent_id,
35            flags,
36            level,
37            traceparent,
38            tracestate: None,
39        }
40    }
41
42    pub(crate) fn with_tracestate(mut self, tracestate: Option<String>) -> Self {
43        self.tracestate = tracestate;
44        self
45    }
46
47    /// Validated 32-character lowercase trace identifier.
48    #[must_use]
49    pub fn trace_id(&self) -> &str {
50        &self.trace_id
51    }
52
53    /// Incoming 16-character lowercase parent identifier.
54    #[must_use]
55    pub fn parent_id(&self) -> &str {
56        &self.parent_id
57    }
58
59    /// Raw W3C trace flags byte.
60    #[must_use]
61    pub const fn flags(&self) -> u8 {
62        self.flags
63    }
64
65    /// Whether the sampled flag is set.
66    #[must_use]
67    pub const fn sampled(&self) -> bool {
68        self.flags & 1 == 1
69    }
70
71    /// Selected W3C Trace Context level.
72    #[must_use]
73    pub const fn trace_context_level(&self) -> TraceContextLevel {
74        self.level
75    }
76
77    /// Whether the caller marked the trace ID as random in Level 2 mode.
78    ///
79    /// Level 1 does not assign portable meaning to this flag, so it returns
80    /// `None` even when bit one is set.
81    #[must_use]
82    pub const fn trace_id_random(&self) -> Option<bool> {
83        match (self.level, self.version) {
84            (TraceContextLevel::Level2, 0) => Some(self.flags & 2 == 2),
85            (TraceContextLevel::Level1 | TraceContextLevel::Level2, _) => None,
86        }
87    }
88
89    /// Accepted raw `traceparent` bytes.
90    #[must_use]
91    pub fn traceparent_bytes(&self) -> &[u8] {
92        &self.traceparent
93    }
94
95    /// Accepted raw `traceparent` as UTF-8 when its opaque future suffix is text.
96    #[must_use]
97    pub fn traceparent(&self) -> Option<&str> {
98        std::str::from_utf8(&self.traceparent).ok()
99    }
100
101    /// Accepted combined `tracestate`, when valid.
102    #[must_use]
103    pub fn tracestate(&self) -> Option<&str> {
104        self.tracestate.as_deref()
105    }
106}
107
108/// Correlation metadata installed in every observed request's extensions.
109#[derive(Clone, Debug, Eq, PartialEq)]
110pub struct RequestContext {
111    request_id: RequestId,
112    correlation_id: String,
113    trace_context: Option<TraceContext>,
114}
115
116impl<S> axum::extract::FromRequestParts<S> for RequestContext
117where
118    S: Send + Sync,
119{
120    type Rejection = MissingRequestContext;
121
122    async fn from_request_parts(
123        parts: &mut axum::http::request::Parts,
124        _state: &S,
125    ) -> Result<Self, Self::Rejection> {
126        parts
127            .extensions
128            .get::<Self>()
129            .cloned()
130            .ok_or(MissingRequestContext)
131    }
132}
133
134/// Rejection returned when [`RequestContext`] is extracted without the
135/// observability middleware installed.
136///
137/// Handlers can explicitly retain the typed rejection when middleware is
138/// optional during composition:
139///
140/// ```
141/// use axum_observability::{MissingRequestContext, RequestContext};
142///
143/// async fn request_id(
144///     context: Result<RequestContext, MissingRequestContext>,
145/// ) -> Result<String, MissingRequestContext> {
146///     Ok(context?.request_id().to_string())
147/// }
148/// # let _ = request_id;
149/// ```
150#[derive(Clone, Copy, Debug, Eq, PartialEq)]
151#[non_exhaustive]
152pub struct MissingRequestContext;
153
154impl fmt::Display for MissingRequestContext {
155    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
156        formatter.write_str("request context unavailable")
157    }
158}
159
160impl Error for MissingRequestContext {}
161
162impl IntoResponse for MissingRequestContext {
163    fn into_response(self) -> Response {
164        (StatusCode::INTERNAL_SERVER_ERROR, self.to_string()).into_response()
165    }
166}
167
168impl RequestContext {
169    pub(crate) fn new(request_id: RequestId, trace_context: Option<TraceContext>) -> Self {
170        let correlation_id = trace_context.as_ref().map_or_else(
171            || request_id.as_str().to_owned(),
172            |trace| trace.trace_id.clone(),
173        );
174        Self {
175            request_id,
176            correlation_id,
177            trace_context,
178        }
179    }
180
181    /// Validated or generated request identifier.
182    #[must_use]
183    pub const fn request_id(&self) -> &RequestId {
184        &self.request_id
185    }
186
187    /// Valid trace identifier, or the request identifier when no trace exists.
188    #[must_use]
189    pub fn correlation_id(&self) -> &str {
190        &self.correlation_id
191    }
192
193    /// Validated inbound trace context.
194    #[must_use]
195    pub const fn trace_context(&self) -> Option<&TraceContext> {
196        self.trace_context.as_ref()
197    }
198}
199
200/// Stable application operation name attached to request or response
201/// extensions.
202#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
203pub struct OperationId(&'static str);
204
205impl OperationId {
206    /// Creates an operation identifier from static route metadata.
207    ///
208    /// The value should be a stable semantic operation name. It must not
209    /// contain request-derived or user-derived data. Callers are responsible
210    /// for uniqueness within their application.
211    ///
212    /// # Panics
213    ///
214    /// Panics when `value` is empty.
215    ///
216    /// # Examples
217    ///
218    /// ```
219    /// use axum_observability::OperationId;
220    ///
221    /// const LIST_ITEMS: OperationId = OperationId::from_static("list-items");
222    /// assert_eq!(LIST_ITEMS.as_str(), "list-items");
223    /// ```
224    #[track_caller]
225    #[must_use]
226    pub const fn from_static(value: &'static str) -> Self {
227        assert!(!value.is_empty(), "operation ID must not be empty");
228        Self(value)
229    }
230
231    /// Returns the operation identifier.
232    #[must_use]
233    pub fn as_str(&self) -> &str {
234        self.0
235    }
236}
237
238impl AsRef<str> for OperationId {
239    fn as_ref(&self) -> &str {
240        self.as_str()
241    }
242}
243
244impl fmt::Display for OperationId {
245    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
246        formatter.write_str(self.as_str())
247    }
248}
249
250#[cfg(test)]
251mod tests {
252    use super::OperationId;
253
254    #[test]
255    fn operation_id_preserves_nonempty_static_values() {
256        const OPERATION: OperationId = OperationId::from_static("create-item");
257        assert_eq!(OPERATION.as_str(), "create-item");
258        assert_eq!(OPERATION.as_ref(), "create-item");
259        assert_eq!(OPERATION.to_string(), "create-item");
260    }
261
262    #[test]
263    #[should_panic(expected = "operation ID must not be empty")]
264    fn operation_id_rejects_empty_static_values() {
265        let _ = OperationId::from_static("");
266    }
267
268    #[test]
269    fn operation_id_preserves_application_static_controls() {
270        let operation = OperationId::from_static("create-item\nvariant");
271        assert_eq!(operation.as_str(), "create-item\nvariant");
272    }
273}