Skip to main content

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}