Skip to main content

rama_http/layer/traffic_writer/
request.rs

1use super::{WriterMode, write_headers_body_flags};
2use crate::io::write_http_request;
3use crate::{Body, Request, StreamingBody, body::util::BodyExt};
4use rama_core::bytes::Bytes;
5use rama_core::error::{BoxError, ErrorContext as _};
6use rama_core::extensions::{Extension, ExtensionsRef};
7use rama_core::rt::Executor;
8use rama_core::telemetry::tracing::{self, Instrument};
9use rama_core::{Layer, Service};
10use std::fmt::Debug;
11use tokio::io::{AsyncWrite, AsyncWriteExt, stderr, stdout};
12use tokio::sync::mpsc::{Sender, UnboundedSender, channel, unbounded_channel};
13
14/// Write a single request entry (request followed by a `\r\n` separator) to the writer.
15async fn write_request_entry<W>(writer: &mut W, req: Request, write_headers: bool, write_body: bool)
16where
17    W: AsyncWrite + Unpin + Send + Sync + 'static,
18{
19    if let Err(err) = write_http_request(writer, req, write_headers, write_body).await {
20        tracing::error!("failed to write http request to writer: {err:?}")
21    }
22    if let Err(err) = writer.write_all(b"\r\n").await {
23        tracing::error!("failed to write separator to writer: {err:?}")
24    }
25}
26
27/// A trait for writing http requests.
28pub trait RequestWriter: Send + Sync + 'static {
29    /// Write the http request.
30    fn write_request(&self, req: Request) -> impl Future<Output = ()> + Send + '_;
31}
32
33/// Marker struct to indicate that the request should not be printed.
34#[derive(Debug, Clone, Default, Extension)]
35#[extension(tags(http))]
36#[non_exhaustive]
37pub struct DoNotWriteRequest;
38
39impl DoNotWriteRequest {
40    /// Create a new [`DoNotWriteRequest`] marker.
41    #[must_use]
42    pub const fn new() -> Self {
43        Self
44    }
45}
46
47#[derive(Clone)]
48/// Middleware to print Http request in std format.
49///
50/// See the [module docs](super) for more details.
51pub struct RequestWriterService<S, W> {
52    inner: S,
53    writer: W,
54}
55
56impl<S, W> RequestWriterService<S, W> {
57    /// Create a new [`RequestWriterService`] with a custom [`RequestWriter`].
58    pub const fn new(inner: S, writer: W) -> Self {
59        Self { inner, writer }
60    }
61}
62
63impl<S: Debug, W> Debug for RequestWriterService<S, W> {
64    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
65        f.debug_struct("RequestWriterService")
66            .field("inner", &self.inner)
67            .field("writer", &format_args!("{}", std::any::type_name::<W>()))
68            .finish()
69    }
70}
71
72impl<S> RequestWriterService<S, UnboundedSender<Request>> {
73    /// Create a new [`RequestWriterService`] that prints requests to an [`AsyncWrite`]r
74    /// over an unbounded channel
75    pub fn writer_unbounded<W>(
76        inner: S,
77        executor: &Executor,
78        mut writer: W,
79        mode: Option<WriterMode>,
80    ) -> Self
81    where
82        W: AsyncWrite + Unpin + Send + Sync + 'static,
83    {
84        let (tx, mut rx) = unbounded_channel();
85        let (write_headers, write_body) = write_headers_body_flags(mode);
86
87        let span =
88            tracing::trace_root_span!("TrafficWriter::request::unbounded", otel.kind = "consumer");
89
90        executor.spawn_task(
91            async move {
92                while let Some(req) = rx.recv().await {
93                    write_request_entry(&mut writer, req, write_headers, write_body).await;
94                }
95            }
96            .instrument(span),
97        );
98        Self { writer: tx, inner }
99    }
100
101    /// Create a new [`RequestWriterService`] that prints requests to stdout
102    /// over an unbounded channel.
103    #[must_use]
104    pub fn stdout_unbounded(inner: S, executor: &Executor, mode: Option<WriterMode>) -> Self {
105        Self::writer_unbounded(inner, executor, stdout(), mode)
106    }
107
108    /// Create a new [`RequestWriterService`] that prints requests to stderr
109    /// over an unbounded channel.
110    #[must_use]
111    pub fn stderr_unbounded(inner: S, executor: &Executor, mode: Option<WriterMode>) -> Self {
112        Self::writer_unbounded(inner, executor, stderr(), mode)
113    }
114}
115
116impl<S> RequestWriterService<S, Sender<Request>> {
117    /// Create a new [`RequestWriterService`] that prints requests to an [`AsyncWrite`]r
118    /// over a bounded channel with a fixed buffer size.
119    pub fn writer<W>(
120        inner: S,
121        executor: &Executor,
122        mut writer: W,
123        buffer_size: usize,
124        mode: Option<WriterMode>,
125    ) -> Self
126    where
127        W: AsyncWrite + Unpin + Send + Sync + 'static,
128    {
129        let (tx, mut rx) = channel(buffer_size);
130        let (write_headers, write_body) = write_headers_body_flags(mode);
131
132        let span =
133            tracing::trace_root_span!("TrafficWriter::request::bounded", otel.kind = "consumer");
134
135        executor.spawn_task(
136            async move {
137                while let Some(req) = rx.recv().await {
138                    write_request_entry(&mut writer, req, write_headers, write_body).await;
139                }
140            }
141            .instrument(span),
142        );
143        Self { writer: tx, inner }
144    }
145
146    /// Create a new [`RequestWriterService`] that prints requests to stdout
147    /// over a bounded channel with a fixed buffer size.
148    #[must_use]
149    pub fn stdout(
150        inner: S,
151        executor: &Executor,
152        buffer_size: usize,
153        mode: Option<WriterMode>,
154    ) -> Self {
155        Self::writer(inner, executor, stdout(), buffer_size, mode)
156    }
157
158    /// Create a new [`RequestWriterService`] that prints requests to stderr
159    /// over a bounded channel with a fixed buffer size.
160    #[must_use]
161    pub fn stderr(
162        inner: S,
163        executor: &Executor,
164        buffer_size: usize,
165        mode: Option<WriterMode>,
166    ) -> Self {
167        Self::writer(inner, executor, stderr(), buffer_size, mode)
168    }
169}
170
171impl<S, W, ReqBody> Service<Request<ReqBody>> for RequestWriterService<S, W>
172where
173    S: Service<Request, Error: Into<BoxError>>,
174    ReqBody: StreamingBody<Data = Bytes, Error: Into<BoxError>> + Send + Sync + 'static,
175    W: RequestWriter,
176{
177    type Error = BoxError;
178    type Output = S::Output;
179
180    async fn serve(&self, req: Request<ReqBody>) -> Result<Self::Output, Self::Error> {
181        let req = if req.extensions().get_ref::<DoNotWriteRequest>().is_some() {
182            req.map(Body::new)
183        } else {
184            let (parts, body) = req.into_parts();
185            let body_bytes = body
186                .collect()
187                .await
188                .context("printer prepare: collect request body")?
189                .to_bytes();
190            let req = Request::from_parts(parts.clone(), Body::from(body_bytes.clone()));
191            self.writer.write_request(req).await;
192            Request::from_parts(parts, Body::from(body_bytes))
193        };
194
195        self.inner.serve(req).await.into_box_error()
196    }
197}
198
199impl RequestWriter for Sender<Request> {
200    async fn write_request(&self, req: Request) {
201        if let Err(err) = self.send(req).await {
202            tracing::error!("failed to send request to channel: {err:?}")
203        }
204    }
205}
206
207impl RequestWriter for UnboundedSender<Request> {
208    async fn write_request(&self, req: Request) {
209        if let Err(err) = self.send(req) {
210            tracing::error!("failed to send request to unbounded channel: {err:?}")
211        }
212    }
213}
214
215impl<F, Fut> RequestWriter for F
216where
217    F: Fn(Request) -> Fut + Send + Sync + 'static,
218    Fut: Future<Output = ()> + Send + 'static,
219{
220    async fn write_request(&self, req: Request) {
221        self(req).await
222    }
223}
224
225#[derive(Clone)]
226/// Middleware to print Http request in std format.
227///
228/// See the [module docs](super) for more details.
229pub struct RequestWriterLayer<W> {
230    writer: W,
231}
232
233impl<W> RequestWriterLayer<W> {
234    /// Create a new [`RequestWriterLayer`] with a custom [`RequestWriter`].
235    pub const fn new(writer: W) -> Self {
236        Self { writer }
237    }
238}
239
240impl<W> Debug for RequestWriterLayer<W> {
241    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
242        f.debug_struct("RequestWriterLayer")
243            .field("writer", &format_args!("{}", std::any::type_name::<W>()))
244            .finish()
245    }
246}
247
248impl RequestWriterLayer<UnboundedSender<Request>> {
249    /// Create a new [`RequestWriterLayer`] that prints requests to an [`AsyncWrite`]r
250    /// over an unbounded channel
251    pub fn writer_unbounded<W>(executor: &Executor, mut writer: W, mode: Option<WriterMode>) -> Self
252    where
253        W: AsyncWrite + Unpin + Send + Sync + 'static,
254    {
255        let (tx, mut rx) = unbounded_channel();
256        let (write_headers, write_body) = write_headers_body_flags(mode);
257
258        let span =
259            tracing::trace_root_span!("TrafficWriter::request::unbounded", otel.kind = "consumer");
260
261        executor.spawn_task(
262            async move {
263                while let Some(req) = rx.recv().await {
264                    write_request_entry(&mut writer, req, write_headers, write_body).await;
265                }
266            }
267            .instrument(span),
268        );
269        Self { writer: tx }
270    }
271
272    /// Create a new [`RequestWriterService`] that prints requests to stdout
273    /// over an unbounded channel.
274    #[must_use]
275    pub fn stdout_unbounded(executor: &Executor, mode: Option<WriterMode>) -> Self {
276        Self::writer_unbounded(executor, stdout(), mode)
277    }
278
279    /// Create a new [`RequestWriterService`] that prints requests to stderr
280    /// over an unbounded channel.
281    #[must_use]
282    pub fn stderr_unbounded(executor: &Executor, mode: Option<WriterMode>) -> Self {
283        Self::writer_unbounded(executor, stderr(), mode)
284    }
285}
286
287impl RequestWriterLayer<Sender<Request>> {
288    /// Create a new [`RequestWriterLayer`] that prints requests to an [`AsyncWrite`]r
289    /// over a bounded channel with a fixed buffer size.
290    pub fn writer<W>(
291        executor: &Executor,
292        mut writer: W,
293        buffer_size: usize,
294        mode: Option<WriterMode>,
295    ) -> Self
296    where
297        W: AsyncWrite + Unpin + Send + Sync + 'static,
298    {
299        let (tx, mut rx) = channel(buffer_size);
300        let (write_headers, write_body) = write_headers_body_flags(mode);
301
302        let span =
303            tracing::trace_root_span!("TrafficWriter::request::bounded", otel.kind = "consumer");
304
305        executor.spawn_task(
306            async move {
307                while let Some(req) = rx.recv().await {
308                    write_request_entry(&mut writer, req, write_headers, write_body).await;
309                }
310            }
311            .instrument(span),
312        );
313        Self { writer: tx }
314    }
315
316    /// Create a new [`RequestWriterService`] that prints requests to stdout
317    /// over a bounded channel with a fixed buffer size.
318    #[must_use]
319    pub fn stdout(executor: &Executor, buffer_size: usize, mode: Option<WriterMode>) -> Self {
320        Self::writer(executor, stdout(), buffer_size, mode)
321    }
322
323    /// Create a new [`RequestWriterService`] that prints requests to stderr
324    /// over a bounded channel with a fixed buffer size.
325    #[must_use]
326    pub fn stderr(executor: &Executor, buffer_size: usize, mode: Option<WriterMode>) -> Self {
327        Self::writer(executor, stderr(), buffer_size, mode)
328    }
329}
330
331impl<S, W: Clone> Layer<S> for RequestWriterLayer<W> {
332    type Service = RequestWriterService<S, W>;
333
334    fn layer(&self, inner: S) -> Self::Service {
335        RequestWriterService {
336            inner,
337            writer: self.writer.clone(),
338        }
339    }
340
341    fn into_layer(self, inner: S) -> Self::Service {
342        RequestWriterService {
343            inner,
344            writer: self.writer,
345        }
346    }
347}