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#[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 #[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 #[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 #[must_use]
120 pub const fn trace_context_level(&self) -> TraceContextLevel {
121 self.trace_context_level
122 }
123
124 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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#[derive(Clone, Debug)]
270#[must_use]
271pub struct ObservabilityLayer {
272 config: ObservabilityConfig,
273}
274
275impl ObservabilityLayer {
276 #[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#[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}