axum_observability/
context.rs1use std::{error::Error, fmt};
2
3use axum::{
4 http::StatusCode,
5 response::{IntoResponse, Response},
6};
7
8use crate::{RequestId, TraceContextLevel};
9
10#[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 #[must_use]
49 pub fn trace_id(&self) -> &str {
50 &self.trace_id
51 }
52
53 #[must_use]
55 pub fn parent_id(&self) -> &str {
56 &self.parent_id
57 }
58
59 #[must_use]
61 pub const fn flags(&self) -> u8 {
62 self.flags
63 }
64
65 #[must_use]
67 pub const fn sampled(&self) -> bool {
68 self.flags & 1 == 1
69 }
70
71 #[must_use]
73 pub const fn trace_context_level(&self) -> TraceContextLevel {
74 self.level
75 }
76
77 #[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 #[must_use]
91 pub fn traceparent_bytes(&self) -> &[u8] {
92 &self.traceparent
93 }
94
95 #[must_use]
97 pub fn traceparent(&self) -> Option<&str> {
98 std::str::from_utf8(&self.traceparent).ok()
99 }
100
101 #[must_use]
103 pub fn tracestate(&self) -> Option<&str> {
104 self.tracestate.as_deref()
105 }
106}
107
108#[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#[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 #[must_use]
183 pub const fn request_id(&self) -> &RequestId {
184 &self.request_id
185 }
186
187 #[must_use]
189 pub fn correlation_id(&self) -> &str {
190 &self.correlation_id
191 }
192
193 #[must_use]
195 pub const fn trace_context(&self) -> Option<&TraceContext> {
196 self.trace_context.as_ref()
197 }
198}
199
200#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
203pub struct OperationId(&'static str);
204
205impl OperationId {
206 #[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 #[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}