1use std::pin::Pin;
2use std::task::{Context, Poll, ready};
3use tokio::io::{self, AsyncWrite, AsyncWriteExt};
4use tokio::sync::mpsc::{Receiver, Sender};
5
6pub struct PrefixWriter<W: AsyncWrite + Unpin> {
7 inner: W,
8 prefix: Vec<u8>,
9 at_start: bool,
10}
11
12impl<W: AsyncWrite + Unpin> PrefixWriter<W> {
13 pub fn new(writer: W, prefix: Vec<u8>) -> Self {
14 Self {
15 inner: writer,
16 prefix,
17 at_start: true,
18 }
19 }
20
21 pub fn into_inner(self) -> W {
22 self.inner
23 }
24}
25
26impl<W: AsyncWrite + Unpin> AsyncWrite for PrefixWriter<W> {
27 fn poll_write(
28 mut self: Pin<&mut Self>,
29 cx: &mut Context<'_>,
30 buf: &[u8],
31 ) -> Poll<io::Result<usize>> {
32 let mut written = 0;
33 let mut offset = 0;
34 while offset < buf.len() {
35 if self.at_start {
37 let prefix = self.prefix.clone();
38 let n = ready!(Pin::new(&mut self.inner).poll_write(cx, &prefix))?;
39 if n < self.prefix.len() {
40 return Poll::Ready(Ok(0));
41 }
42 self.at_start = false;
43 }
44 if let Some(pos) = buf[offset..].iter().position(|&b| b == b'\n') {
46 let end = offset + pos + 1;
47 let n = ready!(Pin::new(&mut self.inner).poll_write(cx, &buf[offset..end]))?;
48 written += n;
49 offset = end;
50 self.at_start = true;
51 } else {
52 let n = ready!(Pin::new(&mut self.inner).poll_write(cx, &buf[offset..]))?;
53 written += n;
54 offset = buf.len();
55 }
56 }
57 Poll::Ready(Ok(written))
58 }
59
60 fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
61 Pin::new(&mut self.inner).poll_flush(cx)
62 }
63
64 fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
65 Pin::new(&mut self.inner).poll_shutdown(cx)
66 }
67}
68
69pub struct WriteLine {
70 pub label: String,
71 pub line: Vec<u8>,
72}
73
74impl WriteLine {
75 pub fn new(label: String, line: Vec<u8>) -> Self {
76 Self { label, line }
77 }
78}
79
80pub struct LineWriter {
81 tx: Sender<Vec<u8>>,
82 buf: Vec<u8>,
83}
84
85impl LineWriter {
86 pub fn new(tx: Sender<Vec<u8>>) -> Self {
87 Self {
88 tx,
89 buf: Vec::new(),
90 }
91 }
92
93 pub fn into_inner(self) -> Sender<Vec<u8>> {
94 self.tx
95 }
96}
97
98impl AsyncWrite for LineWriter {
99 fn poll_write(
100 mut self: Pin<&mut Self>,
101 _cx: &mut Context<'_>,
102 buf: &[u8],
103 ) -> Poll<io::Result<usize>> {
104 self.buf.extend_from_slice(buf);
106 let mut start = 0;
107 while let Some(pos) = self.buf[start..].iter().position(|&b| b == b'\n') {
109 let end = start + pos + 1;
110 let line = self.buf[..end].to_vec();
111 let _ = self.tx.try_send(line);
112 start = end;
113 }
114 if start > 0 {
116 self.buf.drain(..start);
117 }
118 Poll::Ready(Ok(buf.len()))
119 }
120
121 fn poll_flush(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
122 if !self.buf.is_empty() {
123 let remaining = self.buf.split_off(0);
124 let _ = self.tx.try_send(remaining);
125 }
126 Poll::Ready(Ok(()))
127 }
128
129 fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
130 let _ = self.poll_flush(cx);
132 Poll::Ready(Ok(()))
133 }
134}
135
136pub struct TerminalWriter {
137 is_tty: bool,
138 stdout_tx: Sender<Vec<u8>>,
139 stdout_rx: Receiver<Vec<u8>>,
140 stderr_tx: Sender<Vec<u8>>,
141 stderr_rx: Receiver<Vec<u8>>,
142}
143
144impl Default for TerminalWriter {
145 fn default() -> Self {
146 Self::new(true, 1024)
147 }
148}
149
150impl TerminalWriter {
151 pub fn new(is_tty: bool, buffer: usize) -> Self {
152 let (stdout_tx, stdout_rx) = tokio::sync::mpsc::channel(buffer);
153 let (stderr_tx, stderr_rx) = tokio::sync::mpsc::channel(buffer);
154 Self {
155 is_tty,
156 stdout_tx,
157 stdout_rx,
158 stderr_tx,
159 stderr_rx,
160 }
161 }
162
163 pub fn stdout_raw(&self) -> Sender<Vec<u8>> {
164 self.stdout_tx.clone()
165 }
166
167 pub fn stdout(&self) -> LineWriter {
169 LineWriter::new(self.stdout_tx.clone())
170 }
171
172 pub fn stdout_with_label(&self, label: String) -> PrefixWriter<LineWriter> {
173 let stdout = self.stdout();
174 let mut prefix = Vec::new();
175 if self.is_tty {
176 prefix.extend_from_slice(b"\x1b[1;34m"); }
178 prefix.extend_from_slice(label.as_bytes());
179 prefix.extend_from_slice(b": ");
180 if self.is_tty {
181 prefix.extend_from_slice(b"\x1b[0m"); }
183
184 PrefixWriter::new(stdout, prefix)
185 }
186
187 pub fn stderr_raw(&self) -> Sender<Vec<u8>> {
188 self.stderr_tx.clone()
189 }
190
191 pub fn stderr(&self) -> LineWriter {
193 LineWriter::new(self.stderr_tx.clone())
194 }
195
196 pub fn stderr_with_label(&self, label: String) -> PrefixWriter<LineWriter> {
197 let stderr = self.stderr();
198 let mut prefix = Vec::new();
199 if self.is_tty {
200 prefix.extend_from_slice(b"\x1b[1;31m"); }
202 prefix.extend_from_slice(label.as_bytes());
203 prefix.extend_from_slice(b": ");
204 if self.is_tty {
205 prefix.extend_from_slice(b"\x1b[0m"); }
207
208 PrefixWriter::new(stderr, prefix)
209 }
210
211 pub async fn run(self) {
213 self.run_with(io::stdout(), io::stderr()).await;
214 }
215
216 pub async fn run_with<O, E>(self, mut out: O, mut err: E)
218 where
219 O: AsyncWrite + Unpin,
220 E: AsyncWrite + Unpin,
221 {
222 let TerminalWriter {
224 stdout_tx,
225 stderr_tx,
226 stdout_rx,
227 stderr_rx,
228 ..
229 } = self;
230
231 drop(stdout_tx);
233 drop(stderr_tx);
234
235 let stdout_fut = async {
237 let mut rx = stdout_rx;
238 while let Some(fragment) = rx.recv().await {
239 let _ = out.write_all(&fragment).await;
240 }
241 };
242
243 let stderr_fut = async {
244 let mut rx = stderr_rx;
245 while let Some(fragment) = rx.recv().await {
246 let _ = err.write_all(&fragment).await;
247 }
248 };
249
250 tokio::join!(stdout_fut, stderr_fut);
251 }
252}
253
254#[cfg(test)]
255mod tests {
256 use super::*;
257 use std::pin::Pin;
258 use std::task::{Context, Poll};
259 use tokio::io;
260 use tokio::io::AsyncWriteExt;
261
262 struct VecWriter {
264 pub data: Vec<u8>,
265 }
266 impl VecWriter {
267 fn new() -> Self {
268 VecWriter { data: Vec::new() }
269 }
270 }
271 impl io::AsyncWrite for VecWriter {
272 fn poll_write(
273 self: Pin<&mut Self>,
274 _cx: &mut Context<'_>,
275 buf: &[u8],
276 ) -> Poll<io::Result<usize>> {
277 let this = self.get_mut();
278 this.data.extend_from_slice(buf);
279 Poll::Ready(Ok(buf.len()))
280 }
281 fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
282 Poll::Ready(Ok(()))
283 }
284 fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
285 Poll::Ready(Ok(()))
286 }
287 }
288 impl Unpin for VecWriter {}
289
290 #[tokio::test]
291 async fn test_prefix_writer_line() {
292 let fake = VecWriter::new();
293 let mut writer = PrefixWriter::new(fake, "test: ".as_bytes().to_vec());
294 writer.write_all(b"foo\nbar").await.unwrap();
295 let fake = writer.into_inner();
296 let output = String::from_utf8(fake.data).unwrap();
297 assert_eq!(output, "test: foo\ntest: bar");
298 }
299
300 #[tokio::test]
301 async fn test_line_writer() {
302 let (tx, mut rx) = tokio::sync::mpsc::channel(4);
303 let mut lw = LineWriter::new(tx);
304 lw.write_all(b"foo\nba").await.unwrap();
305 lw.write_all(b"r\nbaz").await.unwrap();
306 lw.flush().await.unwrap();
307 drop(lw);
308 let mut lines = Vec::new();
309 while let Some(line) = rx.recv().await {
310 lines.push(String::from_utf8(line).unwrap());
311 }
312 assert_eq!(lines, vec!["foo\n", "bar\n", "baz"]);
313 }
314
315 #[tokio::test]
316 async fn test_prefix_multiple_lines() {
317 let fake = VecWriter::new();
318 let mut writer = PrefixWriter::new(fake, "prefix: ".as_bytes().to_vec());
319 writer.write_all(b"line1\nline2\nline3").await.unwrap();
320 let fake = writer.into_inner();
321 let output = String::from_utf8(fake.data).unwrap();
322 assert_eq!(output, "prefix: line1\nprefix: line2\nprefix: line3");
323 }
324
325 #[tokio::test]
326 async fn test_run() {
327 let writer = TerminalWriter::new(false, 4);
328 let mut tx_out = writer.stdout();
329 let mut tx_err = writer.stderr();
330 tx_out.write_all(b"OUT: hello\n").await.unwrap();
332 tx_err.write_all(b"ERR: world\n").await.unwrap();
333 drop(tx_out);
335 drop(tx_err);
336
337 let mut out_buf = VecWriter::new();
338 let mut err_buf = VecWriter::new();
339 writer.run_with(&mut out_buf, &mut err_buf).await;
340
341 assert_eq!(String::from_utf8(out_buf.data).unwrap(), "OUT: hello\n");
342 assert_eq!(String::from_utf8(err_buf.data).unwrap(), "ERR: world\n");
343 }
344
345 #[tokio::test]
346 async fn test_with_label() {
347 let writer = TerminalWriter::new(false, 4);
348 let mut tx_out = writer.stdout_with_label("OUT".to_string());
349
350 let mut cursor = {
351 let buf = b"hello\nworld\n".to_vec();
352 std::io::Cursor::new(buf)
353 };
354
355 let mut out_buf = VecWriter::new();
356 let mut err_buf = VecWriter::new();
357
358 tokio::join!(
359 async move {
360 tokio::io::copy(&mut cursor, &mut tx_out).await.unwrap();
361 drop(tx_out);
362 },
363 writer.run_with(&mut out_buf, &mut err_buf)
364 );
365
366 assert_eq!(
367 String::from_utf8(out_buf.data).unwrap(),
368 "OUT: hello\nOUT: world\n"
369 );
370 }
371}