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() {
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;
}
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>> {
self.buf.extend_from_slice(buf);
let mut start = 0;
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;
}
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<()>> {
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()
}
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"); }
prefix.extend_from_slice(label.as_bytes());
prefix.extend_from_slice(b": ");
if self.is_tty {
prefix.extend_from_slice(b"\x1b[0m"); }
PrefixWriter::new(stdout, prefix)
}
pub fn stderr_raw(&self) -> Sender<Vec<u8>> {
self.stderr_tx.clone()
}
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"); }
prefix.extend_from_slice(label.as_bytes());
prefix.extend_from_slice(b": ");
if self.is_tty {
prefix.extend_from_slice(b"\x1b[0m"); }
PrefixWriter::new(stderr, prefix)
}
pub async fn run(self) {
self.run_with(io::stdout(), io::stderr()).await;
}
pub async fn run_with<O, E>(self, mut out: O, mut err: E)
where
O: AsyncWrite + Unpin,
E: AsyncWrite + Unpin,
{
let TerminalWriter {
stdout_tx,
stderr_tx,
stdout_rx,
stderr_rx,
..
} = self;
drop(stdout_tx);
drop(stderr_tx);
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;
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();
tx_out.write_all(b"OUT: hello\n").await.unwrap();
tx_err.write_all(b"ERR: world\n").await.unwrap();
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"
);
}
}