rama_http/layer/traffic_writer/
request.rs1use 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
14async 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
27pub trait RequestWriter: Send + Sync + 'static {
29 fn write_request(&self, req: Request) -> impl Future<Output = ()> + Send + '_;
31}
32
33#[derive(Debug, Clone, Default, Extension)]
35#[extension(tags(http))]
36#[non_exhaustive]
37pub struct DoNotWriteRequest;
38
39impl DoNotWriteRequest {
40 #[must_use]
42 pub const fn new() -> Self {
43 Self
44 }
45}
46
47#[derive(Clone)]
48pub struct RequestWriterService<S, W> {
52 inner: S,
53 writer: W,
54}
55
56impl<S, W> RequestWriterService<S, W> {
57 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 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 #[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 #[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 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 #[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 #[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)]
226pub struct RequestWriterLayer<W> {
230 writer: W,
231}
232
233impl<W> RequestWriterLayer<W> {
234 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 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 #[must_use]
275 pub fn stdout_unbounded(executor: &Executor, mode: Option<WriterMode>) -> Self {
276 Self::writer_unbounded(executor, stdout(), mode)
277 }
278
279 #[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 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 #[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 #[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}