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
18pub 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#[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 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 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 WriteState::Ready(..) => WriteState::Closed(WriteError::Closed),
233
234 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 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}