Skip to main content

wasmtime_wasi/cli/
stdout.rs

1use crate::cli::{IsTerminal, StdoutStream, stream_error_from};
2use crate::p2;
3use bytes::Bytes;
4use std::io::{self, Write};
5use std::pin::Pin;
6use std::task::{Context, Poll};
7use tokio::io::AsyncWrite;
8use wasmtime_wasi_io::streams::OutputStream;
9
10// Implementation for tokio::io::Stdout
11impl IsTerminal for tokio::io::Stdout {
12    fn is_terminal(&self) -> bool {
13        std::io::stdout().is_terminal()
14    }
15}
16impl StdoutStream for tokio::io::Stdout {
17    fn p2_stream(&self) -> Box<dyn OutputStream> {
18        Box::new(StdioOutputStream::Stdout)
19    }
20    fn async_stream(&self) -> Box<dyn AsyncWrite + Send + Sync> {
21        Box::new(StdioOutputStream::Stdout)
22    }
23}
24
25// Implementation for std::io::Stdout
26impl IsTerminal for std::io::Stdout {
27    fn is_terminal(&self) -> bool {
28        std::io::IsTerminal::is_terminal(self)
29    }
30}
31impl StdoutStream for std::io::Stdout {
32    fn p2_stream(&self) -> Box<dyn OutputStream> {
33        Box::new(StdioOutputStream::Stdout)
34    }
35    fn async_stream(&self) -> Box<dyn AsyncWrite + Send + Sync> {
36        Box::new(StdioOutputStream::Stdout)
37    }
38}
39
40// Implementation for tokio::io::Stderr
41impl IsTerminal for tokio::io::Stderr {
42    fn is_terminal(&self) -> bool {
43        std::io::stderr().is_terminal()
44    }
45}
46impl StdoutStream for tokio::io::Stderr {
47    fn p2_stream(&self) -> Box<dyn OutputStream> {
48        Box::new(StdioOutputStream::Stderr)
49    }
50    fn async_stream(&self) -> Box<dyn AsyncWrite + Send + Sync> {
51        Box::new(StdioOutputStream::Stderr)
52    }
53}
54
55// Implementation for std::io::Stderr
56impl IsTerminal for std::io::Stderr {
57    fn is_terminal(&self) -> bool {
58        std::io::IsTerminal::is_terminal(self)
59    }
60}
61impl StdoutStream for std::io::Stderr {
62    fn p2_stream(&self) -> Box<dyn OutputStream> {
63        Box::new(StdioOutputStream::Stderr)
64    }
65    fn async_stream(&self) -> Box<dyn AsyncWrite + Send + Sync> {
66        Box::new(StdioOutputStream::Stderr)
67    }
68}
69
70enum StdioOutputStream {
71    Stdout,
72    Stderr,
73}
74
75/// The number of bytes a single `write` is permitted to carry, as reported by
76/// `check_write`.
77const WRITE_BUDGET: usize = 1024 * 1024;
78
79impl OutputStream for StdioOutputStream {
80    fn write(&mut self, bytes: Bytes) -> p2::StreamResult<()> {
81        if bytes.len() > WRITE_BUDGET {
82            return Err(p2::StreamError::Trap(wasmtime::format_err!(
83                "write exceeded budget"
84            )));
85        }
86        match self {
87            StdioOutputStream::Stdout => std::io::stdout().write_all(&bytes),
88            StdioOutputStream::Stderr => std::io::stderr().write_all(&bytes),
89        }
90        .map_err(|e| stream_error_from(e))
91    }
92
93    fn flush(&mut self) -> p2::StreamResult<()> {
94        match self {
95            StdioOutputStream::Stdout => std::io::stdout().flush(),
96            StdioOutputStream::Stderr => std::io::stderr().flush(),
97        }
98        .map_err(|e| stream_error_from(e))
99    }
100
101    fn check_write(&mut self) -> p2::StreamResult<usize> {
102        Ok(WRITE_BUDGET)
103    }
104}
105
106#[cfg(test)]
107mod tests {
108    use super::*;
109
110    /// A write larger than the budget `check_write` reports has to trap instead
111    /// of being written out.
112    #[test]
113    fn write_larger_than_budget_traps() {
114        let mut stream = StdioOutputStream::Stdout;
115        let err = stream
116            .write(Bytes::from(vec![0; WRITE_BUDGET + 1]))
117            .unwrap_err();
118        assert!(matches!(err, p2::StreamError::Trap(_)), "{err:?}");
119    }
120}
121
122impl AsyncWrite for StdioOutputStream {
123    fn poll_write(
124        self: Pin<&mut Self>,
125        _cx: &mut Context<'_>,
126        buf: &[u8],
127    ) -> Poll<io::Result<usize>> {
128        Poll::Ready(match *self {
129            StdioOutputStream::Stdout => std::io::stdout().write(buf),
130            StdioOutputStream::Stderr => std::io::stderr().write(buf),
131        })
132    }
133    fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
134        Poll::Ready(match *self {
135            StdioOutputStream::Stdout => std::io::stdout().flush(),
136            StdioOutputStream::Stderr => std::io::stderr().flush(),
137        })
138    }
139    fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
140        Poll::Ready(Ok(()))
141    }
142}
143
144#[async_trait::async_trait]
145impl p2::Pollable for StdioOutputStream {
146    async fn ready(&mut self) {}
147}