rama_http/layer/trace/mod.rs
1//! Middleware that adds high level [tracing] to a [`Service`].
2//!
3//! # Example
4//!
5//! Adding tracing to your service can be as simple as:
6//!
7//! ```rust
8//! use rama_http::{Body, Request, Response};
9//! use rama_core::service::service_fn;
10//! use rama_core::{Layer, Service};
11//! use rama_http::layer::trace::TraceLayer;
12//! use std::convert::Infallible;
13//!
14//! async fn handle(request: Request) -> Result<Response, Infallible> {
15//! Ok(Response::new(Body::from("foo")))
16//! }
17//!
18//! # #[tokio::main]
19//! # async fn main() -> Result<(), Box<dyn std::error::Error>> {
20//! // Setup tracing
21//! tracing_subscriber::fmt::init();
22//!
23//! let mut service = TraceLayer::new_for_http().into_layer(service_fn(handle));
24//!
25//! let request = Request::new(Body::from("foo"));
26//!
27//! let response = service
28//! .serve(request)
29//! .await?;
30//! # Ok(())
31//! # }
32//! ```
33//!
34//! If you run this application with `RUST_LOG=rama=trace cargo run` you should see logs like:
35//!
36//! ```text
37//! Mar 05 20:50:28.523 DEBUG request{method=GET path="/foo"}: rama_http::layer::trace::on_request: started processing request
38//! Mar 05 20:50:28.524 DEBUG request{method=GET path="/foo"}: rama_http::layer::trace::on_response: finished processing request latency=1 ms status=200
39//! ```
40//!
41//! # Customization
42//!
43//! [`Trace`] comes with good defaults but also supports customizing many aspects of the output.
44//!
45//! The default behaviour supports some customization:
46//!
47//! ```rust
48//! use rama_http::{Body, Request, Response, HeaderMap, StatusCode};
49//! use rama_core::service::service_fn;
50//! use rama_core::{Service, Layer};
51//! use rama_core::telemetry::tracing::Level;
52//! use rama_http::layer::trace::{
53//! TraceLayer, DefaultMakeSpan, DefaultOnRequest, DefaultOnResponse,
54//! };
55//! use rama_utils::latency::LatencyUnit;
56//! use std::time::Duration;
57//! use std::convert::Infallible;
58//!
59//! # async fn handle(request: Request) -> Result<Response, Infallible> {
60//! # Ok(Response::new(Body::from("foo")))
61//! # }
62//! # #[tokio::main]
63//! # async fn main() -> Result<(), Box<dyn std::error::Error>> {
64//! # tracing_subscriber::fmt::init();
65//! #
66//! let service = (
67//! TraceLayer::new_for_http()
68//! .make_span_with(
69//! DefaultMakeSpan::new().with_include_headers(true)
70//! )
71//! .on_request(
72//! DefaultOnRequest::new().with_level(Level::INFO)
73//! )
74//! .on_response(
75//! DefaultOnResponse::new()
76//! .with_level(Level::INFO)
77//! .with_latency_unit(LatencyUnit::Micros)
78//! ),
79//! // and so on for `on_eos`, `on_body_chunk`, and `on_failure`
80//! ).into_layer(service_fn(handle));
81//! # let mut service = service;
82//! # let response = service
83//! # .serve(Request::new(Body::from("foo")))
84//! # .await?;
85//! # Ok(())
86//! # }
87//! ```
88//!
89//! However for maximum control you can provide callbacks:
90//!
91//! ```rust
92//! use rama_http::{Body, Request, Response, HeaderMap, StatusCode};
93//! use rama_core::service::service_fn;
94//! use rama_core::{Service, Layer};
95//! use rama_core::telemetry::tracing::{self, Span};
96//! use rama_http::layer::{classify::ServerErrorsFailureClass, trace::TraceLayer};
97//! use std::time::Duration;
98//! use std::convert::Infallible;
99//! use rama_core::bytes::Bytes;
100//!
101//! # async fn handle(request: Request) -> Result<Response, Infallible> {
102//! # Ok(Response::new(Body::from("foo")))
103//! # }
104//! # #[tokio::main]
105//! # async fn main() -> Result<(), Box<dyn std::error::Error>> {
106//! # tracing_subscriber::fmt::init();
107//! #
108//! let service = (
109//! TraceLayer::new_for_http()
110//! .make_span_with(|request: &Request| {
111//! tracing::debug_span!("http-request")
112//! })
113//! .on_request(|request: &Request, _span: &Span| {
114//! tracing::debug!(
115//! http.method = %request.method(),
116//! url.path = %request.uri().path_or_root().as_ref(),
117//! "started request",
118//! )
119//! })
120//! .on_response(|response: &Response, latency: Duration, _span: &Span| {
121//! tracing::debug!("response generated in {:?}", latency)
122//! })
123//! .on_body_chunk(|chunk: &Bytes, latency: Duration, _span: &Span| {
124//! tracing::debug!("sending {} bytes", chunk.len())
125//! })
126//! .on_eos(|trailers: Option<&HeaderMap>, stream_duration: Duration, _span: &Span| {
127//! tracing::debug!("stream closed after {:?}", stream_duration)
128//! })
129//! .on_failure(|error: ServerErrorsFailureClass, latency: Duration, _span: &Span| {
130//! tracing::debug!("something went wrong")
131//! })
132//! ).into_layer(service_fn(handle));
133//! # let mut service = service;
134//! # let response = service
135//! # .serve(Request::new(Body::from("foo")))
136//! # .await?;
137//! # Ok(())
138//! # }
139//! ```
140//!
141//! ## Disabling something
142//!
143//! Setting the behaviour to `()` will be disable that particular step:
144//!
145//! ```rust
146//! use rama_http::{Body, Request, Response, StatusCode};
147//! use rama_core::service::service_fn;
148//! use rama_core::{Service, Layer};
149//! use rama_core::telemetry::tracing::{self, Span};
150//! use rama_http::layer::{classify::ServerErrorsFailureClass, trace::TraceLayer};
151//! use std::time::Duration;
152//! # use std::convert::Infallible;
153//!
154//! # async fn handle(request: Request) -> Result<Response, Infallible> {
155//! # Ok(Response::new(Body::from("foo")))
156//! # }
157//! # #[tokio::main]
158//! # async fn main() -> Result<(), Box<dyn std::error::Error>> {
159//! # tracing_subscriber::fmt::init();
160//! #
161//! let service = (
162//! // This configuration will only emit events on failures
163//! TraceLayer::new_for_http()
164//! .on_request(())
165//! .on_response(())
166//! .on_body_chunk(())
167//! .on_eos(())
168//! .on_failure(|error: ServerErrorsFailureClass, latency: Duration, _span: &Span| {
169//! tracing::debug!("something went wrong")
170//! })
171//! ).into_layer(service_fn(handle));
172//! # let mut service = service;
173//! # let response = service
174//! # .serve(Request::new(Body::from("foo")))
175//! # .await?;
176//! # Ok(())
177//! # }
178//! ```
179//!
180//! # When the callbacks are called
181//!
182//! ### `on_request`
183//!
184//! The `on_request` callback is called when the request arrives at the
185//! middleware in [`Service::serve`] just prior to passing the request to the
186//! inner service.
187//!
188//! ### `on_response`
189//!
190//! The `on_response` callback is called when the inner service's response
191//! future completes with `Ok(response)` regardless if the response is
192//! classified as a success or a failure.
193//!
194//! For example if you're using [`ServerErrorsAsFailures`] as your classifier
195//! and the inner service responds with `500 Internal Server Error` then the
196//! `on_response` callback is still called. `on_failure` would _also_ be called
197//! in this case since the response was classified as a failure.
198//!
199//! ### `on_body_chunk`
200//!
201//! The `on_body_chunk` callback is called when the response body produces a new
202//! chunk, that is when [`crate::StreamingBody::poll_frame`] returns `Poll::Ready(Some(Ok(chunk)))`.
203//!
204//! `on_body_chunk` is called even if the chunk is empty.
205//!
206//! ### `on_eos`
207//!
208//! The `on_eos` callback is called when a streaming response body ends, that is
209//! when `http_body::Body::poll_frame` returns `Poll::Ready(None)`.
210//!
211//! `on_eos` is called even if the trailers produced are `None`.
212//!
213//! ### `on_failure`
214//!
215//! The `on_failure` callback is called when:
216//!
217//! - The inner [`Service`]'s response future resolves to an error.
218//! - A response is classified as a failure.
219//! - [`crate::StreamingBody::poll_frame`] returns an error.
220//! - An end-of-stream is classified as a failure.
221//!
222//! # Recording fields on the span
223//!
224//! All callbacks receive a reference to the [tracing] [`Span`], corresponding to this request,
225//! produced by the closure passed to [`TraceLayer::make_span_with`]. It can be used to [record
226//! field values][record] that weren't known when the span was created.
227//!
228//! ```rust
229//! use rama_http::{Body, Request, Response, HeaderMap, StatusCode};
230//! use rama_core::service::service_fn;
231//! use rama_core::Layer;
232//! use rama_http::layer::trace::TraceLayer;
233//! use rama_core::telemetry::tracing::{self, Span};
234//! use std::time::Duration;
235//! use std::convert::Infallible;
236//!
237//! # async fn handle(request: Request) -> Result<Response, Infallible> {
238//! # Ok(Response::new(Body::from("foo")))
239//! # }
240//! # #[tokio::main]
241//! # async fn main() -> Result<(), Box<dyn std::error::Error>> {
242//! # tracing_subscriber::fmt::init();
243//! #
244//! let service = (
245//! TraceLayer::new_for_http()
246//! .make_span_with(|request: &Request| {
247//! tracing::debug_span!(
248//! "http-request",
249//! status_code = tracing::field::Empty,
250//! )
251//! })
252//! .on_response(|response: &Response, _latency: Duration, span: &Span| {
253//! span.record("status_code", &tracing::field::display(response.status()));
254//!
255//! tracing::debug!("response generated")
256//! }),
257//! ).into_layer(service_fn(handle));
258//! # Ok(())
259//! # }
260//! ```
261//!
262//! # Providing classifiers
263//!
264//! Tracing requires determining if a response is a success or failure. [`MakeClassifier`] is used
265//! to create a classifier for the incoming request. See the docs for [`MakeClassifier`] and
266//! [`ClassifyResponse`] for more details on classification.
267//!
268//! A [`MakeClassifier`] can be provided when creating a [`TraceLayer`]:
269//!
270//! ```rust
271//! use rama_http::{Body, Request, Response};
272//! use rama_core::service::service_fn;
273//! use rama_core::Layer;
274//! use rama_http::layer::{
275//! trace::TraceLayer,
276//! classify::{
277//! MakeClassifier, ClassifyResponse, ClassifiedResponse, NeverClassifyEos,
278//! SharedClassifier,
279//! },
280//! };
281//! use std::convert::Infallible;
282//!
283//! # async fn handle(request: Request) -> Result<Response, Infallible> {
284//! # Ok(Response::new(Body::from("foo")))
285//! # }
286//! # #[tokio::main]
287//! # async fn main() -> Result<(), Box<dyn std::error::Error>> {
288//! # tracing_subscriber::fmt::init();
289//! #
290//! // Our `MakeClassifier` that always crates `MyClassifier` classifiers.
291//! #[derive(Copy, Clone)]
292//! struct MyMakeClassify;
293//!
294//! impl MakeClassifier for MyMakeClassify {
295//! type Classifier = MyClassifier;
296//! type FailureClass = &'static str;
297//! type ClassifyEos = NeverClassifyEos<&'static str>;
298//!
299//! fn make_classifier<B>(&self, req: &Request<B>) -> Self::Classifier {
300//! MyClassifier
301//! }
302//! }
303//!
304//! // A classifier that classifies failures as `"something went wrong..."`.
305//! #[derive(Copy, Clone)]
306//! struct MyClassifier;
307//!
308//! impl ClassifyResponse for MyClassifier {
309//! type FailureClass = &'static str;
310//! type ClassifyEos = NeverClassifyEos<&'static str>;
311//!
312//! fn classify_response<B>(
313//! self,
314//! res: &Response<B>
315//! ) -> ClassifiedResponse<Self::FailureClass, Self::ClassifyEos> {
316//! // Classify based on the status code.
317//! if res.status().is_server_error() {
318//! ClassifiedResponse::Ready(Err("something went wrong..."))
319//! } else {
320//! ClassifiedResponse::Ready(Ok(()))
321//! }
322//! }
323//!
324//! fn classify_error<E>(self, error: &E) -> Self::FailureClass
325//! where
326//! E: std::fmt::Display,
327//! {
328//! "something went wrong..."
329//! }
330//! }
331//!
332//! let service = (
333//! // Create a trace layer that uses our classifier.
334//! TraceLayer::new(MyMakeClassify),
335//! ).into_layer(service_fn(handle));
336//!
337//! // Since `MyClassifier` is `Clone` we can also use `SharedClassifier`
338//! // to avoid having to define a separate `MakeClassifier`.
339//! let service = TraceLayer::new(SharedClassifier::new(MyClassifier)).into_layer(service_fn(handle));
340//! # Ok(())
341//! # }
342//! ```
343//!
344//! [`TraceLayer`] comes with convenience methods for using common classifiers:
345//!
346//! - [`TraceLayer::new_for_http`] classifies based on the status code. It doesn't consider
347//! streaming responses.
348//! - [`TraceLayer::new_for_grpc`] classifies based on the gRPC protocol and supports streaming
349//! responses.
350//!
351//! [tracing]: https://crates.io/crates/tracing
352//! [`Service`]: rama_core::Service
353//! [`Service::serve`]: rama_core::Service::serve
354//! [`MakeClassifier`]: crate::layer::classify::MakeClassifier
355//! [`ClassifyResponse`]: crate::layer::classify::ClassifyResponse
356//! [record]: https://docs.rs/tracing/latest/tracing/span/struct.Span.html#method.record
357//! [`TraceLayer::make_span_with`]: crate::layer::trace::TraceLayer::make_span_with
358//! [`Span`]: rama_core::telemetry::tracing::Span
359//! [`ServerErrorsAsFailures`]: crate::layer::classify::ServerErrorsAsFailures
360
361use rama_core::telemetry::tracing::Level;
362use std::{fmt, time::Duration};
363
364#[doc(inline)]
365pub use self::{
366 body::ResponseBody,
367 layer::TraceLayer,
368 make_span::{DefaultMakeSpan, MakeSpan},
369 on_body_chunk::{DefaultOnBodyChunk, OnBodyChunk},
370 on_eos::{DefaultOnEos, OnEos},
371 on_failure::{DefaultOnFailure, OnFailure},
372 on_request::{DefaultOnRequest, OnRequest},
373 on_response::{DefaultOnResponse, OnResponse},
374 service::Trace,
375};
376
377use crate::layer::classify::{GrpcErrorsAsFailures, ServerErrorsAsFailures, SharedClassifier};
378use rama_utils::latency::LatencyUnit;
379
380/// MakeClassifier for HTTP requests.
381pub type HttpMakeClassifier = SharedClassifier<ServerErrorsAsFailures>;
382
383/// MakeClassifier for gRPC requests.
384pub type GrpcMakeClassifier = SharedClassifier<GrpcErrorsAsFailures>;
385
386macro_rules! event_dynamic_lvl {
387 ( target: $target:expr, parent: $parent:expr, $lvl:expr, $($tt:tt)* ) => {
388 {
389 use ::rama_core::telemetry::tracing;
390 match $lvl {
391 tracing::Level::ERROR => {
392 tracing::event!(target: $target, parent: $parent, tracing::Level::ERROR, $($tt)*);
393 }
394 tracing::Level::WARN => {
395 tracing::event!(target: $target, parent: $parent, tracing::Level::WARN, $($tt)*);
396 }
397 tracing::Level::INFO => {
398 tracing::event!(target: $target, parent: $parent, tracing::Level::INFO, $($tt)*);
399 }
400 tracing::Level::DEBUG => {
401 tracing::event!(target: $target, parent: $parent, tracing::Level::INFO, $($tt)*);
402 }
403 tracing::Level::TRACE => {
404 tracing::event!(target: $target, parent: $parent, tracing::Level::TRACE, $($tt)*);
405 }
406 }
407 }
408 };
409 ( target: $target:expr, $lvl:expr, $($tt:tt)* ) => {
410 use ::rama_core::telemetry::tracing;
411 match $lvl {
412 tracing::Level::ERROR => {
413 tracing::event!(target: $target, tracing::Level::ERROR, $($tt)*);
414 }
415 tracing::Level::WARN => {
416 tracing::event!(target: $target, tracing::Level::WARN, $($tt)*);
417 }
418 tracing::Level::INFO => {
419 tracing::event!(target: $target, tracing::Level::INFO, $($tt)*);
420 }
421 tracing::Level::DEBUG => {
422 tracing::event!(target: $target, tracing::Level::DEBUG, $($tt)*);
423 }
424 tracing::Level::TRACE => {
425 tracing::event!(target: $target, tracing::Level::TRACE, $($tt)*);
426 }
427 }
428 };
429 ( parent: $parent:expr, $lvl:expr, $($tt:tt)* ) => {
430 use ::rama_core::telemetry::tracing;
431 match $lvl {
432 tracing::Level::ERROR => {
433 tracing::event!(parent: $parent, tracing::Level::ERROR, $($tt)*);
434 }
435 tracing::Level::WARN => {
436 tracing::event!(parent: $parent, tracing::Level::WARN, $($tt)*);
437 }
438 tracing::Level::INFO => {
439 tracing::event!(parent: $parent, tracing::Level::INFO, $($tt)*);
440 }
441 tracing::Level::DEBUG => {
442 tracing::event!(parent: $parent, tracing::Level::DEBUG, $($tt)*);
443 }
444 tracing::Level::TRACE => {
445 tracing::event!(parent: $parent, tracing::Level::TRACE, $($tt)*);
446 }
447 }
448 };
449 ( $lvl:expr, $($tt:tt)* ) => {
450 {
451 use ::rama_core::telemetry::tracing;
452 match $lvl {
453 tracing::Level::ERROR => {
454 tracing::event!(tracing::Level::ERROR, $($tt)*);
455 }
456 tracing::Level::WARN => {
457 tracing::event!(tracing::Level::WARN, $($tt)*);
458 }
459 tracing::Level::INFO => {
460 tracing::event!(tracing::Level::INFO, $($tt)*);
461 }
462 tracing::Level::DEBUG => {
463 tracing::event!(tracing::Level::DEBUG, $($tt)*);
464 }
465 tracing::Level::TRACE => {
466 tracing::event!(tracing::Level::TRACE, $($tt)*);
467 }
468 }
469 }
470 };
471}
472
473mod body;
474mod layer;
475mod make_span;
476mod on_body_chunk;
477mod on_eos;
478mod on_failure;
479mod on_request;
480mod on_response;
481mod service;
482
483const DEFAULT_MESSAGE_LEVEL: Level = Level::DEBUG;
484const DEFAULT_ERROR_LEVEL: Level = Level::ERROR;
485
486struct Latency {
487 unit: LatencyUnit,
488 duration: Duration,
489}
490
491impl fmt::Display for Latency {
492 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
493 match self.unit {
494 LatencyUnit::Seconds => write!(f, "{} s", self.duration.as_secs_f64()),
495 LatencyUnit::Millis => write!(f, "{} ms", self.duration.as_millis()),
496 LatencyUnit::Micros => write!(f, "{} μs", self.duration.as_micros()),
497 LatencyUnit::Nanos => write!(f, "{} ns", self.duration.as_nanos()),
498 }
499 }
500}
501
502#[cfg(test)]
503mod tests {
504 use super::*;
505
506 use crate::body::util::BodyExt as _;
507 use crate::layer::classify::{
508 ClassifiedResponse, ClassifyEos, ClassifyResponse, MakeClassifier, ServerErrorsFailureClass,
509 };
510 use crate::{Body, HeaderMap, Request, Response};
511 use rama_core::bytes::Bytes;
512 use rama_core::error::BoxError;
513 use rama_core::service::service_fn;
514 use rama_core::telemetry::tracing::{self, Span};
515 use rama_core::{Layer, Service};
516 use std::sync::OnceLock;
517 use std::{
518 sync::atomic::{AtomicU32, Ordering},
519 time::Duration,
520 };
521
522 macro_rules! lazy_atomic_u32 {
523 ($($name:ident),+) => {
524 $(
525 #[expect(non_snake_case)]
526 fn $name() -> &'static AtomicU32 {
527 static $name: OnceLock<AtomicU32> = OnceLock::new();
528 $name.get_or_init(|| AtomicU32::new(0))
529 }
530 )+
531 };
532 }
533
534 #[tokio::test]
535 async fn unary_request() {
536 lazy_atomic_u32!(
537 ON_REQUEST_COUNT,
538 ON_RESPONSE_COUNT,
539 ON_BODY_CHUNK_COUNT,
540 ON_EOS,
541 ON_FAILURE
542 );
543
544 let trace_layer = TraceLayer::new_for_http()
545 .make_span_with(|_req: &Request| {
546 tracing::info_span!("test-span", foo = tracing::field::Empty)
547 })
548 .on_request(|_req: &Request, span: &Span| {
549 span.record("foo", 42);
550 ON_REQUEST_COUNT().fetch_add(1, Ordering::AcqRel);
551 })
552 .on_response(|_res: &Response, _latency: Duration, _span: &Span| {
553 ON_RESPONSE_COUNT().fetch_add(1, Ordering::AcqRel);
554 })
555 .on_body_chunk(|_chunk: &Bytes, _latency: Duration, _span: &Span| {
556 ON_BODY_CHUNK_COUNT().fetch_add(1, Ordering::AcqRel);
557 })
558 .on_eos(
559 |_trailers: Option<&HeaderMap>, _latency: Duration, _span: &Span| {
560 ON_EOS().fetch_add(1, Ordering::AcqRel);
561 },
562 )
563 .on_failure(
564 |_class: ServerErrorsFailureClass, _latency: Duration, _span: &Span| {
565 ON_FAILURE().fetch_add(1, Ordering::AcqRel);
566 },
567 );
568
569 let svc = trace_layer.into_layer(service_fn(echo));
570
571 let res = svc.serve(Request::new(Body::from("foobar"))).await.unwrap();
572
573 assert_eq!(1, ON_REQUEST_COUNT().load(Ordering::Acquire), "request");
574 assert_eq!(1, ON_RESPONSE_COUNT().load(Ordering::Acquire), "request");
575 assert_eq!(
576 0,
577 ON_BODY_CHUNK_COUNT().load(Ordering::Acquire),
578 "body chunk"
579 );
580 assert_eq!(0, ON_EOS().load(Ordering::Acquire), "eos");
581 assert_eq!(0, ON_FAILURE().load(Ordering::Acquire), "failure");
582
583 res.into_body().collect().await.unwrap().to_bytes();
584 assert_eq!(
585 1,
586 ON_BODY_CHUNK_COUNT().load(Ordering::Acquire),
587 "body chunk"
588 );
589 assert_eq!(1, ON_EOS().load(Ordering::Acquire), "eos");
590 assert_eq!(0, ON_FAILURE().load(Ordering::Acquire), "failure");
591 }
592
593 #[tokio::test]
594 async fn streaming_response() {
595 lazy_atomic_u32!(
596 ON_REQUEST_COUNT,
597 ON_RESPONSE_COUNT,
598 ON_BODY_CHUNK_COUNT,
599 ON_EOS,
600 ON_FAILURE
601 );
602
603 let trace_layer = TraceLayer::new_for_http()
604 .on_request(|_req: &Request, _span: &Span| {
605 ON_REQUEST_COUNT().fetch_add(1, Ordering::AcqRel);
606 })
607 .on_response(|_res: &Response, _latency: Duration, _span: &Span| {
608 ON_RESPONSE_COUNT().fetch_add(1, Ordering::AcqRel);
609 })
610 .on_body_chunk(|_chunk: &Bytes, _latency: Duration, _span: &Span| {
611 ON_BODY_CHUNK_COUNT().fetch_add(1, Ordering::AcqRel);
612 })
613 .on_eos(
614 |_trailers: Option<&HeaderMap>, _latency: Duration, _span: &Span| {
615 ON_EOS().fetch_add(1, Ordering::AcqRel);
616 },
617 )
618 .on_failure(
619 |_class: ServerErrorsFailureClass, _latency: Duration, _span: &Span| {
620 ON_FAILURE().fetch_add(1, Ordering::AcqRel);
621 },
622 );
623
624 let svc = trace_layer.into_layer(service_fn(streaming_body));
625
626 let res = svc.serve(Request::new(Body::empty())).await.unwrap();
627
628 assert_eq!(1, ON_REQUEST_COUNT().load(Ordering::Acquire), "request");
629 assert_eq!(1, ON_RESPONSE_COUNT().load(Ordering::Acquire), "request");
630 assert_eq!(
631 0,
632 ON_BODY_CHUNK_COUNT().load(Ordering::Acquire),
633 "body chunk"
634 );
635 assert_eq!(0, ON_EOS().load(Ordering::Acquire), "eos");
636 assert_eq!(0, ON_FAILURE().load(Ordering::Acquire), "failure");
637
638 res.into_body().collect().await.unwrap().to_bytes();
639 assert_eq!(
640 3,
641 ON_BODY_CHUNK_COUNT().load(Ordering::Acquire),
642 "body chunk"
643 );
644 assert_eq!(1, ON_EOS().load(Ordering::Acquire), "eos");
645 assert_eq!(0, ON_FAILURE().load(Ordering::Acquire), "failure");
646 }
647
648 #[tokio::test]
649 async fn classify_eos_on_trailers_success() {
650 lazy_atomic_u32!(ON_EOS, ON_FAILURE);
651
652 let trace_layer = TraceLayer::new(TestClassify::new(false))
653 .on_eos(
654 |_trailers: Option<&HeaderMap>, _latency: Duration, _span: &Span| {
655 ON_EOS().fetch_add(1, Ordering::SeqCst);
656 },
657 )
658 .on_failure(|_class: &'static str, _latency: Duration, _span: &Span| {
659 ON_FAILURE().fetch_add(1, Ordering::SeqCst);
660 });
661
662 let svc = trace_layer.into_layer(service_fn(body_with_trailers));
663
664 let res = svc.serve(Request::new(Body::empty())).await.unwrap();
665
666 res.into_body().collect().await.unwrap().to_bytes();
667 assert_eq!(1, ON_EOS().load(Ordering::SeqCst), "eos");
668 assert_eq!(0, ON_FAILURE().load(Ordering::SeqCst), "failure");
669 }
670
671 #[tokio::test]
672 async fn classify_eos_on_trailers_failure() {
673 lazy_atomic_u32!(ON_EOS, ON_FAILURE);
674
675 let trace_layer = TraceLayer::new(TestClassify::new(true))
676 .on_eos(
677 |_trailers: Option<&HeaderMap>, _latency: Duration, _span: &Span| {
678 ON_EOS().fetch_add(1, Ordering::SeqCst);
679 },
680 )
681 .on_failure(|_class: &'static str, _latency: Duration, _span: &Span| {
682 ON_FAILURE().fetch_add(1, Ordering::SeqCst);
683 });
684
685 let svc = trace_layer.into_layer(service_fn(body_with_trailers));
686
687 let res = svc.serve(Request::new(Body::empty())).await.unwrap();
688
689 res.into_body().collect().await.unwrap().to_bytes();
690 assert_eq!(1, ON_EOS().load(Ordering::SeqCst), "eos");
691 assert_eq!(1, ON_FAILURE().load(Ordering::SeqCst), "failure");
692 }
693
694 #[tokio::test]
695 async fn classify_eos_on_empty_stream() {
696 lazy_atomic_u32!(ON_EOS, ON_FAILURE);
697
698 let trace_layer = TraceLayer::new(TestClassify::new(true))
699 .on_eos(
700 |_trailers: Option<&HeaderMap>, _latency: Duration, _span: &Span| {
701 ON_EOS().fetch_add(1, Ordering::SeqCst);
702 },
703 )
704 .on_failure(|_class: &'static str, _latency: Duration, _span: &Span| {
705 ON_FAILURE().fetch_add(1, Ordering::SeqCst);
706 });
707
708 let svc = trace_layer.into_layer(service_fn(streaming_body));
709
710 let res = svc.serve(Request::new(Body::empty())).await.unwrap();
711
712 res.into_body().collect().await.unwrap().to_bytes();
713 assert_eq!(1, ON_EOS().load(Ordering::SeqCst), "eos");
714 assert_eq!(0, ON_FAILURE().load(Ordering::SeqCst), "failure");
715 }
716
717 #[tokio::test]
718 async fn on_eos_fires_for_content_length_body() {
719 // Simulates the scenario where a consumer stops polling after receiving
720 // all bytes (as hyper does when Content-Length is exact). We poll only
721 // the data frame and never poll to None.
722
723 lazy_atomic_u32!(ON_BODY_CHUNK_COUNT, ON_EOS);
724
725 let trace_layer = TraceLayer::new_for_http()
726 .on_body_chunk(|_chunk: &Bytes, _latency: Duration, _span: &Span| {
727 ON_BODY_CHUNK_COUNT().fetch_add(1, Ordering::SeqCst);
728 })
729 .on_eos(
730 |_trailers: Option<&HeaderMap>, _latency: Duration, _span: &Span| {
731 ON_EOS().fetch_add(1, Ordering::SeqCst);
732 },
733 );
734
735 let svc = trace_layer.into_layer(service_fn(echo));
736
737 let res = svc.serve(Request::new(Body::from("hello"))).await.unwrap();
738
739 let mut body = res.into_body();
740
741 // Poll only the data frame (simulating a content-length aware consumer)
742 let frame = body.frame().await.unwrap().unwrap();
743 assert!(frame.data_ref().is_some());
744
745 // on_eos should have fired immediately after the data frame since
746 // is_end_stream() is true for Full bodies after yielding their data.
747 assert_eq!(
748 1,
749 ON_BODY_CHUNK_COUNT().load(Ordering::SeqCst),
750 "body chunk"
751 );
752 assert_eq!(1, ON_EOS().load(Ordering::SeqCst), "eos");
753 }
754
755 #[tokio::test]
756 async fn on_eos_fires_for_streaming_body_on_none() {
757 // Streaming bodies (no content-length) don't report is_end_stream()
758 // until polled to None. Verify on_eos still fires via the None path.
759 lazy_atomic_u32!(ON_EOS);
760
761 let trace_layer = TraceLayer::new_for_http().on_eos(
762 |_trailers: Option<&HeaderMap>, _latency: Duration, _span: &Span| {
763 ON_EOS().fetch_add(1, Ordering::SeqCst);
764 },
765 );
766
767 let svc = trace_layer.into_layer(service_fn(streaming_body));
768
769 let res = svc.serve(Request::new(Body::empty())).await.unwrap();
770
771 res.into_body().collect().await.unwrap();
772 assert_eq!(1, ON_EOS().load(Ordering::SeqCst), "eos");
773 }
774
775 #[tokio::test]
776 async fn on_eos_not_called_twice() {
777 // When is_end_stream() fires on_eos after a data frame, a subsequent
778 // poll returning None must not fire on_eos again.
779 lazy_atomic_u32!(ON_EOS);
780
781 let trace_layer = TraceLayer::new_for_http().on_eos(
782 |_trailers: Option<&HeaderMap>, _latency: Duration, _span: &Span| {
783 ON_EOS().fetch_add(1, Ordering::SeqCst);
784 },
785 );
786
787 let svc = trace_layer.into_layer(service_fn(echo));
788
789 let res = svc.serve(Request::new(Body::from("hello"))).await.unwrap();
790
791 // Consume the body fully (polls data frame then None)
792 res.into_body().collect().await.unwrap();
793
794 // on_eos should fire exactly once, not twice
795 assert_eq!(1, ON_EOS().load(Ordering::SeqCst), "eos");
796 }
797
798 async fn echo(req: Request) -> Result<Response, BoxError> {
799 Ok(Response::new(req.into_body()))
800 }
801
802 async fn streaming_body(_req: Request) -> Result<Response, BoxError> {
803 use rama_core::futures::stream::iter;
804
805 let stream = iter(vec![
806 Ok::<_, BoxError>(Bytes::from("one")),
807 Ok::<_, BoxError>(Bytes::from("two")),
808 Ok::<_, BoxError>(Bytes::from("three")),
809 ]);
810
811 let body = Body::from_stream(stream);
812
813 Ok(Response::new(body))
814 }
815
816 async fn body_with_trailers(_req: Request<Body>) -> Result<Response<Body>, BoxError> {
817 let mut trailers = HeaderMap::new();
818 trailers.insert("x-test-trailer", "value".parse().unwrap());
819 let body = Body::from(Bytes::from("data")).with_trailer_headers(trailers);
820 Ok(Response::new(body))
821 }
822
823 #[derive(Clone)]
824 struct TestClassify {
825 reject: bool,
826 }
827
828 impl TestClassify {
829 fn new(reject: bool) -> Self {
830 Self { reject }
831 }
832 }
833
834 impl MakeClassifier for TestClassify {
835 type FailureClass = &'static str;
836 type ClassifyEos = TestClassifyEos;
837 type Classifier = TestClassifyResponse;
838
839 fn make_classifier<B>(&self, _req: &Request<B>) -> Self::Classifier {
840 TestClassifyResponse {
841 reject: self.reject,
842 }
843 }
844 }
845
846 #[derive(Clone)]
847 struct TestClassifyResponse {
848 reject: bool,
849 }
850
851 impl ClassifyResponse for TestClassifyResponse {
852 type FailureClass = &'static str;
853 type ClassifyEos = TestClassifyEos;
854
855 fn classify_response<B>(
856 self,
857 _res: &Response<B>,
858 ) -> ClassifiedResponse<Self::FailureClass, Self::ClassifyEos> {
859 ClassifiedResponse::RequiresEos(TestClassifyEos {
860 reject: self.reject,
861 })
862 }
863
864 fn classify_error<E>(self, _error: &E) -> Self::FailureClass
865 where
866 E: std::fmt::Display,
867 {
868 "error"
869 }
870 }
871
872 #[derive(Clone)]
873 struct TestClassifyEos {
874 reject: bool,
875 }
876
877 impl ClassifyEos for TestClassifyEos {
878 type FailureClass = &'static str;
879
880 fn classify_eos(self, _trailers: Option<&HeaderMap>) -> Result<(), Self::FailureClass> {
881 if self.reject {
882 Err("classified as failure")
883 } else {
884 Ok(())
885 }
886 }
887
888 fn classify_error<E>(self, _error: &E) -> Self::FailureClass
889 where
890 E: std::fmt::Display,
891 {
892 "error"
893 }
894 }
895
896 /// Regression test for https://github.com/tower-rs/tower-http/issues/655
897 ///
898 /// Reproduces the reported bug: when a subscriber's filter disables the
899 /// request span but still enables events, the events appear without any
900 /// span context. This happens because `Span::enter()` on a disabled span
901 /// is a no-op, so events relying on ambient context have no parent.
902 ///
903 /// The fix (using explicit `parent: span`) ensures events always reference
904 /// the request span, even when it's disabled. A subscriber that records
905 /// disabled spans will still see the correct parent relationship.
906 #[test]
907 fn events_have_span_context_when_span_is_disabled() {
908 use parking_lot::Mutex;
909 use std::sync::Arc;
910 use tracing::subscriber::with_default;
911 use tracing_subscriber::{Layer as _, layer::SubscriberExt, registry::LookupSpan};
912
913 /// A filter that disables spans (by rejecting at the span level)
914 /// but allows all events through. This simulates the scenario where
915 /// a per-layer EnvFilter disables the request span's callsite.
916 struct DisableSpansFilter;
917
918 impl<S: tracing::Subscriber> tracing_subscriber::layer::Filter<S> for DisableSpansFilter {
919 fn enabled(
920 &self,
921 meta: &tracing::Metadata<'_>,
922 _cx: &tracing_subscriber::layer::Context<'_, S>,
923 ) -> bool {
924 // Disable spans, keep events
925 !meta.is_span()
926 }
927 }
928
929 /// Records (event_message, has_any_parent) pairs.
930 #[derive(Clone)]
931 struct RecordingLayer {
932 events: Arc<Mutex<Vec<(String, bool)>>>,
933 }
934
935 impl<S> tracing_subscriber::Layer<S> for RecordingLayer
936 where
937 S: tracing::Subscriber + for<'a> LookupSpan<'a>,
938 {
939 fn on_event(
940 &self,
941 event: &tracing::Event<'_>,
942 ctx: tracing_subscriber::layer::Context<'_, S>,
943 ) {
944 let mut msg = String::new();
945 event.record(&mut MessageVisitor(&mut msg));
946
947 // Check if the event has ANY parent: explicit or contextual
948 let has_parent = event.parent().is_some() || ctx.event_span(event).is_some();
949
950 self.events.lock().push((msg, has_parent));
951 }
952 }
953
954 struct MessageVisitor<'a>(&'a mut String);
955 impl tracing::field::Visit for MessageVisitor<'_> {
956 fn record_debug(&mut self, field: &tracing::field::Field, value: &dyn std::fmt::Debug) {
957 if field.name() == "message" {
958 *self.0 = format!("{value:?}");
959 }
960 }
961 }
962
963 let events = Arc::new(Mutex::new(Vec::new()));
964 let layer = RecordingLayer {
965 events: events.clone(),
966 };
967 let subscriber = tracing_subscriber::registry().with(layer.with_filter(DisableSpansFilter));
968
969 // Use with_default to guarantee cleanup even on panic, avoiding
970 // cross-test subscriber pollution.
971 with_default(subscriber, || {
972 let rt = tokio::runtime::Builder::new_current_thread()
973 .enable_all()
974 .build()
975 .unwrap();
976
977 rt.block_on(async {
978 let svc = TraceLayer::new_for_http().into_layer(service_fn(echo));
979
980 let res = svc.serve(Request::new(Body::from("test"))).await.unwrap();
981
982 res.into_body().collect().await.unwrap();
983 });
984 });
985
986 let events = events.lock();
987 let request_events: Vec<_> = events
988 .iter()
989 .filter(|(msg, _)| {
990 msg.contains("started processing request")
991 || msg.contains("finished processing request")
992 })
993 .collect();
994
995 assert!(
996 request_events.len() >= 2,
997 "expected on_request and on_response events to fire"
998 );
999
1000 // The bug: without explicit parent, these events have no span context
1001 // at all when the request span is disabled. With the fix, they still
1002 // reference the span (even though it's disabled).
1003 for (msg, has_parent) in &request_events {
1004 assert!(
1005 *has_parent,
1006 "event {msg:?} has no span context. When the request span is \
1007 disabled by a filter, events must still reference it via \
1008 explicit parent so subscribers can associate them correctly."
1009 );
1010 }
1011 }
1012}