Skip to main content

wasmtime_wasi/p2/
tcp.rs

1use crate::MAX_READ_SIZE_ALLOC;
2use crate::p2::bindings::sockets::network::ErrorCode;
3use crate::p2::{
4    DynInputStream, DynOutputStream, InputStream, OutputStream, Pollable, SocketResult, StreamError,
5};
6use crate::sockets::{
7    MaybeReady, TcpListenStream, TcpReceiveStream, TcpSendStream, TcpSocket as P3Socket, noop_cx,
8};
9use std::future::poll_fn;
10use std::mem;
11use std::net::Shutdown;
12use std::sync::Arc;
13use std::sync::Mutex;
14use std::task::{Poll, ready};
15use wasmtime::Result;
16use wasmtime_wasi_io::streams::StreamResult;
17
18/// A TCP socket + associated p2 bookkeeping.
19pub struct TcpSocket {
20    pub(crate) inner: P3Socket,
21    pub(crate) in_progress_operation: Option<AsyncOperation>,
22    pub(crate) listener: Option<TcpListenStream>,
23    reader: Option<TcpReader>,
24    writer: Option<TcpWriter>,
25}
26
27#[derive(Clone, Copy, Debug, PartialEq, Eq)]
28pub(crate) enum AsyncOperation {
29    Bind,
30    Connect,
31    Listen,
32}
33
34impl TcpSocket {
35    pub(crate) fn new(inner: P3Socket) -> Self {
36        Self {
37            inner,
38            in_progress_operation: None,
39            listener: None,
40            reader: None,
41            writer: None,
42        }
43    }
44    pub(crate) fn take_streams(&mut self) -> SocketResult<(DynInputStream, DynOutputStream)> {
45        let reader = TcpReader::new(self.inner.take_receive_stream()?);
46        let writer = TcpWriter::new(self.inner.take_send_stream()?);
47        self.reader = Some(reader.clone());
48        self.writer = Some(writer.clone());
49        let input: DynInputStream = Box::new(reader);
50        let output: DynOutputStream = Box::new(writer);
51        Ok((input, output))
52    }
53    pub(crate) fn shutdown(&mut self, how: Shutdown) -> SocketResult<()> {
54        let reader = self.reader.as_mut().ok_or(ErrorCode::InvalidState)?;
55        let writer = self.writer.as_mut().ok_or(ErrorCode::InvalidState)?;
56
57        if let Shutdown::Both | Shutdown::Read = how {
58            reader.0.lock().unwrap().shutdown();
59        }
60
61        if let Shutdown::Both | Shutdown::Write = how {
62            writer.0.lock().unwrap().shutdown();
63        }
64
65        Ok(())
66    }
67}
68
69enum ReadState {
70    Open(TcpReceiveStream),
71    Closed,
72}
73impl ReadState {
74    fn read(&mut self, size: usize) -> StreamResult<bytes::Bytes> {
75        let Self::Open(stream) = self else {
76            return Err(StreamError::Closed);
77        };
78        if size == 0 {
79            return Ok(bytes::Bytes::new());
80        }
81        let mut buf = bytes::BytesMut::zeroed(size.min(crate::MAX_READ_SIZE_ALLOC));
82        let n = match stream.poll_read(&mut noop_cx(), &mut buf) {
83            Poll::Pending => 0,
84            Poll::Ready(Ok(0)) => {
85                *self = ReadState::Closed;
86                return Err(StreamError::Closed);
87            }
88            Poll::Ready(Ok(n)) => n,
89            Poll::Ready(Err(e)) => {
90                *self = ReadState::Closed;
91                return Err(StreamError::LastOperationFailed(e.into()));
92            }
93        };
94
95        buf.truncate(n);
96        Ok(buf.freeze())
97    }
98
99    fn shutdown(&mut self) {
100        *self = ReadState::Closed;
101    }
102
103    fn poll_ready(&mut self, cx: &mut std::task::Context<'_>) -> Poll<()> {
104        match self {
105            Self::Open(stream) => stream.poll_ready(cx),
106            Self::Closed => Poll::Ready(()),
107        }
108    }
109}
110
111#[derive(Clone)]
112struct TcpReader(Arc<Mutex<ReadState>>);
113impl TcpReader {
114    fn new(stream: TcpReceiveStream) -> Self {
115        Self(Arc::new(Mutex::new(ReadState::Open(stream))))
116    }
117}
118
119#[async_trait::async_trait]
120impl InputStream for TcpReader {
121    fn read(&mut self, size: usize) -> StreamResult<bytes::Bytes> {
122        self.0.lock().unwrap().read(size)
123    }
124}
125
126#[async_trait::async_trait]
127impl Pollable for TcpReader {
128    async fn ready(&mut self) {
129        std::future::poll_fn(|cx| self.0.lock().unwrap().poll_ready(cx)).await
130    }
131}
132
133/// A cloneable subset of StreamError
134#[derive(Debug, Clone)]
135enum WriteError {
136    Closed,
137    LastOperationFailed(ErrorCode),
138}
139impl From<WriteError> for StreamError {
140    fn from(err: WriteError) -> Self {
141        match err {
142            WriteError::Closed => StreamError::Closed,
143            WriteError::LastOperationFailed(e) => StreamError::LastOperationFailed(e.into()),
144        }
145    }
146}
147
148enum WriteState {
149    Ready(TcpSendStream, usize),
150    Writing(MaybeReady<Result<TcpSendStream, WriteError>>),
151    Closing(MaybeReady<Result<(), WriteError>>),
152    Closed(WriteError),
153}
154
155impl WriteState {
156    fn take(&mut self) -> WriteState {
157        mem::replace(self, WriteState::Closed(WriteError::Closed))
158    }
159
160    fn check_write(&mut self) -> StreamResult<usize> {
161        match self.poll_ready(&mut noop_cx()) {
162            Poll::Pending => Ok(0),
163            Poll::Ready(Ok((_, permit))) => {
164                *permit = MAX_READ_SIZE_ALLOC;
165                Ok(*permit)
166            }
167            Poll::Ready(Err(e)) => Err(e),
168        }
169    }
170
171    fn write(&mut self, mut bytes: bytes::Bytes) -> StreamResult<()> {
172        let mut stream = match self {
173            WriteState::Ready(_, permit) if bytes.len() <= *permit => {
174                if bytes.is_empty() {
175                    return Ok(());
176                }
177
178                let WriteState::Ready(stream, _) = self.take() else {
179                    unreachable!()
180                };
181                stream
182            }
183            WriteState::Closed(e) => {
184                return Err(e.clone().into());
185            }
186            _ => {
187                return Err(StreamError::Trap(wasmtime::format_err!(
188                    "not permitted to write {} bytes",
189                    bytes.len()
190                )));
191            }
192        };
193
194        *self = WriteState::Writing(MaybeReady::poll_or_spawn(async move {
195            while !bytes.is_empty() {
196                match stream.write(&bytes).await {
197                    Ok(n) => {
198                        let _ = bytes.split_to(n);
199                    }
200                    Err(crate::sockets::ErrorCode::ConnectionBroken) => {
201                        return Err(WriteError::Closed);
202                    }
203                    Err(e) => {
204                        return Err(WriteError::LastOperationFailed(e.into()));
205                    }
206                }
207            }
208
209            Ok(stream)
210        }));
211
212        // Attempt to finish the write, surfacing potential errors immediately:
213        match self.poll_ready(&mut noop_cx()) {
214            Poll::Pending | Poll::Ready(Ok(_)) => Ok(()),
215            Poll::Ready(Err(e)) => Err(e),
216        }
217    }
218
219    fn flush(&mut self) -> StreamResult<()> {
220        // `flush` is a no-op here. Writes happen on background tasks and will
221        // always be delivered to the OS as soon as possible. There's nothing
222        // for `flush` to do here that will speed up that process.
223        match self {
224            WriteState::Ready(..) | WriteState::Writing(_) | WriteState::Closing(_) => Ok(()),
225            WriteState::Closed(e) => Err(e.clone().into()),
226        }
227    }
228
229    pub(crate) fn shutdown(&mut self) {
230        *self = match self.take() {
231            // No write in progress, immediately drop the inner stream:
232            WriteState::Ready(..) => WriteState::Closed(WriteError::Closed),
233
234            // Schedule the shutdown after the current write has finished:
235            WriteState::Writing(write) => {
236                WriteState::Closing(MaybeReady::poll_or_spawn(async move {
237                    _ = write.into_future().await?;
238                    Ok(())
239                }))
240            }
241
242            s => s,
243        };
244    }
245
246    fn poll_ready(
247        &mut self,
248        cx: &mut std::task::Context<'_>,
249    ) -> Poll<StreamResult<(&mut TcpSendStream, &mut usize)>> {
250        match self {
251            WriteState::Writing(write) => {
252                ready!(write.poll_ready(cx));
253                let WriteState::Writing(write) = self.take() else {
254                    unreachable!()
255                };
256                *self = match write.unwrap_ready() {
257                    Ok(stream) => WriteState::Ready(stream, 0),
258                    Err(err) => WriteState::Closed(err),
259                };
260            }
261            WriteState::Closing(close) => {
262                ready!(close.poll_ready(cx));
263                let WriteState::Closing(close) = self.take() else {
264                    unreachable!()
265                };
266                *self = match close.unwrap_ready() {
267                    Ok(()) => WriteState::Closed(WriteError::Closed),
268                    Err(err) => WriteState::Closed(err),
269                };
270            }
271            _ => {}
272        }
273
274        match self {
275            WriteState::Ready(stream, permit) => match stream.poll_ready(cx) {
276                Poll::Ready(()) => Poll::Ready(Ok((stream, permit))),
277                Poll::Pending => Poll::Pending,
278            },
279            WriteState::Writing(..) | WriteState::Closing(..) => Poll::Pending,
280            WriteState::Closed(e) => Poll::Ready(Err(e.clone().into())),
281        }
282    }
283}
284
285#[derive(Clone)]
286struct TcpWriter(Arc<Mutex<WriteState>>);
287impl TcpWriter {
288    fn new(stream: TcpSendStream) -> Self {
289        Self(Arc::new(Mutex::new(WriteState::Ready(stream, 0))))
290    }
291}
292
293#[async_trait::async_trait]
294impl OutputStream for TcpWriter {
295    fn write(&mut self, bytes: bytes::Bytes) -> StreamResult<()> {
296        self.0.lock().unwrap().write(bytes)
297    }
298
299    fn flush(&mut self) -> StreamResult<()> {
300        self.0.lock().unwrap().flush()
301    }
302
303    fn check_write(&mut self) -> StreamResult<usize> {
304        self.0.lock().unwrap().check_write()
305    }
306
307    async fn cancel(&mut self) {
308        // Wait for background writes to finish in order to prevent silently
309        // dropping data that (from the guest's perspective) was already written.
310        self.ready().await
311    }
312}
313
314#[async_trait::async_trait]
315impl Pollable for TcpWriter {
316    async fn ready(&mut self) {
317        poll_fn(|cx| self.0.lock().unwrap().poll_ready(cx).map(|_| ())).await;
318    }
319}