Skip to main content

kumosh_core/
io.rs

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            // write prefix if at start of line
36            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            // find newline
45            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        // Append to internal buffer
105        self.buf.extend_from_slice(buf);
106        let mut start = 0;
107        // Extract complete lines
108        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        // Remove sent lines
115        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        // Flush any pending data
131        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    /// Cloneable sender for stdout lines.
168    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"); // Blue color for stdout
177        }
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"); // Reset color
182        }
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    /// Cloneable sender for stderr lines.
192    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"); // Red color for stderr
201        }
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"); // Reset color
206        }
207
208        PrefixWriter::new(stderr, prefix)
209    }
210
211    /// Run writer with real stdout/stderr
212    pub async fn run(self) {
213        self.run_with(io::stdout(), io::stderr()).await;
214    }
215
216    /// Run writer writing to provided out/err AsyncWrite impls.
217    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        // Destructure to get both senders and receivers
223        let TerminalWriter {
224            stdout_tx,
225            stderr_tx,
226            stdout_rx,
227            stderr_rx,
228            ..
229        } = self;
230
231        // Close the channel so that recv() sees None when all buffered messages drained
232        drop(stdout_tx);
233        drop(stderr_tx);
234
235        // Prepare futures for stdout and stderr processing
236        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    /// A simple in-memory writer to capture output.
263    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        // send fragments
331        tx_out.write_all(b"OUT: hello\n").await.unwrap();
332        tx_err.write_all(b"ERR: world\n").await.unwrap();
333        // close channels
334        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}