Skip to main content

wasmtime_wasi/sockets/
tcp.rs

1use crate::runtime::with_ambient_tokio_runtime;
2use crate::sockets::{
3    ErrorCode, MaybeReady, SocketAddrCheck, SocketAddrUse, SocketAddressFamily, WasiSocketsCtx,
4    get_receive_buffer_size, get_send_buffer_size, get_unicast_hop_limit, is_valid_address_family,
5    is_valid_remote_address, is_valid_unicast_address, set_receive_buffer_size,
6    set_send_buffer_size, set_unicast_hop_limit, unspecified_addr,
7};
8use rustix::fd::AsFd;
9use rustix::io::Errno;
10use rustix::net::sockopt;
11use std::fmt::Debug;
12use std::future::poll_fn;
13use std::mem;
14use std::net::SocketAddr;
15use std::sync::Arc;
16use std::task::{Poll, ready};
17use std::time::Duration;
18
19/// Value taken from rust std library.
20const DEFAULT_BACKLOG: u32 = 128;
21
22const NANOS_PER_SEC: u64 = 1_000_000_000;
23
24/// The state of a TCP socket.
25///
26/// This represents the various states a socket can be in during the
27/// activities of listening, accepting, and connecting.
28enum TcpState {
29    /// The initial state for a newly-created socket.
30    ///
31    /// The socket may be bound to a local address in this state, but doesn't
32    /// have to.
33    ///
34    /// From here a socket can transition to `Listening` or `Connecting`.
35    Default(tokio::net::TcpSocket),
36
37    /// The socket is now listening and waiting for an incoming connection.
38    ///
39    /// Sockets will not leave this state.
40    Listening(Arc<tokio::net::TcpListener>),
41
42    /// An outgoing connection is started.
43    ///
44    /// This is created via the `start_connect` method. The payload is a future
45    /// for the eventual result of the connect.
46    ///
47    /// From here a socket can transition to `Connected` or `Closed`.
48    Connecting(MaybeReady<Result<tokio::net::TcpStream, ErrorCode>>),
49
50    /// A connection has been established.
51    ///
52    /// This is created either via `finish_connect` or for freshly accepted
53    /// sockets from a TCP listener.
54    ///
55    /// A socket will not transition out of this state.
56    Connected {
57        stream: Arc<tokio::net::TcpStream>,
58        /// Cached peer address, returned by `accept` or the first successful
59        /// `peer_addr` query. The stream may no longer report it after a reset.
60        peer: Option<SocketAddr>,
61        receive_taken: bool,
62        send_taken: bool,
63    },
64
65    /// The socket is closed and no more operations can be performed.
66    Closed(ErrorCode),
67}
68impl TcpState {
69    fn connected(stream: tokio::net::TcpStream, peer: Option<SocketAddr>) -> Self {
70        TcpState::Connected {
71            stream: Arc::new(stream),
72            peer,
73            receive_taken: false,
74            send_taken: false,
75        }
76    }
77    fn take(&mut self) -> Self {
78        mem::replace(self, TcpState::Closed(ErrorCode::Other))
79    }
80}
81impl Debug for TcpState {
82    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
83        match self {
84            Self::Default(_) => f.debug_tuple("Default").finish(),
85            Self::Listening { .. } => f.debug_tuple("Listening").finish(),
86            Self::Connecting(..) => f.debug_tuple("Connecting").finish(),
87            Self::Connected { .. } => f.debug_tuple("Connected").finish(),
88            Self::Closed(..) => write!(f, "Closed"),
89        }
90    }
91}
92
93/// A host TCP socket, plus associated bookkeeping.
94pub struct TcpSocket {
95    /// The current state in the bind/listen/accept/connect progression.
96    tcp_state: TcpState,
97
98    /// The desired listen queue size.
99    listen_backlog_size: u32,
100
101    family: SocketAddressFamily,
102
103    /// The checks to perform before doing any noteworthy syscall.
104    permissions: SocketAddrCheck,
105
106    /// Persisted socket options to manually apply to newly accepted client
107    /// sockets on platforms that don't inherit socket options from the listener.
108    listener_options: NonInheritedOptions,
109
110    /// Cached value of whether the socket is bound. Various methods use the
111    /// `.is_bound()` method, so we cache it to avoid redundant syscalls.
112    is_bound: bool,
113}
114
115impl TcpSocket {
116    /// Create a new socket in the given family.
117    pub(crate) fn new(
118        ctx: &WasiSocketsCtx,
119        family: SocketAddressFamily,
120    ) -> Result<Self, ErrorCode> {
121        ctx.allowed_network_uses.check_allowed_tcp()?;
122
123        let socket = with_ambient_tokio_runtime(|| socket(family))?;
124
125        Ok(Self {
126            tcp_state: TcpState::Default(socket),
127            listen_backlog_size: DEFAULT_BACKLOG,
128            family,
129            is_bound: false,
130            listener_options: Default::default(),
131            permissions: ctx.socket_addr_check.clone(),
132        })
133    }
134
135    fn as_fd(&self) -> Result<rustix::fd::BorrowedFd<'_>, ErrorCode> {
136        match &self.tcp_state {
137            TcpState::Default(socket) => Ok(socket.as_fd()),
138            TcpState::Connected { stream, .. } => Ok(stream.as_fd()),
139            TcpState::Listening(listener) => Ok(listener.as_fd()),
140            TcpState::Connecting(..) => Err(ErrorCode::InvalidState),
141            TcpState::Closed(err) => Err(*err),
142        }
143    }
144
145    pub(crate) fn is_bound(&mut self) -> bool {
146        // Once bound, a TCP socket can never become unbound again. So we can
147        // skip all work after a previous call has already determined the
148        // socket to be bound.
149        if !self.is_bound {
150            self.is_bound = match &self.tcp_state {
151                TcpState::Default(socket) => socket
152                    .local_addr()
153                    .is_ok_and(|addr| addr != unspecified_addr(self.family)),
154                _ => true,
155            };
156        }
157        self.is_bound
158    }
159
160    pub(crate) async fn bind(&mut self, addr: SocketAddr) -> Result<(), ErrorCode> {
161        if self.is_bound() {
162            return Err(ErrorCode::InvalidState);
163        }
164        let TcpState::Default(sock) = &self.tcp_state else {
165            return Err(ErrorCode::InvalidState);
166        };
167
168        if !is_valid_unicast_address(addr.ip()) || !is_valid_address_family(addr.ip(), self.family)
169        {
170            return Err(ErrorCode::InvalidArgument);
171        }
172
173        self.permissions.check(addr, SocketAddrUse::TcpBind).await?;
174        bind(sock, addr)?;
175        Ok(())
176    }
177
178    pub(crate) fn start_connect(&mut self, addr: SocketAddr) -> Result<(), ErrorCode> {
179        let TcpState::Default(_) = &self.tcp_state else {
180            return Err(ErrorCode::InvalidState);
181        };
182
183        let permissions = self.permissions.clone();
184        let family = self.family;
185        let already_bound = self.is_bound();
186
187        if !is_valid_unicast_address(addr.ip())
188            || !is_valid_remote_address(addr)
189            || !is_valid_address_family(addr.ip(), family)
190        {
191            return Err(ErrorCode::InvalidArgument);
192        };
193
194        let TcpState::Default(sock) = self.tcp_state.take() else {
195            unreachable!();
196        };
197
198        self.tcp_state = TcpState::Connecting(MaybeReady::new(Box::pin(async move {
199            // Perform all checks before doing any syscalls.
200            {
201                if !already_bound {
202                    // If not explicitly bound, the OS will implicitly bind the
203                    // socket to an ephemeral port when connecting. Unlike other
204                    // operations (e.g. `listen`), we will *not* do the implicit
205                    // bind ourselves because that may accelerate port exhaustion.
206                    // For more info, see IP_BIND_ADDRESS_NO_PORT (Linux) or
207                    // SO_REUSE_UNICASTPORT (Windows).
208                    //
209                    // Instead we check the permission to bind, but not perform
210                    // the actual bind:
211                    let implicit = unspecified_addr(family);
212                    permissions.check(implicit, SocketAddrUse::TcpBind).await?;
213                }
214
215                permissions.check(addr, SocketAddrUse::TcpConnect).await?;
216            }
217
218            let stream = sock.connect(addr).await?;
219            Ok(stream)
220        })));
221
222        Ok(())
223    }
224
225    pub(crate) fn poll_finish_connect(
226        &mut self,
227        cx: &mut std::task::Context<'_>,
228    ) -> Poll<Result<(), ErrorCode>> {
229        match &mut self.tcp_state {
230            TcpState::Connecting(connect) => {
231                ready!(with_ambient_tokio_runtime(|| connect.poll_ready(cx)));
232            }
233            TcpState::Connected { .. } => return Poll::Ready(Ok(())),
234            TcpState::Closed(e) => return Poll::Ready(Err(*e)),
235            _ => return Poll::Ready(Err(ErrorCode::InvalidState)),
236        }
237        let TcpState::Connecting(connect) = self.tcp_state.take() else {
238            unreachable!();
239        };
240
241        match connect.unwrap_ready() {
242            Ok(stream) => {
243                self.tcp_state = TcpState::connected(stream, None);
244                Poll::Ready(Ok(()))
245            }
246            Err(err) => {
247                self.tcp_state = TcpState::Closed(err);
248                Poll::Ready(Err(err))
249            }
250        }
251    }
252
253    pub(crate) async fn listen(&mut self) -> Result<TcpListenStream, ErrorCode> {
254        let already_bound = self.is_bound();
255        let sock = match self.tcp_state.take() {
256            TcpState::Default(sock) => sock,
257            tcp_state => {
258                self.tcp_state = tcp_state;
259                return Err(ErrorCode::InvalidState);
260            }
261        };
262
263        // Perform all checks before doing any syscalls.
264        {
265            if already_bound {
266                self.permissions
267                    .check(sock.local_addr()?, SocketAddrUse::TcpListen)
268                    .await?;
269            } else {
270                let implicit = unspecified_addr(self.family);
271                self.permissions
272                    .check(implicit, SocketAddrUse::TcpBind)
273                    .await?;
274                self.permissions
275                    .check(implicit, SocketAddrUse::TcpListen)
276                    .await?;
277            }
278        }
279
280        // Some platforms automatically perform an implicit bind as part of
281        // the `listen` syscall. However this is not ubiquitous behavior:
282        // - Linux mentions it in their docs [0] that they perform an
283        //   implicit bind. This behavior has been experimentally verified.
284        // - Windows requires a `bind` before `listen`. This is both
285        //   documented [1] and experimentally verified.
286        // - Other platforms (e.g. macOS, FreeBSD) do not explicitly
287        //   document it either way and instead leave it up to the
288        //   individual protocol to decide [2]. However, experiments
289        //   show that MacOS in fact _does_ perform an implicit bind.
290        //
291        // Thus to ensure consistent behavior across all platforms, we
292        // perform the implicit bind ourselves here for unbound sockets.
293        //
294        // [0]: https://man7.org/linux/man-pages/man7/ip.7.html
295        // > An ephemeral port is allocated to a socket in the following
296        // > circumstances: (...) listen(2) is called on a stream socket
297        // > that was not previously bound;
298        //
299        // [1]: https://learn.microsoft.com/en-us/windows/win32/api/winsock2/nf-winsock2-listen
300        // > WSAEINVAL: The socket has not been bound with bind.
301        //
302        // [2]: https://pubs.opengroup.org/onlinepubs/9699919799/functions/listen.html
303        // > EDESTADDRREQ: The socket is not bound to a local address,
304        // > and the protocol does not support listening on an unbound
305        // > socket.
306        if !already_bound {
307            let implicit = unspecified_addr(self.family);
308            bind(&sock, implicit)?;
309        }
310
311        let listener = sock.listen(self.listen_backlog_size).map_err(|err| {
312            match Errno::from_io_error(&err) {
313                // See: https://learn.microsoft.com/en-us/windows/win32/api/winsock2/nf-winsock2-listen#:~:text=WSAEMFILE
314                // According to the docs, `listen` can return EMFILE on Windows.
315                // This is odd, because we're not trying to create a new socket
316                // or file descriptor of any kind. So we rewrite it to less
317                // surprising error code.
318                //
319                // At the time of writing, this behavior has never been experimentally
320                // observed by any of the wasmtime authors, so we're relying fully
321                // on Microsoft's documentation here.
322                #[cfg(windows)]
323                Some(Errno::MFILE) => Errno::NOBUFS.into(),
324
325                _ => err,
326            }
327        })?;
328        let listener = Arc::new(listener);
329        self.tcp_state = TcpState::Listening(listener.clone());
330
331        Ok(TcpListenStream {
332            inner: listener,
333            listener_options: self.listener_options.clone(),
334            family: self.family,
335            permissions: self.permissions.clone(),
336            pending_accept: None,
337        })
338    }
339
340    pub(crate) fn take_send_stream(&mut self) -> Result<TcpSendStream, ErrorCode> {
341        match &mut self.tcp_state {
342            TcpState::Connected {
343                stream, send_taken, ..
344            } if !*send_taken => {
345                *send_taken = true;
346                Ok(TcpSendStream {
347                    inner: stream.clone(),
348                })
349            }
350            TcpState::Closed(err) => Err(*err),
351            _ => Err(ErrorCode::InvalidState),
352        }
353    }
354
355    pub(crate) fn take_receive_stream(&mut self) -> Result<TcpReceiveStream, ErrorCode> {
356        match &mut self.tcp_state {
357            TcpState::Connected {
358                stream,
359                receive_taken,
360                ..
361            } if !*receive_taken => {
362                *receive_taken = true;
363                Ok(TcpReceiveStream {
364                    inner: stream.clone(),
365                })
366            }
367            TcpState::Closed(err) => Err(*err),
368            _ => Err(ErrorCode::InvalidState),
369        }
370    }
371
372    pub(crate) fn local_address(&mut self) -> Result<SocketAddr, ErrorCode> {
373        if !self.is_bound() {
374            return Err(ErrorCode::InvalidState);
375        }
376
377        match &self.tcp_state {
378            TcpState::Default(socket) => Ok(socket.local_addr()?),
379            TcpState::Connecting(_) => Err(ErrorCode::InvalidState),
380            TcpState::Connected { stream, .. } => Ok(stream.local_addr()?),
381            TcpState::Listening(listener) => Ok(listener.local_addr()?),
382            TcpState::Closed(err) => Err(*err),
383        }
384    }
385
386    pub(crate) fn remote_address(&mut self) -> Result<SocketAddr, ErrorCode> {
387        match &mut self.tcp_state {
388            TcpState::Connected {
389                peer: Some(peer), ..
390            } => Ok(*peer),
391            TcpState::Connected { stream, peer, .. } => {
392                let addr = stream.peer_addr()?;
393                *peer = Some(addr);
394                Ok(addr)
395            }
396            TcpState::Closed(err) => Err(*err),
397            _ => Err(ErrorCode::InvalidState),
398        }
399    }
400
401    pub(crate) fn is_listening(&self) -> bool {
402        matches!(self.tcp_state, TcpState::Listening(_))
403    }
404
405    pub(crate) fn address_family(&self) -> SocketAddressFamily {
406        self.family
407    }
408
409    pub(crate) fn set_listen_backlog_size(&mut self, value: u64) -> Result<(), ErrorCode> {
410        const MIN_BACKLOG: u32 = 1;
411        const MAX_BACKLOG: u32 = i32::MAX as u32; // OS'es will most likely limit it down even further.
412
413        if value == 0 {
414            return Err(ErrorCode::InvalidArgument);
415        }
416        // Silently clamp backlog size. This is OK for us to do, because operating systems do this too.
417        let value = value
418            .try_into()
419            .unwrap_or(MAX_BACKLOG)
420            .clamp(MIN_BACKLOG, MAX_BACKLOG);
421        match &self.tcp_state {
422            TcpState::Default(..) => {
423                // Socket not listening yet. Stash value for first invocation to `listen`.
424                self.listen_backlog_size = value;
425                Ok(())
426            }
427            TcpState::Listening(listener) => {
428                // Try to update the backlog by calling `listen` again.
429                // Not all platforms support this. We'll only update our own value if the OS supports changing the backlog size after the fact.
430                if rustix::net::listen(&listener, value.try_into().unwrap_or(i32::MAX)).is_err() {
431                    return Err(ErrorCode::NotSupported);
432                }
433                self.listen_backlog_size = value;
434                Ok(())
435            }
436            TcpState::Closed(err) => Err(*err),
437            _ => Err(ErrorCode::InvalidState),
438        }
439    }
440
441    pub(crate) fn keep_alive_enabled(&self) -> Result<bool, ErrorCode> {
442        let fd = self.as_fd()?;
443        let v = sockopt::socket_keepalive(fd)?;
444        Ok(v)
445    }
446
447    pub(crate) fn set_keep_alive_enabled(&self, value: bool) -> Result<(), ErrorCode> {
448        let fd = self.as_fd()?;
449        sockopt::set_socket_keepalive(fd, value)?;
450        Ok(())
451    }
452
453    pub(crate) fn keep_alive_idle_time(&self) -> Result<u64, ErrorCode> {
454        let fd = self.as_fd()?;
455        let v = sockopt::tcp_keepidle(fd)?;
456        Ok(v.as_nanos().try_into().unwrap_or(u64::MAX))
457    }
458
459    pub(crate) fn set_keep_alive_idle_time(&mut self, value: u64) -> Result<(), ErrorCode> {
460        if value == 0 {
461            // WIT: "If the provided value is 0, an `invalid-argument` error is returned."
462            return Err(ErrorCode::InvalidArgument);
463        }
464        let fd = self.as_fd()?;
465        let value = clamp_keep_alive_time(value);
466        sockopt::set_tcp_keepidle(fd, Duration::from_nanos(value))?;
467        self.listener_options.set_keep_alive_idle_time(value);
468        Ok(())
469    }
470
471    pub(crate) fn keep_alive_interval(&self) -> Result<u64, ErrorCode> {
472        let fd = self.as_fd()?;
473        let v = sockopt::tcp_keepintvl(fd)?;
474        Ok(v.as_nanos().try_into().unwrap_or(u64::MAX))
475    }
476
477    pub(crate) fn set_keep_alive_interval(&self, value: u64) -> Result<(), ErrorCode> {
478        if value == 0 {
479            // WIT: "If the provided value is 0, an `invalid-argument` error is returned."
480            return Err(ErrorCode::InvalidArgument);
481        }
482        let fd = self.as_fd()?;
483        let value = clamp_keep_alive_time(value);
484        sockopt::set_tcp_keepintvl(fd, Duration::from_nanos(value))?;
485        Ok(())
486    }
487
488    pub(crate) fn keep_alive_count(&self) -> Result<u32, ErrorCode> {
489        let fd = self.as_fd()?;
490        let v = sockopt::tcp_keepcnt(fd)?;
491        Ok(v)
492    }
493
494    pub(crate) fn set_keep_alive_count(&self, value: u32) -> Result<(), ErrorCode> {
495        if value == 0 {
496            // WIT: "If the provided value is 0, an `invalid-argument` error is returned."
497            return Err(ErrorCode::InvalidArgument);
498        }
499        let value = clamp_keep_alive_count(value);
500        let fd = self.as_fd()?;
501        sockopt::set_tcp_keepcnt(fd, value)?;
502        Ok(())
503    }
504
505    pub(crate) fn hop_limit(&self) -> Result<u8, ErrorCode> {
506        let fd = self.as_fd()?;
507        let n = get_unicast_hop_limit(fd, self.family)?;
508        Ok(n)
509    }
510
511    pub(crate) fn set_hop_limit(&mut self, value: u8) -> Result<(), ErrorCode> {
512        {
513            let fd = self.as_fd()?;
514            set_unicast_hop_limit(fd, self.family, value)?;
515        }
516        self.listener_options.set_hop_limit(value);
517        Ok(())
518    }
519
520    pub(crate) fn receive_buffer_size(&self) -> Result<u64, ErrorCode> {
521        let fd = self.as_fd()?;
522        let n = get_receive_buffer_size(fd)?;
523        Ok(n)
524    }
525
526    pub(crate) fn set_receive_buffer_size(&mut self, value: u64) -> Result<(), ErrorCode> {
527        let res = {
528            let fd = self.as_fd()?;
529            set_receive_buffer_size(fd, value)?
530        };
531        self.listener_options.set_receive_buffer_size(res);
532        Ok(())
533    }
534
535    pub(crate) fn send_buffer_size(&self) -> Result<u64, ErrorCode> {
536        let fd = self.as_fd()?;
537        let n = get_send_buffer_size(fd)?;
538        Ok(n)
539    }
540
541    pub(crate) fn set_send_buffer_size(&mut self, value: u64) -> Result<(), ErrorCode> {
542        let res = {
543            let fd = self.as_fd()?;
544            set_send_buffer_size(fd, value)?
545        };
546        self.listener_options.set_send_buffer_size(res);
547        Ok(())
548    }
549}
550
551pub(crate) struct TcpListenStream {
552    inner: Arc<tokio::net::TcpListener>,
553    family: SocketAddressFamily,
554    listener_options: NonInheritedOptions,
555    permissions: SocketAddrCheck,
556    pending_accept: Option<MaybeReady<Result<(tokio::net::TcpStream, SocketAddr), ErrorCode>>>,
557}
558impl TcpListenStream {
559    pub(crate) fn poll_accept(&mut self, cx: &mut std::task::Context<'_>) -> Poll<TcpSocket> {
560        ready!(self.poll_ready(cx));
561        let result = self.pending_accept.take().unwrap().unwrap_ready();
562        Poll::Ready(TcpSocket {
563            tcp_state: match result {
564                Ok((client, peer)) => {
565                    self.listener_options.apply(self.family, &client);
566                    TcpState::connected(client, Some(peer))
567                }
568                Err(err) => TcpState::Closed(err),
569            },
570            listen_backlog_size: DEFAULT_BACKLOG,
571            family: self.family,
572            is_bound: true,
573            listener_options: Default::default(),
574            permissions: self.permissions.clone(),
575        })
576    }
577
578    pub(crate) fn poll_ready(&mut self, cx: &mut std::task::Context<'_>) -> Poll<()> {
579        if self.pending_accept.is_none() {
580            let listener = self.inner.clone();
581            let permissions = self.permissions.clone();
582
583            self.pending_accept = Some(MaybeReady::new(Box::pin(async move {
584                loop {
585                    match accept(&listener).await {
586                        Ok((client, addr)) => {
587                            if permissions
588                                .check(addr, SocketAddrUse::TcpAccept)
589                                .await
590                                .is_ok()
591                            {
592                                return Ok((client, addr));
593                            } else {
594                                reset(client);
595                                continue;
596                            }
597                        }
598                        Err(err) => {
599                            return Err(err.into());
600                        }
601                    }
602                }
603            })));
604        }
605
606        with_ambient_tokio_runtime(|| {
607            self.pending_accept
608                .as_mut()
609                .unwrap()
610                .poll_ready(cx)
611                .map(|_| ())
612        })
613    }
614}
615
616pub(crate) struct TcpSendStream {
617    inner: Arc<tokio::net::TcpStream>,
618}
619impl TcpSendStream {
620    pub(crate) fn poll_ready(&mut self, cx: &mut std::task::Context<'_>) -> Poll<()> {
621        self.inner.poll_write_ready(cx).map(|_| ())
622    }
623
624    pub(crate) fn poll_write(
625        &mut self,
626        cx: &mut std::task::Context<'_>,
627        buf: &[u8],
628    ) -> Poll<Result<usize, ErrorCode>> {
629        loop {
630            return match self.inner.try_write(buf) {
631                Ok(n) => Poll::Ready(Ok(n)),
632                Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
633                    match self.inner.poll_write_ready(cx) {
634                        Poll::Ready(Ok(())) => continue,
635                        Poll::Ready(Err(e)) => Poll::Ready(Err(e.into())),
636                        Poll::Pending => Poll::Pending,
637                    }
638                }
639                Err(e) => Poll::Ready(Err(match Errno::from_io_error(&e) {
640                    #[cfg(windows)]
641                    Some(Errno::SHUTDOWN) | Some(Errno::CONNABORTED) => ErrorCode::ConnectionBroken,
642                    #[cfg(not(windows))]
643                    Some(Errno::PIPE) => ErrorCode::ConnectionBroken,
644
645                    _ => e.into(),
646                })),
647            };
648        }
649    }
650    pub(crate) async fn write(&mut self, buf: &[u8]) -> Result<usize, ErrorCode> {
651        poll_fn(|cx| self.poll_write(cx, buf)).await
652    }
653}
654impl Drop for TcpSendStream {
655    fn drop(&mut self) {
656        _ = rustix::net::shutdown(&self.inner, rustix::net::Shutdown::Write);
657    }
658}
659
660pub(crate) struct TcpReceiveStream {
661    inner: Arc<tokio::net::TcpStream>,
662}
663impl TcpReceiveStream {
664    pub(crate) fn poll_ready(&mut self, cx: &mut std::task::Context<'_>) -> Poll<()> {
665        self.inner.poll_read_ready(cx).map(|_| ())
666    }
667
668    pub(crate) fn poll_read(
669        &mut self,
670        cx: &mut std::task::Context<'_>,
671        buf: &mut [u8],
672    ) -> Poll<Result<usize, ErrorCode>> {
673        if buf.is_empty() {
674            return Poll::Ready(Ok(0));
675        }
676        loop {
677            return match self.inner.try_read(buf) {
678                Ok(0) => Poll::Ready(Ok(0)),
679                Ok(n) => Poll::Ready(Ok(n)),
680                Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
681                    match self.inner.poll_read_ready(cx) {
682                        Poll::Ready(Ok(())) => continue,
683                        Poll::Ready(Err(e)) => Poll::Ready(Err(e.into())),
684                        Poll::Pending => Poll::Pending,
685                    }
686                }
687                Err(e) => Poll::Ready(Err(e.into())),
688            };
689        }
690    }
691}
692impl Drop for TcpReceiveStream {
693    fn drop(&mut self) {
694        _ = rustix::net::shutdown(&self.inner, rustix::net::Shutdown::Read);
695    }
696}
697
698#[cfg(not(target_os = "macos"))]
699pub use inherits_option::*;
700#[cfg(not(target_os = "macos"))]
701mod inherits_option {
702    use crate::sockets::SocketAddressFamily;
703    use tokio::net::TcpStream;
704
705    #[derive(Default, Clone)]
706    pub struct NonInheritedOptions;
707
708    impl NonInheritedOptions {
709        pub fn set_keep_alive_idle_time(&mut self, _value: u64) {}
710
711        pub fn set_hop_limit(&mut self, _value: u8) {}
712
713        pub fn set_receive_buffer_size(&mut self, _value: usize) {}
714
715        pub fn set_send_buffer_size(&mut self, _value: usize) {}
716
717        pub(crate) fn apply(&self, _family: SocketAddressFamily, _stream: &TcpStream) {}
718    }
719}
720
721#[cfg(target_os = "macos")]
722pub use does_not_inherit_options::*;
723#[cfg(target_os = "macos")]
724mod does_not_inherit_options {
725    use crate::sockets::SocketAddressFamily;
726    use rustix::net::sockopt;
727    use std::sync::Arc;
728    use std::sync::atomic::{AtomicU8, AtomicU64, AtomicUsize, Ordering::Relaxed};
729    use std::time::Duration;
730    use tokio::net::TcpStream;
731
732    // The socket options below are not automatically inherited from the listener
733    // on all platforms. So we keep track of which options have been explicitly
734    // set and manually apply those values to newly accepted clients.
735    #[derive(Default, Clone)]
736    pub struct NonInheritedOptions(Arc<Inner>);
737
738    #[derive(Default)]
739    struct Inner {
740        receive_buffer_size: AtomicUsize,
741        send_buffer_size: AtomicUsize,
742        hop_limit: AtomicU8,
743        keep_alive_idle_time: AtomicU64, // nanoseconds
744    }
745
746    impl NonInheritedOptions {
747        pub fn set_keep_alive_idle_time(&mut self, value: u64) {
748            self.0.keep_alive_idle_time.store(value, Relaxed);
749        }
750
751        pub fn set_hop_limit(&mut self, value: u8) {
752            self.0.hop_limit.store(value, Relaxed);
753        }
754
755        pub fn set_receive_buffer_size(&mut self, value: usize) {
756            self.0.receive_buffer_size.store(value, Relaxed);
757        }
758
759        pub fn set_send_buffer_size(&mut self, value: usize) {
760            self.0.send_buffer_size.store(value, Relaxed);
761        }
762
763        pub(crate) fn apply(&self, family: SocketAddressFamily, stream: &TcpStream) {
764            // Manually inherit socket options from listener. We only have to
765            // do this on platforms that don't already do this automatically
766            // and only if a specific value was explicitly set on the listener.
767
768            let receive_buffer_size = self.0.receive_buffer_size.load(Relaxed);
769            if receive_buffer_size > 0 {
770                // Ignore potential error.
771                _ = sockopt::set_socket_recv_buffer_size(&stream, receive_buffer_size);
772            }
773
774            let send_buffer_size = self.0.send_buffer_size.load(Relaxed);
775            if send_buffer_size > 0 {
776                // Ignore potential error.
777                _ = sockopt::set_socket_send_buffer_size(&stream, send_buffer_size);
778            }
779
780            // For some reason, IP_TTL is inherited, but IPV6_UNICAST_HOPS isn't.
781            if family == SocketAddressFamily::Ipv6 {
782                let hop_limit = self.0.hop_limit.load(Relaxed);
783                if hop_limit > 0 {
784                    // Ignore potential error.
785                    _ = sockopt::set_ipv6_unicast_hops(&stream, Some(hop_limit));
786                }
787            }
788
789            let keep_alive_idle_time = self.0.keep_alive_idle_time.load(Relaxed);
790            if keep_alive_idle_time > 0 {
791                // Ignore potential error.
792                _ = sockopt::set_tcp_keepidle(&stream, Duration::from_nanos(keep_alive_idle_time));
793            }
794        }
795    }
796}
797
798fn socket(family: SocketAddressFamily) -> std::io::Result<tokio::net::TcpSocket> {
799    match family {
800        SocketAddressFamily::Ipv4 => tokio::net::TcpSocket::new_v4(),
801        SocketAddressFamily::Ipv6 => {
802            let socket = tokio::net::TcpSocket::new_v6()?;
803
804            // From the WASI spec:
805            // > On IPv6 sockets, IPV6_V6ONLY is enabled by default and can't
806            // > be configured otherwise.
807            sockopt::set_ipv6_v6only(&socket, true)?;
808            Ok(socket)
809        }
810    }
811}
812
813fn bind(socket: &tokio::net::TcpSocket, local_address: SocketAddr) -> Result<(), ErrorCode> {
814    // From the WASI spec:
815    // > The bind operation shouldn't be affected by the TIME_WAIT state of a
816    // > recently closed socket on the same local address. In practice this
817    // > means that the SO_REUSEADDR socket option should be set implicitly on
818    // > all platforms, except on Windows where this is the default behavior
819    // > and SO_REUSEADDR performs something different.
820    #[cfg(not(windows))]
821    {
822        _ = sockopt::set_socket_reuseaddr(&socket, true);
823    }
824
825    // Perform the OS bind call.
826    socket
827        .bind(local_address)
828        .map_err(|err| match Errno::from_io_error(&err) {
829            // From https://pubs.opengroup.org/onlinepubs/9699919799/functions/bind.html:
830            // > [EAFNOSUPPORT] The specified address is not a valid address for the address family of the specified socket
831            //
832            // The most common reasons for this error should have already
833            // been handled by our own validation.. This error mapping is here
834            // just in case there is an edge case we didn't catch.
835            Some(Errno::AFNOSUPPORT) => ErrorCode::InvalidArgument,
836            // See: https://learn.microsoft.com/en-us/windows/win32/api/winsock2/nf-winsock2-bind#:~:text=WSAENOBUFS
837            // Windows returns WSAENOBUFS when the ephemeral ports have been exhausted.
838            #[cfg(windows)]
839            Some(Errno::NOBUFS) => ErrorCode::AddressInUse,
840            _ => err.into(),
841        })
842}
843
844async fn accept(
845    listener: &tokio::net::TcpListener,
846) -> std::io::Result<(tokio::net::TcpStream, SocketAddr)> {
847    listener
848        .accept()
849        .await
850        .map_err(|err| match Errno::from_io_error(&err) {
851            // From: https://learn.microsoft.com/en-us/windows/win32/api/winsock2/nf-winsock2-accept#:~:text=WSAEINPROGRESS
852            // > WSAEINPROGRESS: A blocking Windows Sockets 1.1 call is in progress,
853            // > or the service provider is still processing a callback function.
854            //
855            // wasi-sockets doesn't have an equivalent to the EINPROGRESS error,
856            // because in POSIX this error is only returned by a non-blocking
857            // `connect` and wasi-sockets has a different solution for that.
858            #[cfg(windows)]
859            Some(Errno::INPROGRESS) => Errno::INTR.into(),
860
861            // Normalize Linux' non-standard behavior.
862            //
863            // From https://man7.org/linux/man-pages/man2/accept.2.html:
864            // > Linux accept() passes already-pending network errors on the
865            // > new socket as an error code from accept(). This behavior
866            // > differs from other BSD socket implementations. (...)
867            #[cfg(target_os = "linux")]
868            Some(
869                Errno::CONNRESET
870                | Errno::NETRESET
871                | Errno::HOSTUNREACH
872                | Errno::HOSTDOWN
873                | Errno::NETDOWN
874                | Errno::NETUNREACH
875                | Errno::PROTO
876                | Errno::NOPROTOOPT
877                | Errno::NONET
878                | Errno::OPNOTSUPP,
879            ) => Errno::CONNABORTED.into(),
880
881            _ => err,
882        })
883}
884
885fn reset(socket: tokio::net::TcpStream) {
886    _ = socket.set_zero_linger();
887    drop(socket);
888}
889
890fn clamp_keep_alive_time(value: u64) -> u64 {
891    // Ensure that the value passed to the actual syscall never gets rounded down to 0.
892    const MIN: u64 = 1 * NANOS_PER_SEC;
893
894    // Cap it at Linux' maximum, which appears to have the lowest limit across our supported platforms.
895    const MAX: u64 = (i16::MAX as u64) * NANOS_PER_SEC;
896
897    value.clamp(MIN, MAX)
898}
899
900fn clamp_keep_alive_count(value: u32) -> u32 {
901    const MIN_CNT: u32 = 1;
902    // Cap it at Linux' maximum, which appears to have the lowest limit across our supported platforms.
903    const MAX_CNT: u32 = i8::MAX as u32;
904
905    value.clamp(MIN_CNT, MAX_CNT)
906}
907
908#[cfg(all(test, any(target_os = "linux", target_os = "macos")))]
909mod tests {
910    use super::*;
911    use crate::WasiCtxBuilder;
912
913    #[tokio::test]
914    async fn accepted_remote_address_survives_reset() {
915        let mut ctx = WasiCtxBuilder::new();
916        ctx.inherit_network().allow_tcp(true);
917        let ctx = ctx.build();
918        let mut socket = TcpSocket::new(&ctx.sockets, SocketAddressFamily::Ipv4).unwrap();
919        socket.bind("127.0.0.1:0".parse().unwrap()).await.unwrap();
920        let mut listener = socket.listen().await.unwrap();
921
922        let client = tokio::net::TcpStream::connect(socket.local_address().unwrap())
923            .await
924            .unwrap();
925        let peer = client.local_addr().unwrap();
926
927        // Complete the OS accept, but leave the result pending for poll_accept.
928        poll_fn(|cx| listener.poll_ready(cx)).await;
929        client.set_zero_linger().unwrap();
930        drop(client);
931
932        let mut accepted = poll_fn(|cx| listener.poll_accept(cx)).await;
933        let mut input = accepted.take_receive_stream().unwrap();
934        let mut byte = [0];
935        let read = tokio::time::timeout(
936            Duration::from_secs(5),
937            poll_fn(|cx| input.poll_read(cx, &mut byte)),
938        )
939        .await
940        .unwrap();
941        assert!(matches!(read, Err(ErrorCode::ConnectionReset)));
942
943        // The OS no longer answers getpeername, but accept already had the
944        // address and the accepted socket must still report it.
945        let TcpState::Connected { stream, .. } = &accepted.tcp_state else {
946            panic!("expected an accepted connection");
947        };
948        assert!(stream.peer_addr().is_err());
949        assert_eq!(accepted.remote_address().unwrap(), peer);
950    }
951
952    #[tokio::test]
953    async fn connected_remote_address_survives_reset() {
954        let mut ctx = WasiCtxBuilder::new();
955        ctx.inherit_network().allow_tcp(true);
956        let ctx = ctx.build();
957        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
958        let peer = listener.local_addr().unwrap();
959
960        let mut client = TcpSocket::new(&ctx.sockets, SocketAddressFamily::Ipv4).unwrap();
961        client.start_connect(peer).unwrap();
962        poll_fn(|cx| client.poll_finish_connect(cx)).await.unwrap();
963        let (server, _) = listener.accept().await.unwrap();
964
965        // Query the outgoing socket's peer while the connection is healthy.
966        assert_eq!(client.remote_address().unwrap(), peer);
967        server.set_zero_linger().unwrap();
968        drop(server);
969
970        let mut input = client.take_receive_stream().unwrap();
971        let mut byte = [0];
972        let read = tokio::time::timeout(
973            Duration::from_secs(5),
974            poll_fn(|cx| input.poll_read(cx, &mut byte)),
975        )
976        .await
977        .unwrap();
978        assert!(matches!(read, Err(ErrorCode::ConnectionReset)));
979
980        // Subsequent queries must use the previously observed peer address.
981        let TcpState::Connected { stream, .. } = &client.tcp_state else {
982            panic!("expected an outgoing connection");
983        };
984        assert!(stream.peer_addr().is_err());
985        assert_eq!(client.remote_address().unwrap(), peer);
986    }
987}