Skip to main content

axum_observability/
middleware.rs

1use std::{
2    collections::BTreeMap,
3    fmt,
4    future::Future,
5    panic::{AssertUnwindSafe, catch_unwind},
6    pin::Pin,
7    sync::Arc,
8    task::{Context, Poll},
9    time::{Duration, Instant},
10};
11
12#[cfg(feature = "peer-ip")]
13use axum::extract::ConnectInfo;
14use axum::{
15    body::{Body, Bytes},
16    extract::MatchedPath,
17    http::{HeaderMap, HeaderName, HeaderValue, Request, Response, StatusCode, header::USER_AGENT},
18};
19use http_body::{Body as HttpBody, Frame, SizeHint};
20use pin_project_lite::pin_project;
21use serde::Serialize;
22use serde_json::Value;
23#[cfg(feature = "peer-ip")]
24use std::net::SocketAddr;
25use tower_layer::Layer;
26use tower_service::Service;
27use tracing::{Instrument, Level, Span};
28use uuid::Uuid;
29
30use crate::{
31    FieldConvention, JsonLayer, OperationId, RequestContext, RequestId, TraceContext,
32    TraceContextLevel,
33    request_id::native_field_content,
34    trace_context::{parse_traceparent_with_level, parse_tracestate_with_level},
35};
36
37type Generator = Arc<dyn Fn() -> Option<RequestId> + Send + Sync>;
38type Validator = Arc<dyn Fn(&str) -> bool + Send + Sync>;
39type LevelMapper = Arc<dyn Fn(StatusCode) -> Level + Send + Sync>;
40type Clock = Arc<dyn Fn() -> Instant + Send + Sync>;
41type Enricher = Arc<dyn Fn(&RequestContext) -> BTreeMap<String, Value> + Send + Sync>;
42
43/// Configuration for [`ObservabilityLayer`].
44#[derive(Clone)]
45#[must_use]
46#[allow(
47    clippy::struct_excessive_bools,
48    reason = "independent opt-in capture and response policies are explicit configuration"
49)]
50pub struct ObservabilityConfig {
51    pub(crate) field_convention: FieldConvention,
52    trace_context_level: TraceContextLevel,
53    request_id_header: HeaderName,
54    response_header: bool,
55    raw_path: bool,
56    #[cfg(feature = "peer-ip")]
57    peer_ip: bool,
58    user_agent: bool,
59    generator: Generator,
60    validator: Option<Validator>,
61    level_mapper: LevelMapper,
62    clock: Clock,
63    enricher: Enricher,
64}
65
66impl fmt::Debug for ObservabilityConfig {
67    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
68        let mut debug = formatter.debug_struct("ObservabilityConfig");
69        debug
70            .field("field_convention", &self.field_convention)
71            .field("trace_context_level", &self.trace_context_level)
72            .field("request_id_header", &self.request_id_header)
73            .field("response_header", &self.response_header)
74            .field("raw_path", &self.raw_path);
75        #[cfg(feature = "peer-ip")]
76        debug.field("peer_ip", &self.peer_ip);
77        debug
78            .field("user_agent", &self.user_agent)
79            .finish_non_exhaustive()
80    }
81}
82
83impl Default for ObservabilityConfig {
84    fn default() -> Self {
85        Self {
86            field_convention: FieldConvention::Generic,
87            trace_context_level: TraceContextLevel::Level1,
88            request_id_header: HeaderName::from_static("x-request-id"),
89            response_header: true,
90            raw_path: false,
91            #[cfg(feature = "peer-ip")]
92            peer_ip: false,
93            user_agent: false,
94            generator: Arc::new(|| Some(random_request_id())),
95            validator: None,
96            level_mapper: Arc::new(default_level),
97            clock: Arc::new(Instant::now),
98            enricher: Arc::new(|_| BTreeMap::new()),
99        }
100    }
101}
102
103impl ObservabilityConfig {
104    /// Selects the provider field convention.
105    #[must_use = "configuration builders return a new value"]
106    pub fn with_field_convention(mut self, convention: FieldConvention) -> Self {
107        self.field_convention = convention;
108        self
109    }
110
111    /// Selects the W3C Trace Context level used for inbound requests.
112    #[must_use = "configuration builders return a new value"]
113    pub const fn with_trace_context_level(mut self, level: TraceContextLevel) -> Self {
114        self.trace_context_level = level;
115        self
116    }
117
118    /// Returns the resolved W3C Trace Context level.
119    #[must_use]
120    pub const fn trace_context_level(&self) -> TraceContextLevel {
121        self.trace_context_level
122    }
123
124    /// Sets the request and response correlation header name.
125    ///
126    /// Use [`HeaderName::from_static`] for a known lowercase name, or
127    /// [`HeaderName::try_from`] when configuration supplies the value:
128    ///
129    /// ```
130    /// use axum::http::HeaderName;
131    /// use axum_observability::ObservabilityConfig;
132    ///
133    /// let static_config = ObservabilityConfig::default()
134    ///     .with_request_id_header(HeaderName::from_static("x-correlation-id"));
135    /// let dynamic_name = HeaderName::try_from("x-runtime-correlation-id")?;
136    /// let dynamic_config = ObservabilityConfig::default()
137    ///     .with_request_id_header(dynamic_name);
138    /// # let _ = (static_config, dynamic_config);
139    /// # Ok::<(), axum::http::header::InvalidHeaderName>(())
140    /// ```
141    #[must_use = "configuration builders return a new value"]
142    pub fn with_request_id_header(mut self, name: HeaderName) -> Self {
143        self.request_id_header = name;
144        self
145    }
146
147    /// Enables or disables adding the request ID response header.
148    #[must_use = "configuration builders return a new value"]
149    pub fn with_response_header(mut self, enabled: bool) -> Self {
150        self.response_header = enabled;
151        self
152    }
153
154    /// Enables or disables exact query-free raw path capture.
155    ///
156    /// Every nonempty path component exposed by
157    /// [`Uri::path`](axum::http::Uri::path) is retained without applying a
158    /// second path grammar.
159    ///
160    /// Enabling this can record identifying data and changes the application's
161    /// privacy posture. Query strings are never captured.
162    #[must_use = "configuration builders return a new value"]
163    pub fn with_raw_path(mut self, enabled: bool) -> Self {
164        self.raw_path = enabled;
165        self
166    }
167
168    /// Enables or disables capture of Axum's trusted socket peer extension.
169    ///
170    /// Enabling this can record identifying data and changes the application's
171    /// privacy posture. Forwarding headers are never inspected.
172    #[cfg(feature = "peer-ip")]
173    #[must_use = "configuration builders return a new value"]
174    pub fn with_peer_ip(mut self, enabled: bool) -> Self {
175        self.peer_ip = enabled;
176        self
177    }
178
179    /// Enables or disables capture of one unambiguous text User-Agent value.
180    ///
181    /// Enabling this can record identifying data and changes the application's
182    /// privacy posture.
183    #[must_use = "configuration builders return a new value"]
184    pub fn with_user_agent(mut self, enabled: bool) -> Self {
185        self.user_agent = enabled;
186        self
187    }
188
189    /// Sets a fallible request ID generator. It is invoked once per replacement
190    /// request before the crate falls back to a package-owned random identifier.
191    #[must_use = "configuration builders return a new value"]
192    pub fn with_request_id_generator(
193        mut self,
194        generator: impl Fn() -> Option<RequestId> + Send + Sync + 'static,
195    ) -> Self {
196        self.generator = Arc::new(generator);
197        self
198    }
199
200    /// Adds an application validator for one runtime-valid caller header.
201    /// Generated identifiers always retain the package's default grammar.
202    #[must_use = "configuration builders return a new value"]
203    pub fn with_request_id_validator(
204        mut self,
205        validator: impl Fn(&str) -> bool + Send + Sync + 'static,
206    ) -> Self {
207        self.validator = Some(Arc::new(validator));
208        self
209    }
210
211    /// Sets the mapping from final response status to access-event level.
212    #[must_use = "configuration builders return a new value"]
213    pub fn with_status_level_mapper(
214        mut self,
215        mapper: impl Fn(StatusCode) -> Level + Send + Sync + 'static,
216    ) -> Self {
217        self.level_mapper = Arc::new(mapper);
218        self
219    }
220
221    /// Sets a monotonic clock seam, primarily for deterministic testing.
222    ///
223    /// Clock panics are contained when the application uses Rust's default
224    /// `panic = "unwind"` behavior. Rust code cannot recover from
225    /// `panic = "abort"`.
226    #[must_use = "configuration builders return a new value"]
227    pub fn with_clock(mut self, clock: impl Fn() -> Instant + Send + Sync + 'static) -> Self {
228        self.clock = Arc::new(clock);
229        self
230    }
231
232    /// Adds controlled fields to terminal access records. Reserved package
233    /// fields cannot be overwritten.
234    #[must_use = "configuration builders return a new value"]
235    pub fn with_access_enricher(
236        mut self,
237        enricher: impl Fn(&RequestContext) -> BTreeMap<String, Value> + Send + Sync + 'static,
238    ) -> Self {
239        self.enricher = Arc::new(enricher);
240        self
241    }
242
243    /// Creates a composable JSON layer using this configuration's field convention.
244    #[must_use = "configuration builders return a new value"]
245    pub fn json_layer<W>(&self, writer: W) -> JsonLayer<W> {
246        JsonLayer::from_convention(writer, self.field_convention)
247    }
248
249    fn accepts_request_id(&self, value: &str) -> bool {
250        self.validator.as_ref().map_or_else(
251            || RequestId::parse(value).is_ok(),
252            |validator| catch_unwind(AssertUnwindSafe(|| validator(value))).unwrap_or(false),
253        )
254    }
255
256    fn generate_request_id(&self) -> RequestId {
257        if let Some(value) = catch_unwind(AssertUnwindSafe(|| (self.generator)()))
258            .ok()
259            .flatten()
260        {
261            return value;
262        }
263
264        random_request_id()
265    }
266}
267
268/// Cloneable Tower layer that installs correlation and terminal access logs.
269#[derive(Clone, Debug)]
270#[must_use]
271pub struct ObservabilityLayer {
272    config: ObservabilityConfig,
273}
274
275impl ObservabilityLayer {
276    /// Creates a layer from an explicit configuration.
277    #[must_use = "configuration builders return a new value"]
278    pub const fn new(config: ObservabilityConfig) -> Self {
279        Self { config }
280    }
281}
282
283impl Default for ObservabilityLayer {
284    fn default() -> Self {
285        Self::new(ObservabilityConfig::default())
286    }
287}
288
289impl<S> Layer<S> for ObservabilityLayer {
290    type Service = ObservabilityService<S>;
291
292    fn layer(&self, inner: S) -> Self::Service {
293        ObservabilityService {
294            inner,
295            config: self.config.clone(),
296        }
297    }
298}
299
300/// Service produced by [`ObservabilityLayer`].
301#[derive(Clone, Debug)]
302pub struct ObservabilityService<S> {
303    inner: S,
304    config: ObservabilityConfig,
305}
306
307impl<S> Service<Request<Body>> for ObservabilityService<S>
308where
309    S: Service<Request<Body>, Response = Response<Body>> + Send + 'static,
310    S::Future: Send + 'static,
311    S::Error: Send + 'static,
312{
313    type Response = Response<Body>;
314    type Error = S::Error;
315    type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
316
317    fn poll_ready(&mut self, context: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
318        self.inner.poll_ready(context)
319    }
320
321    fn call(&mut self, mut request: Request<Body>) -> Self::Future {
322        let request_id = select_request_id(request.headers_mut(), &self.config);
323        let trace_context =
324            select_trace_context(request.headers(), self.config.trace_context_level);
325        let request_context = RequestContext::new(request_id, trace_context);
326        let metadata = RequestMetadata::from_request(&request, &self.config);
327        let span = request_span(&request_context);
328        let started = catch_unwind(AssertUnwindSafe(|| (self.config.clock)()))
329            .unwrap_or_else(|_| Instant::now());
330        request.extensions_mut().insert(request_context.clone());
331        let future = {
332            let _entered = span.enter();
333            self.inner.call(request)
334        };
335        let config = self.config.clone();
336        let guard_span = span.clone();
337        let guard = TerminalGuard::new(
338            metadata,
339            request_context,
340            started,
341            config.clone(),
342            guard_span,
343        );
344
345        Box::pin(
346            async move {
347                let mut guard = guard;
348                match future.await {
349                    Ok(response) => {
350                        let (mut parts, body) = response.into_parts();
351                        guard.set_status(parts.status);
352                        if let Some(operation_id) = parts.extensions.get::<OperationId>() {
353                            guard.set_operation_id(operation_id);
354                        }
355                        if config.response_header {
356                            let value = HeaderValue::from_str(guard.request_id().as_str())
357                                .expect("validated request ID is always a header value");
358                            parts
359                                .headers
360                                .insert(config.request_id_header.clone(), value);
361                        }
362                        Ok(Response::from_parts(
363                            parts,
364                            Body::new(ObservedBody::new(body, guard)),
365                        ))
366                    }
367                    Err(error) => {
368                        guard.finish(TerminalOutcome::ServiceError);
369                        Err(error)
370                    }
371                }
372            }
373            .instrument(span),
374        )
375    }
376}
377
378fn select_request_id(headers: &mut HeaderMap, config: &ObservabilityConfig) -> RequestId {
379    let mut values = headers.get_all(&config.request_id_header).iter();
380    let first = values
381        .next()
382        .and_then(|value| std::str::from_utf8(value.as_bytes()).ok());
383    let request_id = if values.next().is_none() {
384        first
385            .filter(|value| native_field_content(value))
386            .filter(|value| config.accepts_request_id(value))
387            .and_then(RequestId::from_native_header)
388            .unwrap_or_else(|| config.generate_request_id())
389    } else {
390        config.generate_request_id()
391    };
392    let header_value = HeaderValue::from_str(request_id.as_str())
393        .expect("validated request ID is always a header value");
394    headers.insert(config.request_id_header.clone(), header_value);
395    request_id
396}
397
398fn random_request_id() -> RequestId {
399    RequestId::parse(&Uuid::new_v4().simple().to_string())
400        .expect("UUID simple formatting satisfies the request-ID contract")
401}
402
403fn select_trace_context(headers: &HeaderMap, level: TraceContextLevel) -> Option<TraceContext> {
404    let mut parents = headers.get_all("traceparent").iter();
405    let first = parents.next()?.as_bytes();
406    if parents.next().is_some() {
407        return None;
408    }
409    let trace = parse_traceparent_with_level(first, level)?;
410
411    let states = headers
412        .get_all("tracestate")
413        .iter()
414        .map(HeaderValue::to_str)
415        .collect::<Result<Vec<_>, _>>();
416    let tracestate = states
417        .ok()
418        .and_then(|values| parse_tracestate_with_level(values, level));
419    Some(trace.with_tracestate(tracestate))
420}
421
422fn request_span(context: &RequestContext) -> Span {
423    if let Some(trace) = context.trace_context() {
424        let flags = format!("{:02x}", trace.flags());
425        if let Some(random) = trace.trace_id_random() {
426            tracing::info_span!(
427                target: "axum_observability::request",
428                "request",
429                request_id = %context.request_id(),
430                correlation_id = context.correlation_id(),
431                trace_id = trace.trace_id(),
432                parent_id = trace.parent_id(),
433                trace_flags = flags.as_str(),
434                trace_sampled = trace.sampled(),
435                trace_id_random = random,
436            )
437        } else {
438            tracing::info_span!(
439                target: "axum_observability::request",
440                "request",
441                request_id = %context.request_id(),
442                correlation_id = context.correlation_id(),
443                trace_id = trace.trace_id(),
444                parent_id = trace.parent_id(),
445                trace_flags = flags.as_str(),
446                trace_sampled = trace.sampled(),
447                trace_id_random = tracing::field::Empty,
448            )
449        }
450    } else {
451        tracing::info_span!(
452            target: "axum_observability::request",
453            "request",
454            request_id = %context.request_id(),
455            correlation_id = context.correlation_id(),
456            trace_id = tracing::field::Empty,
457            parent_id = tracing::field::Empty,
458            trace_flags = tracing::field::Empty,
459            trace_sampled = tracing::field::Empty,
460            trace_id_random = tracing::field::Empty,
461        )
462    }
463}
464
465fn default_level(status: StatusCode) -> Level {
466    if status.is_server_error() {
467        Level::ERROR
468    } else if status.is_client_error() {
469        Level::WARN
470    } else {
471        Level::INFO
472    }
473}
474
475#[derive(Debug)]
476struct RequestMetadata {
477    method: String,
478    path: Option<String>,
479    path_template: Option<String>,
480    operation_id: Option<String>,
481    peer_ip: Option<String>,
482    user_agent: Option<String>,
483}
484
485impl RequestMetadata {
486    fn from_request(request: &Request<Body>, config: &ObservabilityConfig) -> Self {
487        Self {
488            method: request.method().to_string(),
489            path: config
490                .raw_path
491                .then(|| canonical_raw_path(request.uri().path()))
492                .flatten(),
493            path_template: request
494                .extensions()
495                .get::<MatchedPath>()
496                .and_then(|path| canonical_route_template(path.as_str())),
497            operation_id: request
498                .extensions()
499                .get::<OperationId>()
500                .map(|operation| operation.as_str().to_owned()),
501            peer_ip: peer_ip(request, config),
502            user_agent: config
503                .user_agent
504                .then(|| exactly_one_header(request.headers(), &USER_AGENT))
505                .flatten(),
506        }
507    }
508}
509
510fn canonical_raw_path(value: &str) -> Option<String> {
511    (!value.is_empty()).then(|| value.to_owned())
512}
513
514fn canonical_route_template(native: &str) -> Option<String> {
515    (!native.is_empty()).then(|| native.to_owned())
516}
517
518#[cfg(test)]
519mod route_template_tests {
520    use super::{canonical_raw_path, canonical_route_template};
521
522    #[test]
523    fn canonical_raw_path_preserves_every_nonempty_native_path_component() {
524        for (raw, expected) in [
525            ("/", Some("/")),
526            ("/objects/a%2Fb/%E2%9C%93", Some("/objects/a%2Fb/%E2%9C%93")),
527            ("objects/no-leading-slash", Some("objects/no-leading-slash")),
528            ("*", Some("*")),
529            ("/objects/bad%2", Some("/objects/bad%2")),
530            ("/objects/bad%GG", Some("/objects/bad%GG")),
531            ("/objects/bad%2G", Some("/objects/bad%2G")),
532            ("/objects/bad%G2", Some("/objects/bad%G2")),
533            ("/a%20%G2", Some("/a%20%G2")),
534            ("", None),
535        ] {
536            assert_eq!(canonical_raw_path(raw).as_deref(), expected, "{raw}");
537        }
538    }
539
540    #[test]
541    fn canonical_route_template_preserves_nonempty_authoritative_matched_paths() {
542        for (native, expected) in [
543            ("/health".to_owned(), Some("/health".to_owned())),
544            (
545                "/items/{item_id}".to_owned(),
546                Some("/items/{item_id}".to_owned()),
547            ),
548            (
549                "/files/{*path}".to_owned(),
550                Some("/files/{*path}".to_owned()),
551            ),
552            (
553                format!("/items/{{{}}}", "a".repeat(64)),
554                Some(format!("/items/{{{}}}", "a".repeat(64))),
555            ),
556            (
557                format!("/items/{{{}}}", "a".repeat(65)),
558                Some(format!("/items/{{{}}}", "a".repeat(65))),
559            ),
560            (
561                "/items/{0item}".to_owned(),
562                Some("/items/{0item}".to_owned()),
563            ),
564            (
565                "/items/{item-id}".to_owned(),
566                Some("/items/{item-id}".to_owned()),
567            ),
568            ("/literal*star".to_owned(), Some("/literal*star".to_owned())),
569            (String::new(), None),
570        ] {
571            assert_eq!(canonical_route_template(&native), expected, "{native}");
572        }
573    }
574}
575
576#[cfg(feature = "peer-ip")]
577fn peer_ip(request: &Request<Body>, config: &ObservabilityConfig) -> Option<String> {
578    if !config.peer_ip {
579        return None;
580    }
581    request
582        .extensions()
583        .get::<ConnectInfo<SocketAddr>>()
584        .map(|connect| connect.0.ip().to_string())
585}
586
587#[cfg(not(feature = "peer-ip"))]
588fn unavailable_peer_ip(_request: &Request<Body>, _config: &ObservabilityConfig) -> Option<String> {
589    None
590}
591
592#[cfg(not(feature = "peer-ip"))]
593use unavailable_peer_ip as peer_ip;
594
595fn exactly_one_header(headers: &HeaderMap, name: &HeaderName) -> Option<String> {
596    let mut values = headers.get_all(name).iter();
597    let first = std::str::from_utf8(values.next()?.as_bytes()).ok()?;
598    if values.next().is_some() || !native_field_content(first) {
599        return None;
600    }
601    Some(first.to_owned())
602}
603
604#[derive(Debug, Serialize)]
605pub(crate) struct AccessRecord {
606    request_id: String,
607    correlation_id: String,
608    #[serde(skip_serializing_if = "Option::is_none")]
609    trace_id: Option<String>,
610    #[serde(skip_serializing_if = "Option::is_none")]
611    parent_id: Option<String>,
612    #[serde(skip_serializing_if = "Option::is_none")]
613    trace_flags: Option<String>,
614    #[serde(skip_serializing_if = "Option::is_none")]
615    trace_sampled: Option<bool>,
616    #[serde(skip_serializing_if = "Option::is_none")]
617    trace_id_random: Option<bool>,
618    method: String,
619    #[serde(skip_serializing_if = "Option::is_none")]
620    path: Option<String>,
621    #[serde(skip_serializing_if = "Option::is_none")]
622    path_template: Option<String>,
623    #[serde(skip_serializing_if = "Option::is_none")]
624    operation_id: Option<String>,
625    #[serde(skip_serializing_if = "Option::is_none")]
626    status: Option<u16>,
627    duration_ms: DurationMilliseconds,
628    #[serde(skip_serializing_if = "Option::is_none")]
629    peer_ip: Option<String>,
630    #[serde(skip_serializing_if = "Option::is_none")]
631    user_agent: Option<String>,
632    #[serde(skip_serializing_if = "Option::is_none")]
633    terminal_reason: Option<String>,
634    enrichment: BTreeMap<String, Value>,
635}
636
637struct TerminalState {
638    metadata: RequestMetadata,
639    request_context: RequestContext,
640    started: Instant,
641    status: Option<StatusCode>,
642    config: ObservabilityConfig,
643    span: Span,
644}
645
646struct TerminalGuard {
647    state: Option<TerminalState>,
648}
649
650#[derive(Clone, Copy, Debug, Eq, PartialEq)]
651enum TerminalOutcome {
652    Completed,
653    ServiceError,
654    BodyError,
655    ResponseDropped,
656}
657
658impl TerminalOutcome {
659    const fn terminal_reason(self) -> Option<&'static str> {
660        match self {
661            Self::Completed => None,
662            Self::ServiceError => Some("service_error"),
663            Self::BodyError => Some("body_error"),
664            Self::ResponseDropped => Some("response_dropped"),
665        }
666    }
667}
668
669impl TerminalGuard {
670    fn new(
671        metadata: RequestMetadata,
672        request_context: RequestContext,
673        started: Instant,
674        config: ObservabilityConfig,
675        span: Span,
676    ) -> Self {
677        Self {
678            state: Some(TerminalState {
679                metadata,
680                request_context,
681                started,
682                status: None,
683                config,
684                span,
685            }),
686        }
687    }
688
689    fn request_id(&self) -> &RequestId {
690        self.state
691            .as_ref()
692            .expect("terminal guard has not completed")
693            .request_context
694            .request_id()
695    }
696
697    fn set_status(&mut self, status: StatusCode) {
698        if let Some(state) = &mut self.state {
699            state.status = Some(status);
700        }
701    }
702
703    fn set_operation_id(&mut self, operation_id: &OperationId) {
704        if let Some(state) = &mut self.state {
705            state.metadata.operation_id = Some(operation_id.as_str().to_owned());
706        }
707    }
708
709    fn finish(&mut self, outcome: TerminalOutcome) {
710        let Some(state) = self.state.take() else {
711            return;
712        };
713        let mapped_level = |status| {
714            catch_unwind(AssertUnwindSafe(|| (state.config.level_mapper)(status)))
715                .unwrap_or_else(|_| default_level(status))
716        };
717        let level = match outcome {
718            TerminalOutcome::Completed => {
719                mapped_level(state.status.expect("completed response has a status"))
720            }
721            TerminalOutcome::ServiceError
722            | TerminalOutcome::BodyError
723            | TerminalOutcome::ResponseDropped => Level::ERROR,
724        };
725        let finished =
726            catch_unwind(AssertUnwindSafe(|| (state.config.clock)())).unwrap_or(state.started);
727        let duration = finished.saturating_duration_since(state.started);
728        let trace = state.request_context.trace_context();
729        let terminal_reason = outcome.terminal_reason();
730        let enrichment = catch_unwind(AssertUnwindSafe(|| {
731            (state.config.enricher)(&state.request_context)
732        }))
733        .unwrap_or_default();
734        let record = AccessRecord {
735            request_id: state.request_context.request_id().as_str().to_owned(),
736            correlation_id: state.request_context.correlation_id().to_owned(),
737            trace_id: trace.map(|trace| trace.trace_id().to_owned()),
738            parent_id: trace.map(|trace| trace.parent_id().to_owned()),
739            trace_flags: trace.map(|trace| format!("{:02x}", trace.flags())),
740            trace_sampled: trace.map(TraceContext::sampled),
741            trace_id_random: trace.and_then(TraceContext::trace_id_random),
742            method: state.metadata.method,
743            path: state.metadata.path,
744            path_template: state.metadata.path_template,
745            operation_id: state.metadata.operation_id,
746            status: state.status.map(|status| status.as_u16()),
747            duration_ms: duration_milliseconds(duration),
748            peer_ip: state.metadata.peer_ip,
749            user_agent: state.metadata.user_agent,
750            terminal_reason: terminal_reason.map(str::to_owned),
751            enrichment,
752        };
753        let serialized = serde_json::to_string(&record).expect("access record is serializable");
754        state.span.in_scope(|| emit_access(level, &serialized));
755    }
756}
757
758#[derive(Debug, Serialize)]
759#[serde(untagged)]
760enum DurationMilliseconds {
761    Integer(u128),
762    Fractional(f64),
763}
764
765fn duration_milliseconds(duration: Duration) -> DurationMilliseconds {
766    if duration.subsec_nanos().is_multiple_of(1_000_000) {
767        DurationMilliseconds::Integer(duration.as_millis())
768    } else {
769        DurationMilliseconds::Fractional(duration.as_secs_f64() * 1_000.0)
770    }
771}
772
773#[cfg(test)]
774mod duration_tests {
775    use std::time::Duration;
776
777    use serde_json::json;
778
779    use super::duration_milliseconds;
780
781    #[test]
782    fn portable_duration_is_not_clamped_by_a_provider_projection_range() {
783        let maximum = serde_json::to_value(duration_milliseconds(Duration::from_hours(87_660_000)))
784            .expect("serialize maximum duration");
785        let overflow =
786            serde_json::to_value(duration_milliseconds(Duration::from_secs(315_576_000_001)))
787                .expect("serialize overflow duration");
788        assert_eq!(maximum, json!(315_576_000_000_000_u64));
789        assert_eq!(overflow, json!(315_576_000_001_000_u64));
790    }
791}
792
793impl Drop for TerminalGuard {
794    fn drop(&mut self) {
795        self.finish(TerminalOutcome::ResponseDropped);
796    }
797}
798
799fn emit_access(level: Level, serialized: &str) {
800    match level {
801        Level::ERROR => tracing::event!(
802            target: "axum_observability::access",
803            Level::ERROR,
804            message = "request completed",
805            "obs.record" = serialized
806        ),
807        Level::WARN => tracing::event!(
808            target: "axum_observability::access",
809            Level::WARN,
810            message = "request completed",
811            "obs.record" = serialized
812        ),
813        Level::INFO => tracing::event!(
814            target: "axum_observability::access",
815            Level::INFO,
816            message = "request completed",
817            "obs.record" = serialized
818        ),
819        Level::DEBUG => tracing::event!(
820            target: "axum_observability::access",
821            Level::DEBUG,
822            message = "request completed",
823            "obs.record" = serialized
824        ),
825        Level::TRACE => tracing::event!(
826            target: "axum_observability::access",
827            Level::TRACE,
828            message = "request completed",
829            "obs.record" = serialized
830        ),
831    }
832}
833
834pin_project! {
835    struct ObservedBody {
836        #[pin]
837        body: Body,
838        guard: TerminalGuard,
839    }
840}
841
842impl ObservedBody {
843    fn new(body: Body, mut guard: TerminalGuard) -> Self {
844        if body.is_end_stream() {
845            guard.finish(TerminalOutcome::Completed);
846        }
847        Self { body, guard }
848    }
849}
850
851impl HttpBody for ObservedBody {
852    type Data = Bytes;
853    type Error = axum::Error;
854
855    fn poll_frame(
856        self: Pin<&mut Self>,
857        context: &mut Context<'_>,
858    ) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
859        let mut this = self.project();
860        match this.body.as_mut().poll_frame(context) {
861            Poll::Ready(None) => {
862                this.guard.finish(TerminalOutcome::Completed);
863                Poll::Ready(None)
864            }
865            Poll::Ready(Some(Err(error))) => {
866                this.guard.finish(TerminalOutcome::BodyError);
867                Poll::Ready(Some(Err(error)))
868            }
869            Poll::Ready(Some(Ok(frame))) => {
870                if this.body.is_end_stream() {
871                    this.guard.finish(TerminalOutcome::Completed);
872                }
873                Poll::Ready(Some(Ok(frame)))
874            }
875            Poll::Pending => Poll::Pending,
876        }
877    }
878
879    fn is_end_stream(&self) -> bool {
880        self.body.is_end_stream()
881    }
882
883    fn size_hint(&self) -> SizeHint {
884        self.body.size_hint()
885    }
886}