wasmtime_wasi/cli/
worker_thread_stdin.rs1use crate::cli::{IsTerminal, StdinStream, stream_error_from};
27use bytes::{Bytes, BytesMut};
28use std::mem;
29use std::pin::Pin;
30use std::sync::{Condvar, Mutex, OnceLock};
31use std::task::{Context, Poll};
32use tokio::io::{self, AsyncRead, ReadBuf};
33use tokio::sync::Notify;
34use tokio::sync::futures::Notified;
35use wasmtime_wasi_io::{
36 poll::Pollable,
37 streams::{InputStream, StreamError},
38};
39
40use crate::MAX_READ_SIZE_ALLOC;
41
42impl IsTerminal for tokio::io::Stdin {
44 fn is_terminal(&self) -> bool {
45 std::io::stdin().is_terminal()
46 }
47}
48impl StdinStream for tokio::io::Stdin {
49 fn p2_stream(&self) -> Box<dyn InputStream> {
50 Box::new(WasiStdin)
51 }
52 fn async_stream(&self) -> Box<dyn AsyncRead + Send + Sync> {
53 Box::new(WasiStdinAsyncRead::Ready)
54 }
55}
56
57impl IsTerminal for std::io::Stdin {
59 fn is_terminal(&self) -> bool {
60 std::io::IsTerminal::is_terminal(self)
61 }
62}
63impl StdinStream for std::io::Stdin {
64 fn p2_stream(&self) -> Box<dyn InputStream> {
65 Box::new(WasiStdin)
66 }
67 fn async_stream(&self) -> Box<dyn AsyncRead + Send + Sync> {
68 Box::new(WasiStdinAsyncRead::Ready)
69 }
70}
71
72#[derive(Default)]
73struct GlobalStdin {
74 state: Mutex<StdinState>,
75 read_requested: Condvar,
76 read_completed: Notify,
77}
78
79#[derive(Default, Debug)]
80enum StdinState {
81 #[default]
82 ReadNotRequested,
83 ReadRequested(usize),
84 Data(BytesMut),
85 Error(std::io::Error),
86 Closed,
87}
88
89impl GlobalStdin {
90 fn get() -> &'static GlobalStdin {
91 static STDIN: OnceLock<GlobalStdin> = OnceLock::new();
92 STDIN.get_or_init(|| create())
93 }
94}
95
96fn create() -> GlobalStdin {
97 std::thread::spawn(|| {
98 let state = GlobalStdin::get();
99 loop {
100 let mut lock = state.state.lock().unwrap();
103 lock = state
104 .read_requested
105 .wait_while(lock, |state| !matches!(state, StdinState::ReadRequested(_)))
106 .unwrap();
107
108 let size_hint = match *lock {
112 StdinState::ReadRequested(size) => size.min(MAX_READ_SIZE_ALLOC).max(1),
113 _ => unreachable!(),
114 };
115 drop(lock);
116
117 let mut bytes = BytesMut::zeroed(size_hint);
118 let (new_state, done) = match read_stdin(&mut bytes) {
119 Ok(0) => (StdinState::Closed, true),
120 Ok(nbytes) => {
121 bytes.truncate(nbytes);
122 (StdinState::Data(bytes), false)
123 }
124 Err(e) => (StdinState::Error(e), true),
125 };
126
127 debug_assert!(matches!(
130 *state.state.lock().unwrap(),
131 StdinState::ReadRequested(_)
132 ));
133 let mut lock = state.state.lock().unwrap();
134 *lock = new_state;
135 state.read_completed.notify_waiters();
136 if done {
137 break;
138 }
139 }
140 });
141
142 GlobalStdin::default()
143}
144
145fn read_stdin(bytes: &mut [u8]) -> std::io::Result<usize> {
149 #[cfg(unix)]
150 {
151 use std::os::fd::AsFd;
152 let stdin = std::io::stdin();
153 let stdin = stdin.lock();
154 rustix::io::read(stdin.as_fd(), bytes).map_err(Into::into)
155 }
156
157 #[cfg(windows)]
158 {
159 use std::io::Read as _;
160 use std::os::windows::io::{AsRawHandle, FromRawHandle};
161
162 let stdin = std::io::stdin();
163 let mut stdin = stdin.lock();
164 if std::io::IsTerminal::is_terminal(&stdin) {
165 return stdin.read(bytes);
166 }
167
168 let mut file = std::mem::ManuallyDrop::new(unsafe {
171 std::fs::File::from_raw_handle(stdin.as_raw_handle())
172 });
173 file.read(bytes)
174 }
175
176 #[cfg(not(any(unix, windows)))]
177 {
178 use std::io::Read as _;
179 std::io::stdin().read(bytes)
180 }
181}
182
183struct WasiStdin;
184
185#[async_trait::async_trait]
186impl InputStream for WasiStdin {
187 fn read(&mut self, size: usize) -> Result<Bytes, StreamError> {
188 if size == 0 {
189 return Ok(Bytes::new());
190 }
191 let g = GlobalStdin::get();
192 let mut locked = g.state.lock().unwrap();
193 match mem::replace(&mut *locked, StdinState::ReadRequested(size)) {
194 StdinState::ReadNotRequested => {
195 g.read_requested.notify_one();
196 Ok(Bytes::new())
197 }
198 StdinState::ReadRequested(prev_size) => {
199 *locked = StdinState::ReadRequested(prev_size.max(size));
202 Ok(Bytes::new())
203 }
204 StdinState::Data(mut data) => {
205 let size = data.len().min(size);
206 let bytes = data.split_to(size);
207 *locked = if data.is_empty() {
208 StdinState::ReadNotRequested
209 } else {
210 StdinState::Data(data)
211 };
212 Ok(bytes.freeze())
213 }
214 StdinState::Error(e) => {
215 *locked = StdinState::Closed;
216 Err(stream_error_from(e))
217 }
218 StdinState::Closed => {
219 *locked = StdinState::Closed;
220 Err(StreamError::Closed)
221 }
222 }
223 }
224}
225
226#[async_trait::async_trait]
227impl Pollable for WasiStdin {
228 async fn ready(&mut self) {
229 let g = GlobalStdin::get();
230
231 let notified = {
234 let mut locked = g.state.lock().unwrap();
235 match *locked {
236 StdinState::ReadNotRequested => {
240 g.read_requested.notify_one();
241 *locked = StdinState::ReadRequested(MAX_READ_SIZE_ALLOC);
242 g.read_completed.notified()
243 }
244 StdinState::ReadRequested(_) => g.read_completed.notified(),
245 StdinState::Data(_) | StdinState::Closed | StdinState::Error(_) => return,
246 }
247 };
248
249 notified.await;
250 }
251}
252
253enum WasiStdinAsyncRead {
254 Ready,
255 Waiting(Notified<'static>),
256}
257
258impl AsyncRead for WasiStdinAsyncRead {
259 fn poll_read(
260 mut self: Pin<&mut Self>,
261 cx: &mut Context<'_>,
262 buf: &mut ReadBuf<'_>,
263 ) -> Poll<io::Result<()>> {
264 let g = GlobalStdin::get();
265
266 let mut locked = g.state.lock().unwrap();
273
274 loop {
278 if let Some(notified) = self.as_mut().notified_future() {
281 match notified.poll(cx) {
282 Poll::Ready(()) => self.set(WasiStdinAsyncRead::Ready),
283 Poll::Pending => break Poll::Pending,
284 }
285 }
286
287 assert!(matches!(*self, WasiStdinAsyncRead::Ready));
288
289 match mem::replace(&mut *locked, StdinState::ReadRequested(buf.remaining())) {
292 StdinState::Data(mut data) => {
294 let size = data.len().min(buf.remaining());
295 let bytes = data.split_to(size);
296 *locked = if data.is_empty() {
297 StdinState::ReadNotRequested
298 } else {
299 StdinState::Data(data)
300 };
301 buf.put_slice(&bytes);
302 break Poll::Ready(Ok(()));
303 }
304
305 StdinState::Error(e) => {
308 *locked = StdinState::Closed;
309 break Poll::Ready(Err(e));
310 }
311
312 StdinState::Closed => {
314 *locked = StdinState::Closed;
315 break Poll::Ready(Ok(()));
316 }
317
318 StdinState::ReadNotRequested => {
322 g.read_requested.notify_one();
323 }
324 StdinState::ReadRequested(prev_size) => {
325 *locked = StdinState::ReadRequested(prev_size.max(buf.remaining()));
327 }
328 }
329
330 self.set(WasiStdinAsyncRead::Waiting(g.read_completed.notified()));
331 }
332 }
333}
334
335impl WasiStdinAsyncRead {
336 fn notified_future(self: Pin<&mut Self>) -> Option<Pin<&mut Notified<'static>>> {
337 unsafe {
341 match self.get_unchecked_mut() {
342 WasiStdinAsyncRead::Ready => None,
343 WasiStdinAsyncRead::Waiting(notified) => Some(Pin::new_unchecked(notified)),
344 }
345 }
346 }
347}