kumosh-core 0.1.1

Kumosh is a cluster shell
Documentation
use std::pin::Pin;
use std::task::{Context, Poll, ready};
use tokio::io::{self, AsyncWrite, AsyncWriteExt};
use tokio::sync::mpsc::{Receiver, Sender};

pub struct PrefixWriter<W: AsyncWrite + Unpin> {
    inner: W,
    prefix: Vec<u8>,
    at_start: bool,
}

impl<W: AsyncWrite + Unpin> PrefixWriter<W> {
    pub fn new(writer: W, prefix: Vec<u8>) -> Self {
        Self {
            inner: writer,
            prefix,
            at_start: true,
        }
    }

    pub fn into_inner(self) -> W {
        self.inner
    }
}

impl<W: AsyncWrite + Unpin> AsyncWrite for PrefixWriter<W> {
    fn poll_write(
        mut self: Pin<&mut Self>,
        cx: &mut Context<'_>,
        buf: &[u8],
    ) -> Poll<io::Result<usize>> {
        let mut written = 0;
        let mut offset = 0;
        while offset < buf.len() {
            // write prefix if at start of line
            if self.at_start {
                let prefix = self.prefix.clone();
                let n = ready!(Pin::new(&mut self.inner).poll_write(cx, &prefix))?;
                if n < self.prefix.len() {
                    return Poll::Ready(Ok(0));
                }
                self.at_start = false;
            }
            // find newline
            if let Some(pos) = buf[offset..].iter().position(|&b| b == b'\n') {
                let end = offset + pos + 1;
                let n = ready!(Pin::new(&mut self.inner).poll_write(cx, &buf[offset..end]))?;
                written += n;
                offset = end;
                self.at_start = true;
            } else {
                let n = ready!(Pin::new(&mut self.inner).poll_write(cx, &buf[offset..]))?;
                written += n;
                offset = buf.len();
            }
        }
        Poll::Ready(Ok(written))
    }

    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
        Pin::new(&mut self.inner).poll_flush(cx)
    }

    fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
        Pin::new(&mut self.inner).poll_shutdown(cx)
    }
}

pub struct WriteLine {
    pub label: String,
    pub line: Vec<u8>,
}

impl WriteLine {
    pub fn new(label: String, line: Vec<u8>) -> Self {
        Self { label, line }
    }
}

pub struct LineWriter {
    tx: Sender<Vec<u8>>,
    buf: Vec<u8>,
}

impl LineWriter {
    pub fn new(tx: Sender<Vec<u8>>) -> Self {
        Self {
            tx,
            buf: Vec::new(),
        }
    }

    pub fn into_inner(self) -> Sender<Vec<u8>> {
        self.tx
    }
}

impl AsyncWrite for LineWriter {
    fn poll_write(
        mut self: Pin<&mut Self>,
        _cx: &mut Context<'_>,
        buf: &[u8],
    ) -> Poll<io::Result<usize>> {
        // Append to internal buffer
        self.buf.extend_from_slice(buf);
        let mut start = 0;
        // Extract complete lines
        while let Some(pos) = self.buf[start..].iter().position(|&b| b == b'\n') {
            let end = start + pos + 1;
            let line = self.buf[..end].to_vec();
            let _ = self.tx.try_send(line);
            start = end;
        }
        // Remove sent lines
        if start > 0 {
            self.buf.drain(..start);
        }
        Poll::Ready(Ok(buf.len()))
    }

    fn poll_flush(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
        if !self.buf.is_empty() {
            let remaining = self.buf.split_off(0);
            let _ = self.tx.try_send(remaining);
        }
        Poll::Ready(Ok(()))
    }

    fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
        // Flush any pending data
        let _ = self.poll_flush(cx);
        Poll::Ready(Ok(()))
    }
}

pub struct TerminalWriter {
    is_tty: bool,
    stdout_tx: Sender<Vec<u8>>,
    stdout_rx: Receiver<Vec<u8>>,
    stderr_tx: Sender<Vec<u8>>,
    stderr_rx: Receiver<Vec<u8>>,
}

impl Default for TerminalWriter {
    fn default() -> Self {
        Self::new(true, 1024)
    }
}

impl TerminalWriter {
    pub fn new(is_tty: bool, buffer: usize) -> Self {
        let (stdout_tx, stdout_rx) = tokio::sync::mpsc::channel(buffer);
        let (stderr_tx, stderr_rx) = tokio::sync::mpsc::channel(buffer);
        Self {
            is_tty,
            stdout_tx,
            stdout_rx,
            stderr_tx,
            stderr_rx,
        }
    }

    pub fn stdout_raw(&self) -> Sender<Vec<u8>> {
        self.stdout_tx.clone()
    }

    /// Cloneable sender for stdout lines.
    pub fn stdout(&self) -> LineWriter {
        LineWriter::new(self.stdout_tx.clone())
    }

    pub fn stdout_with_label(&self, label: String) -> PrefixWriter<LineWriter> {
        let stdout = self.stdout();
        let mut prefix = Vec::new();
        if self.is_tty {
            prefix.extend_from_slice(b"\x1b[1;34m"); // Blue color for stdout
        }
        prefix.extend_from_slice(label.as_bytes());
        prefix.extend_from_slice(b": ");
        if self.is_tty {
            prefix.extend_from_slice(b"\x1b[0m"); // Reset color
        }

        PrefixWriter::new(stdout, prefix)
    }

    pub fn stderr_raw(&self) -> Sender<Vec<u8>> {
        self.stderr_tx.clone()
    }

    /// Cloneable sender for stderr lines.
    pub fn stderr(&self) -> LineWriter {
        LineWriter::new(self.stderr_tx.clone())
    }

    pub fn stderr_with_label(&self, label: String) -> PrefixWriter<LineWriter> {
        let stderr = self.stderr();
        let mut prefix = Vec::new();
        if self.is_tty {
            prefix.extend_from_slice(b"\x1b[1;31m"); // Red color for stderr
        }
        prefix.extend_from_slice(label.as_bytes());
        prefix.extend_from_slice(b": ");
        if self.is_tty {
            prefix.extend_from_slice(b"\x1b[0m"); // Reset color
        }

        PrefixWriter::new(stderr, prefix)
    }

    /// Run writer with real stdout/stderr
    pub async fn run(self) {
        self.run_with(io::stdout(), io::stderr()).await;
    }

    /// Run writer writing to provided out/err AsyncWrite impls.
    pub async fn run_with<O, E>(self, mut out: O, mut err: E)
    where
        O: AsyncWrite + Unpin,
        E: AsyncWrite + Unpin,
    {
        // Destructure to get both senders and receivers
        let TerminalWriter {
            stdout_tx,
            stderr_tx,
            stdout_rx,
            stderr_rx,
            ..
        } = self;

        // Close the channel so that recv() sees None when all buffered messages drained
        drop(stdout_tx);
        drop(stderr_tx);

        // Prepare futures for stdout and stderr processing
        let stdout_fut = async {
            let mut rx = stdout_rx;
            while let Some(fragment) = rx.recv().await {
                let _ = out.write_all(&fragment).await;
            }
        };

        let stderr_fut = async {
            let mut rx = stderr_rx;
            while let Some(fragment) = rx.recv().await {
                let _ = err.write_all(&fragment).await;
            }
        };

        tokio::join!(stdout_fut, stderr_fut);
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use std::pin::Pin;
    use std::task::{Context, Poll};
    use tokio::io;
    use tokio::io::AsyncWriteExt;

    /// A simple in-memory writer to capture output.
    struct VecWriter {
        pub data: Vec<u8>,
    }
    impl VecWriter {
        fn new() -> Self {
            VecWriter { data: Vec::new() }
        }
    }
    impl io::AsyncWrite for VecWriter {
        fn poll_write(
            self: Pin<&mut Self>,
            _cx: &mut Context<'_>,
            buf: &[u8],
        ) -> Poll<io::Result<usize>> {
            let this = self.get_mut();
            this.data.extend_from_slice(buf);
            Poll::Ready(Ok(buf.len()))
        }
        fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
            Poll::Ready(Ok(()))
        }
        fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
            Poll::Ready(Ok(()))
        }
    }
    impl Unpin for VecWriter {}

    #[tokio::test]
    async fn test_prefix_writer_line() {
        let fake = VecWriter::new();
        let mut writer = PrefixWriter::new(fake, "test: ".as_bytes().to_vec());
        writer.write_all(b"foo\nbar").await.unwrap();
        let fake = writer.into_inner();
        let output = String::from_utf8(fake.data).unwrap();
        assert_eq!(output, "test: foo\ntest: bar");
    }

    #[tokio::test]
    async fn test_line_writer() {
        let (tx, mut rx) = tokio::sync::mpsc::channel(4);
        let mut lw = LineWriter::new(tx);
        lw.write_all(b"foo\nba").await.unwrap();
        lw.write_all(b"r\nbaz").await.unwrap();
        lw.flush().await.unwrap();
        drop(lw);
        let mut lines = Vec::new();
        while let Some(line) = rx.recv().await {
            lines.push(String::from_utf8(line).unwrap());
        }
        assert_eq!(lines, vec!["foo\n", "bar\n", "baz"]);
    }

    #[tokio::test]
    async fn test_prefix_multiple_lines() {
        let fake = VecWriter::new();
        let mut writer = PrefixWriter::new(fake, "prefix: ".as_bytes().to_vec());
        writer.write_all(b"line1\nline2\nline3").await.unwrap();
        let fake = writer.into_inner();
        let output = String::from_utf8(fake.data).unwrap();
        assert_eq!(output, "prefix: line1\nprefix: line2\nprefix: line3");
    }

    #[tokio::test]
    async fn test_run() {
        let writer = TerminalWriter::new(false, 4);
        let mut tx_out = writer.stdout();
        let mut tx_err = writer.stderr();
        // send fragments
        tx_out.write_all(b"OUT: hello\n").await.unwrap();
        tx_err.write_all(b"ERR: world\n").await.unwrap();
        // close channels
        drop(tx_out);
        drop(tx_err);

        let mut out_buf = VecWriter::new();
        let mut err_buf = VecWriter::new();
        writer.run_with(&mut out_buf, &mut err_buf).await;

        assert_eq!(String::from_utf8(out_buf.data).unwrap(), "OUT: hello\n");
        assert_eq!(String::from_utf8(err_buf.data).unwrap(), "ERR: world\n");
    }

    #[tokio::test]
    async fn test_with_label() {
        let writer = TerminalWriter::new(false, 4);
        let mut tx_out = writer.stdout_with_label("OUT".to_string());

        let mut cursor = {
            let buf = b"hello\nworld\n".to_vec();
            std::io::Cursor::new(buf)
        };

        let mut out_buf = VecWriter::new();
        let mut err_buf = VecWriter::new();

        tokio::join!(
            async move {
                tokio::io::copy(&mut cursor, &mut tx_out).await.unwrap();
                drop(tx_out);
            },
            writer.run_with(&mut out_buf, &mut err_buf)
        );

        assert_eq!(
            String::from_utf8(out_buf.data).unwrap(),
            "OUT: hello\nOUT: world\n"
        );
    }
}