Skip to main content

rama_http/layer/trace/
service.rs

1use super::{
2    DefaultMakeSpan, DefaultOnBodyChunk, DefaultOnEos, DefaultOnFailure, DefaultOnRequest,
3    DefaultOnResponse, GrpcMakeClassifier, HttpMakeClassifier, MakeSpan, OnBodyChunk, OnEos,
4    OnFailure, OnRequest, OnResponse, ResponseBody,
5};
6use crate::layer::classify::{
7    ClassifiedResponse, ClassifyResponse, GrpcErrorsAsFailures, MakeClassifier,
8    ServerErrorsAsFailures, SharedClassifier,
9};
10use crate::{Request, Response, StreamingBody};
11use rama_core::{Service, telemetry::tracing::Instrument};
12use rama_utils::macros::define_inner_service_accessors;
13use std::{fmt, time::Instant};
14
15/// Middleware that adds high level [tracing] to a [`Service`].
16///
17/// See the [module docs](crate::layer::trace) for an example.
18///
19/// [tracing]: https://crates.io/crates/tracing
20/// [`Service`]: rama_core::Service
21#[derive(Debug, Clone)]
22pub struct Trace<
23    S,
24    M,
25    MakeSpan = DefaultMakeSpan,
26    OnRequest = DefaultOnRequest,
27    OnResponse = DefaultOnResponse,
28    OnBodyChunk = DefaultOnBodyChunk,
29    OnEos = DefaultOnEos,
30    OnFailure = DefaultOnFailure,
31> {
32    pub(crate) inner: S,
33    pub(crate) make_classifier: M,
34    pub(crate) make_span: MakeSpan,
35    pub(crate) on_request: OnRequest,
36    pub(crate) on_response: OnResponse,
37    pub(crate) on_body_chunk: OnBodyChunk,
38    pub(crate) on_eos: OnEos,
39    pub(crate) on_failure: OnFailure,
40}
41
42impl<S, M> Trace<S, M> {
43    /// Create a new [`Trace`] using the given [`MakeClassifier`].
44    pub fn new(inner: S, make_classifier: M) -> Self
45    where
46        M: MakeClassifier,
47    {
48        Self {
49            inner,
50            make_classifier,
51            make_span: DefaultMakeSpan::new(),
52            on_request: DefaultOnRequest::default(),
53            on_response: DefaultOnResponse::default(),
54            on_body_chunk: DefaultOnBodyChunk::default(),
55            on_eos: DefaultOnEos::default(),
56            on_failure: DefaultOnFailure::default(),
57        }
58    }
59}
60
61impl<S, M, MakeSpan, OnRequest, OnResponse, OnBodyChunk, OnEos, OnFailure>
62    Trace<S, M, MakeSpan, OnRequest, OnResponse, OnBodyChunk, OnEos, OnFailure>
63{
64    define_inner_service_accessors!();
65
66    /// Customize what to do when a request is received.
67    ///
68    /// `NewOnRequest` is expected to implement [`OnRequest`].
69    ///
70    /// [`OnRequest`]: super::OnRequest
71    pub fn on_request<NewOnRequest>(
72        self,
73        new_on_request: NewOnRequest,
74    ) -> Trace<S, M, MakeSpan, NewOnRequest, OnResponse, OnBodyChunk, OnEos, OnFailure> {
75        Trace {
76            on_request: new_on_request,
77            inner: self.inner,
78            on_failure: self.on_failure,
79            on_eos: self.on_eos,
80            on_body_chunk: self.on_body_chunk,
81            make_span: self.make_span,
82            on_response: self.on_response,
83            make_classifier: self.make_classifier,
84        }
85    }
86
87    /// Customize what to do when a response has been produced.
88    ///
89    /// `NewOnResponse` is expected to implement [`OnResponse`].
90    ///
91    /// [`OnResponse`]: super::OnResponse
92    pub fn on_response<NewOnResponse>(
93        self,
94        new_on_response: NewOnResponse,
95    ) -> Trace<S, M, MakeSpan, OnRequest, NewOnResponse, OnBodyChunk, OnEos, OnFailure> {
96        Trace {
97            on_response: new_on_response,
98            inner: self.inner,
99            on_request: self.on_request,
100            on_failure: self.on_failure,
101            on_body_chunk: self.on_body_chunk,
102            on_eos: self.on_eos,
103            make_span: self.make_span,
104            make_classifier: self.make_classifier,
105        }
106    }
107
108    /// Customize what to do when a body chunk has been sent.
109    ///
110    /// `NewOnBodyChunk` is expected to implement [`OnBodyChunk`].
111    ///
112    /// [`OnBodyChunk`]: super::OnBodyChunk
113    pub fn on_body_chunk<NewOnBodyChunk>(
114        self,
115        new_on_body_chunk: NewOnBodyChunk,
116    ) -> Trace<S, M, MakeSpan, OnRequest, OnResponse, NewOnBodyChunk, OnEos, OnFailure> {
117        Trace {
118            on_body_chunk: new_on_body_chunk,
119            on_eos: self.on_eos,
120            make_span: self.make_span,
121            inner: self.inner,
122            on_failure: self.on_failure,
123            on_request: self.on_request,
124            on_response: self.on_response,
125            make_classifier: self.make_classifier,
126        }
127    }
128
129    /// Customize what to do when a streaming response has closed.
130    ///
131    /// `NewOnEos` is expected to implement [`OnEos`].
132    ///
133    /// [`OnEos`]: super::OnEos
134    pub fn on_eos<NewOnEos>(
135        self,
136        new_on_eos: NewOnEos,
137    ) -> Trace<S, M, MakeSpan, OnRequest, OnResponse, OnBodyChunk, NewOnEos, OnFailure> {
138        Trace {
139            on_eos: new_on_eos,
140            make_span: self.make_span,
141            inner: self.inner,
142            on_failure: self.on_failure,
143            on_request: self.on_request,
144            on_body_chunk: self.on_body_chunk,
145            on_response: self.on_response,
146            make_classifier: self.make_classifier,
147        }
148    }
149
150    /// Customize what to do when a response has been classified as a failure.
151    ///
152    /// `NewOnFailure` is expected to implement [`OnFailure`].
153    ///
154    /// [`OnFailure`]: super::OnFailure
155    pub fn on_failure<NewOnFailure>(
156        self,
157        new_on_failure: NewOnFailure,
158    ) -> Trace<S, M, MakeSpan, OnRequest, OnResponse, OnBodyChunk, OnEos, NewOnFailure> {
159        Trace {
160            on_failure: new_on_failure,
161            inner: self.inner,
162            make_span: self.make_span,
163            on_body_chunk: self.on_body_chunk,
164            on_request: self.on_request,
165            on_eos: self.on_eos,
166            on_response: self.on_response,
167            make_classifier: self.make_classifier,
168        }
169    }
170
171    /// Customize how to make [`Span`]s that all request handling will be wrapped in.
172    ///
173    /// `NewMakeSpan` is expected to implement [`MakeSpan`].
174    ///
175    /// [`MakeSpan`]: super::MakeSpan
176    /// [`Span`]: rama_core::telemetry::tracing::Span
177    pub fn make_span_with<NewMakeSpan>(
178        self,
179        new_make_span: NewMakeSpan,
180    ) -> Trace<S, M, NewMakeSpan, OnRequest, OnResponse, OnBodyChunk, OnEos, OnFailure> {
181        Trace {
182            make_span: new_make_span,
183            inner: self.inner,
184            on_failure: self.on_failure,
185            on_request: self.on_request,
186            on_body_chunk: self.on_body_chunk,
187            on_response: self.on_response,
188            on_eos: self.on_eos,
189            make_classifier: self.make_classifier,
190        }
191    }
192}
193
194impl<S>
195    Trace<
196        S,
197        HttpMakeClassifier,
198        DefaultMakeSpan,
199        DefaultOnRequest,
200        DefaultOnResponse,
201        DefaultOnBodyChunk,
202        DefaultOnEos,
203        DefaultOnFailure,
204    >
205{
206    /// Create a new [`Trace`] using [`ServerErrorsAsFailures`] which supports classifying
207    /// regular HTTP responses based on the status code.
208    pub fn new_for_http(inner: S) -> Self {
209        Self {
210            inner,
211            make_classifier: SharedClassifier::new(ServerErrorsAsFailures::default()),
212            make_span: DefaultMakeSpan::new(),
213            on_request: DefaultOnRequest::default(),
214            on_response: DefaultOnResponse::default(),
215            on_body_chunk: DefaultOnBodyChunk::default(),
216            on_eos: DefaultOnEos::default(),
217            on_failure: DefaultOnFailure::default(),
218        }
219    }
220}
221
222impl<S>
223    Trace<
224        S,
225        GrpcMakeClassifier,
226        DefaultMakeSpan,
227        DefaultOnRequest,
228        DefaultOnResponse,
229        DefaultOnBodyChunk,
230        DefaultOnEos,
231        DefaultOnFailure,
232    >
233{
234    /// Create a new [`Trace`] using [`GrpcErrorsAsFailures`] which supports classifying
235    /// gRPC responses and streams based on the `grpc-status` header.
236    pub fn new_for_grpc(inner: S) -> Self {
237        Self {
238            inner,
239            make_classifier: SharedClassifier::new(GrpcErrorsAsFailures::default()),
240            make_span: DefaultMakeSpan::new(),
241            on_request: DefaultOnRequest::default(),
242            on_response: DefaultOnResponse::default(),
243            on_body_chunk: DefaultOnBodyChunk::default(),
244            on_eos: DefaultOnEos::default(),
245            on_failure: DefaultOnFailure::default(),
246        }
247    }
248}
249
250impl<S, ReqBody, ResBody, M, OnRequestT, OnResponseT, OnFailureT, OnBodyChunkT, OnEosT, MakeSpanT>
251    Service<Request<ReqBody>>
252    for Trace<S, M, MakeSpanT, OnRequestT, OnResponseT, OnBodyChunkT, OnEosT, OnFailureT>
253where
254    S: Service<Request<ReqBody>, Output = Response<ResBody>, Error: fmt::Display>,
255    ReqBody: StreamingBody + Send + 'static,
256    ResBody: StreamingBody<Error: fmt::Display> + Send + Sync + 'static,
257    M: MakeClassifier<Classifier: Clone>,
258    MakeSpanT: MakeSpan<ReqBody>,
259    OnRequestT: OnRequest<ReqBody>,
260    OnResponseT: OnResponse<ResBody> + Clone,
261    OnBodyChunkT: OnBodyChunk<ResBody::Data> + Clone,
262    OnEosT: OnEos + Clone,
263    OnFailureT: OnFailure<M::FailureClass> + Clone,
264{
265    type Output = Response<ResponseBody<ResBody, M::ClassifyEos, OnBodyChunkT, OnEosT, OnFailureT>>;
266    type Error = S::Error;
267
268    async fn serve(&self, req: Request<ReqBody>) -> Result<Self::Output, Self::Error> {
269        let start = Instant::now();
270
271        let span = self.make_span.make_span(&req);
272
273        let classifier = self.make_classifier.make_classifier(&req);
274
275        span.in_scope(|| self.on_request.on_request(&req, &span));
276        let result = self.inner.serve(req).instrument(span.clone()).await;
277        let latency = start.elapsed();
278
279        match result {
280            Ok(res) => {
281                let classification = classifier.classify_response(&res);
282
283                self.on_response.clone().on_response(&res, latency, &span);
284
285                match classification {
286                    ClassifiedResponse::Ready(classification) => {
287                        if let Err(failure_class) = classification {
288                            self.on_failure.on_failure(failure_class, latency, &span);
289                        }
290
291                        let span = span.clone();
292                        let res = res.map(|body| ResponseBody {
293                            inner: body,
294                            classify_eos: None,
295                            on_eos: Some((self.on_eos.clone(), Instant::now())),
296                            on_body_chunk: self.on_body_chunk.clone(),
297                            on_failure: Some(self.on_failure.clone()),
298                            start,
299                            span,
300                        });
301
302                        Ok(res)
303                    }
304                    ClassifiedResponse::RequiresEos(classify_eos) => {
305                        let span = span.clone();
306                        let res = res.map(|body| ResponseBody {
307                            inner: body,
308                            classify_eos: Some(classify_eos),
309                            on_eos: Some((self.on_eos.clone(), Instant::now())),
310                            on_body_chunk: self.on_body_chunk.clone(),
311                            on_failure: Some(self.on_failure.clone()),
312                            start,
313                            span,
314                        });
315
316                        Ok(res)
317                    }
318                }
319            }
320            Err(err) => {
321                let failure_class: <M as MakeClassifier>::FailureClass =
322                    classifier.classify_error(&err);
323                self.on_failure.on_failure(failure_class, latency, &span);
324
325                Err(err)
326            }
327        }
328    }
329}