From b098df5efc84997e7d4987aeb853f952249f665f Mon Sep 17 00:00:00 2001 From: Adam Bratschi-Kaye Date: Tue, 22 Sep 2026 01:16:49 +0000 Subject: [PATCH 1/4] Implement `net` for WASIp3 --- examples/tcp_echo_server.rs | 8 +- examples/tcp_stream_client.rs | 4 +- examples/udp_echo_server.rs | 8 +- examples/udp_stream_client.rs | 4 +- src/lib.rs | 1 - src/net/mod.rs | 47 ++++- src/net/tcp_listener.rs | 70 +++++-- src/net/tcp_stream.rs | 60 ++++-- src/net/udp.rs | 240 +++++++++++++++++------ test-programs/tests/tcp_echo_server.rs | 29 ++- test-programs/tests/tcp_stream_client.rs | 29 ++- test-programs/tests/udp_echo_server.rs | 29 ++- test-programs/tests/udp_stream_client.rs | 27 ++- 13 files changed, 411 insertions(+), 145 deletions(-) diff --git a/examples/tcp_echo_server.rs b/examples/tcp_echo_server.rs index 224222a..20fb24f 100644 --- a/examples/tcp_echo_server.rs +++ b/examples/tcp_echo_server.rs @@ -1,5 +1,5 @@ -#![cfg_attr(not(all(target_os = "wasi", target_env = "p2")), no_main)] -#![cfg(all(target_os = "wasi", target_env = "p2"))] +#![cfg_attr(not(target_os = "wasi"), no_main)] +#![cfg(target_os = "wasi")] use wstd::io; use wstd::iter::AsyncIterator; @@ -7,9 +7,9 @@ use wstd::net::TcpListener; #[wstd::main] async fn main() -> io::Result<()> { - let mut listener = TcpListener::bind("127.0.0.1:8080").await?; + let mut listener = TcpListener::bind("127.0.0.1:0").await?; println!("Listening on {}", listener.local_addr()?); - println!("type `nc localhost 8080` to create a TCP client"); + println!("type `nc localhost ` to create a TCP client"); let mut incoming = listener.incoming(); while let Some(stream) = incoming.next().await { diff --git a/examples/tcp_stream_client.rs b/examples/tcp_stream_client.rs index a269b8c..0b04efc 100644 --- a/examples/tcp_stream_client.rs +++ b/examples/tcp_stream_client.rs @@ -1,5 +1,5 @@ -#![cfg_attr(not(all(target_os = "wasi", target_env = "p2")), no_main)] -#![cfg(all(target_os = "wasi", target_env = "p2"))] +#![cfg_attr(not(target_os = "wasi"), no_main)] +#![cfg(target_os = "wasi")] use wstd::io::{self, AsyncRead, AsyncWrite}; use wstd::net::TcpStream; diff --git a/examples/udp_echo_server.rs b/examples/udp_echo_server.rs index c441d87..28a3165 100644 --- a/examples/udp_echo_server.rs +++ b/examples/udp_echo_server.rs @@ -1,14 +1,14 @@ -#![cfg_attr(not(all(target_os = "wasi", target_env = "p2")), no_main)] -#![cfg(all(target_os = "wasi", target_env = "p2"))] +#![cfg_attr(not(target_os = "wasi"), no_main)] +#![cfg(target_os = "wasi")] use wstd::io; use wstd::net::UdpSocket; #[wstd::main] async fn main() -> io::Result<()> { - let socket = UdpSocket::bind("127.0.0.1:8080").await?; + let socket = UdpSocket::bind("127.0.0.1:0").await?; println!("Listening on {}", socket.local_addr()?); - println!("type `nc -u localhost 8080` to create a UDP client"); + println!("type `nc -u localhost ` to create a UDP client"); let mut buf = vec![0; 65535]; loop { diff --git a/examples/udp_stream_client.rs b/examples/udp_stream_client.rs index f26f5d5..25b06ab 100644 --- a/examples/udp_stream_client.rs +++ b/examples/udp_stream_client.rs @@ -1,5 +1,5 @@ -#![cfg_attr(not(all(target_os = "wasi", target_env = "p2")), no_main)] -#![cfg(all(target_os = "wasi", target_env = "p2"))] +#![cfg_attr(not(target_os = "wasi"), no_main)] +#![cfg(target_os = "wasi")] use wstd::io; use wstd::net::{UdpSocket, UdpStream}; diff --git a/src/lib.rs b/src/lib.rs index 0660725..f9c4f06 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -62,7 +62,6 @@ pub mod future; pub mod http; pub mod io; pub mod iter; -#[cfg(all(target_os = "wasi", target_env = "p2"))] pub mod net; pub mod rand; pub mod runtime; diff --git a/src/net/mod.rs b/src/net/mod.rs index dc1dc41..f18287e 100644 --- a/src/net/mod.rs +++ b/src/net/mod.rs @@ -1,7 +1,16 @@ //! Async network abstractions. use std::io::{self, ErrorKind}; -use wasip2::sockets::network::{ErrorCode, IpSocketAddress, Ipv4SocketAddress}; +#[cfg(target_env = "p2")] +use wasip2::sockets::{ + network::{ErrorCode, IpSocketAddress, Ipv4SocketAddress, Ipv6SocketAddress}, + tcp_create_socket::create_tcp_socket, + udp_create_socket::create_udp_socket, +}; +#[cfg(target_env = "p3")] +use wasip3::sockets::types::{ + ErrorCode, IpAddressFamily, IpSocketAddress, Ipv4SocketAddress, Ipv6SocketAddress, TcpSocket, +}; mod tcp_listener; mod tcp_stream; @@ -13,26 +22,39 @@ pub use udp::*; fn to_io_err(err: ErrorCode) -> io::Error { match err { - ErrorCode::Unknown => ErrorKind::Other.into(), ErrorCode::AccessDenied => ErrorKind::PermissionDenied.into(), ErrorCode::NotSupported => ErrorKind::Unsupported.into(), ErrorCode::InvalidArgument => ErrorKind::InvalidInput.into(), ErrorCode::OutOfMemory => ErrorKind::OutOfMemory.into(), ErrorCode::Timeout => ErrorKind::TimedOut.into(), - ErrorCode::WouldBlock => ErrorKind::WouldBlock.into(), ErrorCode::InvalidState => ErrorKind::InvalidData.into(), ErrorCode::AddressInUse => ErrorKind::AddrInUse.into(), ErrorCode::ConnectionRefused => ErrorKind::ConnectionRefused.into(), ErrorCode::ConnectionReset => ErrorKind::ConnectionReset.into(), ErrorCode::ConnectionAborted => ErrorKind::ConnectionAborted.into(), - ErrorCode::ConcurrencyConflict => ErrorKind::AlreadyExists.into(), ErrorCode::DatagramTooLarge => ErrorKind::InvalidInput.into(), + + #[cfg(target_env = "p2")] + ErrorCode::Unknown => ErrorKind::Other.into(), + #[cfg(target_env = "p2")] + ErrorCode::WouldBlock => ErrorKind::WouldBlock.into(), + #[cfg(target_env = "p2")] + ErrorCode::ConcurrencyConflict => ErrorKind::AlreadyExists.into(), + #[cfg(target_env = "p2")] _ => ErrorKind::Other.into(), + + #[cfg(target_env = "p3")] + ErrorCode::AddressNotBindable => ErrorKind::AddrNotAvailable.into(), + #[cfg(target_env = "p3")] + ErrorCode::RemoteUnreachable => ErrorKind::HostUnreachable.into(), + #[cfg(target_env = "p3")] + ErrorCode::ConnectionBroken => ErrorKind::BrokenPipe.into(), + #[cfg(target_env = "p3")] + ErrorCode::Other(s) => io::Error::other(s.unwrap_or_default()), } } fn sockaddr_from_wasi(addr: IpSocketAddress) -> std::net::SocketAddr { - use wasip2::sockets::network::Ipv6SocketAddress; match addr { IpSocketAddress::Ipv4(Ipv4SocketAddress { address, port }) => { std::net::SocketAddr::V4(std::net::SocketAddrV4::new( @@ -58,7 +80,6 @@ fn sockaddr_from_wasi(addr: IpSocketAddress) -> std::net::SocketAddr { } fn sockaddr_to_wasi(addr: std::net::SocketAddr) -> IpSocketAddress { - use wasip2::sockets::network::Ipv6SocketAddress; match addr { std::net::SocketAddr::V4(addr) => { let ip = addr.ip().octets(); @@ -78,3 +99,17 @@ fn sockaddr_to_wasi(addr: std::net::SocketAddr) -> IpSocketAddress { } } } + +#[cfg(target_env = "p3")] +fn create_tcp_socket( + family: IpAddressFamily, +) -> Result { + TcpSocket::create(family) +} + +#[cfg(target_env = "p3")] +fn create_udp_socket( + family: IpAddressFamily, +) -> Result { + wasip3::sockets::types::UdpSocket::create(family) +} diff --git a/src/net/tcp_listener.rs b/src/net/tcp_listener.rs index 69a70f3..d4cfdae 100644 --- a/src/net/tcp_listener.rs +++ b/src/net/tcp_listener.rs @@ -1,17 +1,27 @@ +#[cfg(target_env = "p2")] use wasip2::sockets::tcp::{IpAddressFamily, TcpSocket}; +#[cfg(target_env = "p3")] +use wasip3::{ + sockets::types::{IpAddressFamily, TcpSocket}, + wit_bindgen::StreamReader, +}; use crate::io; use crate::iter::AsyncIterator; use std::net::SocketAddr; -use super::{TcpStream, sockaddr_from_wasi, sockaddr_to_wasi, to_io_err}; +use super::{TcpStream, create_tcp_socket, sockaddr_from_wasi, sockaddr_to_wasi, to_io_err}; +#[cfg(target_env = "p2")] use crate::runtime::AsyncPollable; /// A TCP socket server, listening for connections. #[derive(Debug)] pub struct TcpListener { // Field order matters: must drop this child before parent below + #[cfg(target_env = "p2")] pollable: AsyncPollable, + #[cfg(target_env = "p3")] + connections: StreamReader, socket: TcpSocket, } @@ -27,31 +37,42 @@ impl TcpListener { SocketAddr::V4(_) => IpAddressFamily::Ipv4, SocketAddr::V6(_) => IpAddressFamily::Ipv6, }; - let socket = - wasip2::sockets::tcp_create_socket::create_tcp_socket(family).map_err(to_io_err)?; - let network = wasip2::sockets::instance_network::instance_network(); - + let socket = create_tcp_socket(family).map_err(to_io_err)?; let local_address = sockaddr_to_wasi(addr); - socket - .start_bind(&network, local_address) - .map_err(to_io_err)?; - let pollable = AsyncPollable::new(socket.subscribe()); - pollable.wait_for().await; - socket.finish_bind().map_err(to_io_err)?; + #[cfg(target_env = "p2")] + { + let network = wasip2::sockets::instance_network::instance_network(); + socket + .start_bind(&network, local_address) + .map_err(to_io_err)?; + let pollable = AsyncPollable::new(socket.subscribe()); + pollable.wait_for().await; + socket.finish_bind().map_err(to_io_err)?; - socket.start_listen().map_err(to_io_err)?; - pollable.wait_for().await; - socket.finish_listen().map_err(to_io_err)?; - Ok(Self { pollable, socket }) + socket.start_listen().map_err(to_io_err)?; + pollable.wait_for().await; + socket.finish_listen().map_err(to_io_err)?; + Ok(Self { pollable, socket }) + } + #[cfg(target_env = "p3")] + { + socket.bind(local_address).map_err(to_io_err)?; + let connections = socket.listen().map_err(to_io_err)?; + Ok(Self { + connections, + socket, + }) + } } /// Returns the local socket address of this listener. pub fn local_addr(&self) -> io::Result { - self.socket - .local_address() - .map_err(to_io_err) - .map(sockaddr_from_wasi) + #[cfg(target_env = "p2")] + let addr = self.socket.local_address(); + #[cfg(target_env = "p3")] + let addr = self.socket.get_local_address(); + addr.map_err(to_io_err).map(sockaddr_from_wasi) } /// Returns an iterator over the connections being received on this listener. @@ -69,6 +90,7 @@ pub struct Incoming<'a> { impl<'a> AsyncIterator for Incoming<'a> { type Item = io::Result; + #[cfg(target_env = "p2")] async fn next(&mut self) -> Option { self.listener.pollable.wait_for().await; let (socket, input, output) = match self.listener.socket.accept().map_err(to_io_err) { @@ -77,4 +99,14 @@ impl<'a> AsyncIterator for Incoming<'a> { }; Some(Ok(TcpStream::new(input, output, socket))) } + + #[cfg(target_env = "p3")] + async fn next(&mut self) -> Option { + self.listener.connections.next().await.map(|socket| { + let (input, _receive_result) = socket.receive(); + let (output, receiver) = wasip3::wit_stream::new(); + let _send_result = socket.send(receiver); + Ok(TcpStream::new(input, output, socket)) + }) + } } diff --git a/src/net/tcp_stream.rs b/src/net/tcp_stream.rs index 977fb29..04dad77 100644 --- a/src/net/tcp_stream.rs +++ b/src/net/tcp_stream.rs @@ -1,16 +1,26 @@ use std::io::ErrorKind; use std::net::{SocketAddr, ToSocketAddrs}; -use wasip2::sockets::instance_network::instance_network; -use wasip2::sockets::network::Ipv4SocketAddress; -use wasip2::sockets::tcp::{IpAddressFamily, IpSocketAddress}; -use wasip2::sockets::tcp_create_socket::create_tcp_socket; + +#[cfg(target_env = "p2")] use wasip2::{ io::streams::{InputStream, OutputStream}, - sockets::tcp::TcpSocket, + sockets::{ + instance_network::instance_network, + network::Ipv4SocketAddress, + tcp::{IpAddressFamily, IpSocketAddress, TcpSocket}, + }, }; -use super::to_io_err; +#[cfg(target_env = "p3")] +use wasip3::sockets::types::{IpAddressFamily, IpSocketAddress, Ipv4SocketAddress, TcpSocket}; +#[cfg(target_env = "p3")] +type InputStream = wasip3::wit_bindgen::StreamReader; +#[cfg(target_env = "p3")] +type OutputStream = wasip3::wit_bindgen::StreamWriter; + +use super::{create_tcp_socket, to_io_err}; use crate::io::{self, AsyncInputStream, AsyncOutputStream}; +#[cfg(target_env = "p2")] use crate::runtime::AsyncPollable; /// A TCP stream between a local and a remote socket. @@ -59,7 +69,6 @@ impl TcpStream { SocketAddr::V6(_) => IpAddressFamily::Ipv6, }; let socket = create_tcp_socket(family).map_err(to_io_err)?; - let network = instance_network(); let remote_address = match addr { SocketAddr::V4(addr) => { @@ -70,19 +79,33 @@ impl TcpStream { } SocketAddr::V6(_) => todo!("IPv6 not yet supported in `wstd::net::TcpStream`"), }; - socket - .start_connect(&network, remote_address) - .map_err(to_io_err)?; - let pollable = AsyncPollable::new(socket.subscribe()); - pollable.wait_for().await; - let (input, output) = socket.finish_connect().map_err(to_io_err)?; - - Ok(TcpStream::new(input, output, socket)) + #[cfg(target_env = "p2")] + { + let network = instance_network(); + socket + .start_connect(&network, remote_address) + .map_err(to_io_err)?; + let pollable = AsyncPollable::new(socket.subscribe()); + pollable.wait_for().await; + let (input, output) = socket.finish_connect().map_err(to_io_err)?; + Ok(TcpStream::new(input, output, socket)) + } + #[cfg(target_env = "p3")] + { + socket.connect(remote_address).await.map_err(to_io_err)?; + let (input, _receive_result) = socket.receive(); + let (output, receiver) = wasip3::wit_stream::new(); + let _send_result = socket.send(receiver); + Ok(TcpStream::new(input, output, socket)) + } } /// Returns the socket address of the remote peer of this TCP connection. pub fn peer_addr(&self) -> io::Result { + #[cfg(target_env = "p2")] let addr = self.socket.remote_address().map_err(to_io_err)?; + #[cfg(target_env = "p3")] + let addr = self.socket.get_remote_address().map_err(to_io_err)?; Ok(format!("{addr:?}")) } @@ -90,16 +113,19 @@ impl TcpStream { ( ReadHalf { stream: &mut self.input, + #[cfg(target_env = "p2")] socket: &self.socket, }, WriteHalf { stream: &mut self.output, + #[cfg(target_env = "p2")] socket: &self.socket, }, ) } } +#[cfg(target_env = "p2")] impl Drop for TcpStream { fn drop(&mut self) { let _ = self @@ -134,9 +160,11 @@ impl io::AsyncWrite for TcpStream { pub struct ReadHalf<'a> { stream: &'a mut AsyncInputStream, + #[cfg(target_env = "p2")] socket: &'a TcpSocket, } +#[cfg(target_env = "p2")] impl<'a> Drop for ReadHalf<'a> { fn drop(&mut self) { let _ = self @@ -157,6 +185,7 @@ impl<'a> io::AsyncRead for ReadHalf<'a> { pub struct WriteHalf<'a> { stream: &'a mut AsyncOutputStream, + #[cfg(target_env = "p2")] socket: &'a TcpSocket, } @@ -174,6 +203,7 @@ impl<'a> io::AsyncWrite for WriteHalf<'a> { } } +#[cfg(target_env = "p2")] impl<'a> Drop for WriteHalf<'a> { fn drop(&mut self) { let _ = self diff --git a/src/net/udp.rs b/src/net/udp.rs index 59afb0d..7728f30 100644 --- a/src/net/udp.rs +++ b/src/net/udp.rs @@ -1,16 +1,22 @@ use std::io::ErrorKind; use std::net::{SocketAddr, ToSocketAddrs}; +#[cfg(target_env = "p2")] use std::sync::OnceLock; -use wasip2::sockets::instance_network::instance_network; -use wasip2::sockets::udp::{ - IncomingDatagramStream, IpAddressFamily, IpSocketAddress, OutgoingDatagram, - OutgoingDatagramStream, +#[cfg(target_env = "p2")] +use wasip2::sockets::{ + instance_network::instance_network, + udp::{ + IncomingDatagramStream, IpAddressFamily, IpSocketAddress, OutgoingDatagram, + OutgoingDatagramStream, UdpSocket as WasiUdpSocket, + }, }; -use wasip2::sockets::udp_create_socket::create_udp_socket; +#[cfg(target_env = "p3")] +use wasip3::sockets::types::{IpAddressFamily, UdpSocket as WasiUdpSocket}; -use super::{sockaddr_from_wasi, sockaddr_to_wasi, to_io_err}; +use super::{create_udp_socket, sockaddr_from_wasi, sockaddr_to_wasi, to_io_err}; use crate::io; +#[cfg(target_env = "p2")] use crate::runtime::AsyncPollable; /// A UDP socket, bound to a local address. @@ -21,9 +27,11 @@ use crate::runtime::AsyncPollable; /// single remote address instead, giving a [`UdpStream`]. #[derive(Debug)] pub struct UdpSocket { + #[cfg(target_env = "p2")] incoming: AsyncIncomingDatagramStream, + #[cfg(target_env = "p2")] outgoing: AsyncOutgoingDatagramStream, - socket: wasip2::sockets::udp::UdpSocket, + socket: WasiUdpSocket, } impl UdpSocket { @@ -34,29 +42,45 @@ impl UdpSocket { .map_err(|_| io::Error::other("failed to parse string to socket addr"))?; let socket = bind_socket(addr).await?; - // Datagram streams without a remote address may send to, and receive - // from, any address. - let (incoming, outgoing) = socket.stream(None).map_err(to_io_err)?; - Ok(Self { - incoming: AsyncIncomingDatagramStream::new(incoming), - outgoing: AsyncOutgoingDatagramStream::new(outgoing), - socket, - }) + #[cfg(target_env = "p2")] + { + // Datagram streams without a remote address may send to, and receive + // from, any address. + let (incoming, outgoing) = socket.stream(None).map_err(to_io_err)?; + Ok(Self { + incoming: AsyncIncomingDatagramStream::new(incoming), + outgoing: AsyncOutgoingDatagramStream::new(outgoing), + socket, + }) + } + #[cfg(target_env = "p3")] + Ok(Self { socket }) } /// Returns the local socket address of this socket. pub fn local_addr(&self) -> io::Result { - self.socket - .local_address() - .map_err(to_io_err) - .map(sockaddr_from_wasi) + #[cfg(target_env = "p2")] + let addr = self.socket.local_address(); + #[cfg(target_env = "p3")] + let addr = self.socket.get_local_address(); + addr.map_err(to_io_err).map(sockaddr_from_wasi) } /// Sends a datagram to the given address. pub async fn send_to(&self, buf: &[u8], addr: SocketAddr) -> io::Result { - self.outgoing + #[cfg(target_env = "p2")] + return self + .outgoing .send_to(buf, Some(sockaddr_to_wasi(addr))) - .await + .await; + #[cfg(target_env = "p3")] + { + self.socket + .send(buf.to_vec(), Some(sockaddr_to_wasi(addr))) + .await + .map_err(to_io_err)?; + Ok(buf.len()) + } } /// Receives a single datagram. On success, returns the number of bytes @@ -64,7 +88,15 @@ impl UdpSocket { /// /// If `buf` is shorter than the datagram, the excess bytes are discarded. pub async fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { - self.incoming.recv_from(buf).await + #[cfg(target_env = "p2")] + return self.incoming.recv_from(buf).await; + #[cfg(target_env = "p3")] + { + let (datagram, remote_address) = self.socket.receive().await.map_err(to_io_err)?; + let len = datagram.len().min(buf.len()); + buf[..len].copy_from_slice(&datagram[..len]); + Ok((len, sockaddr_from_wasi(remote_address))) + } } /// Associates this socket with a remote address, giving a [`UdpStream`] @@ -73,24 +105,38 @@ impl UdpSocket { /// This only changes the local socket configuration, and does not generate /// any network traffic. pub fn connect(self, addr: SocketAddr) -> io::Result { - // WASI may trap if streams from a previous call to `stream` are still - // live, so drop the unconnected streams before creating connected ones. - let Self { - incoming, - outgoing, - socket, - } = self; - drop((incoming, outgoing)); - - let (incoming, outgoing) = socket - .stream(Some(sockaddr_to_wasi(addr))) - .map_err(to_io_err)?; - Ok(UdpStream::new(incoming, outgoing, socket)) + #[cfg(target_env = "p2")] + { + // WASI may trap if streams from a previous call to `stream` are still + // live, so drop the unconnected streams before creating connected ones. + let Self { + incoming, + outgoing, + socket, + } = self; + drop((incoming, outgoing)); + + let (incoming, outgoing) = socket + .stream(Some(sockaddr_to_wasi(addr))) + .map_err(to_io_err)?; + Ok(UdpStream::new(incoming, outgoing, socket)) + } + #[cfg(target_env = "p3")] + { + self.socket + .connect(sockaddr_to_wasi(addr)) + .map_err(to_io_err)?; + Ok(UdpStream::new(self.socket)) + } } /// Returns the unicast hop limit ("time to live") of this socket. pub fn unicast_hop_limit(&self) -> io::Result { - self.socket.unicast_hop_limit().map_err(to_io_err) + #[cfg(target_env = "p2")] + let result = self.socket.unicast_hop_limit(); + #[cfg(target_env = "p3")] + let result = self.socket.get_unicast_hop_limit(); + result.map_err(to_io_err) } /// Sets the unicast hop limit ("time to live") of this socket. @@ -100,7 +146,11 @@ impl UdpSocket { /// Returns the size of the receive buffer of this socket. pub fn receive_buffer_size(&self) -> io::Result { - self.socket.receive_buffer_size().map_err(to_io_err) + #[cfg(target_env = "p2")] + let result = self.socket.receive_buffer_size(); + #[cfg(target_env = "p3")] + let result = self.socket.get_receive_buffer_size(); + result.map_err(to_io_err) } /// Sets the size of the receive buffer of this socket. This is a hint: the @@ -113,7 +163,11 @@ impl UdpSocket { /// Returns the size of the send buffer of this socket. pub fn send_buffer_size(&self) -> io::Result { - self.socket.send_buffer_size().map_err(to_io_err) + #[cfg(target_env = "p2")] + let result = self.socket.send_buffer_size(); + #[cfg(target_env = "p3")] + let result = self.socket.get_send_buffer_size(); + result.map_err(to_io_err) } /// Sets the size of the send buffer of this socket. This is a hint: the @@ -130,16 +184,19 @@ impl UdpSocket { /// any other address are not received. #[derive(Debug)] pub struct UdpStream { + #[cfg(target_env = "p2")] incoming: AsyncIncomingDatagramStream, + #[cfg(target_env = "p2")] outgoing: AsyncOutgoingDatagramStream, - socket: wasip2::sockets::udp::UdpSocket, + socket: WasiUdpSocket, } impl UdpStream { + #[cfg(target_env = "p2")] fn new( incoming: IncomingDatagramStream, outgoing: OutgoingDatagramStream, - socket: wasip2::sockets::udp::UdpSocket, + socket: WasiUdpSocket, ) -> Self { Self { incoming: AsyncIncomingDatagramStream::new(incoming), @@ -148,6 +205,11 @@ impl UdpStream { } } + #[cfg(target_env = "p3")] + fn new(socket: WasiUdpSocket) -> Self { + Self { socket } + } + /// Associates a UDP socket with a remote host. pub async fn connect(addr: impl ToSocketAddrs) -> io::Result { let addrs = addr.to_socket_addrs()?; @@ -176,31 +238,50 @@ impl UdpStream { }; let socket = bind_socket(local_addr).await?; - let (incoming, outgoing) = socket - .stream(Some(sockaddr_to_wasi(addr))) - .map_err(to_io_err)?; - Ok(Self::new(incoming, outgoing, socket)) + #[cfg(target_env = "p2")] + { + let (incoming, outgoing) = socket + .stream(Some(sockaddr_to_wasi(addr))) + .map_err(to_io_err)?; + Ok(Self::new(incoming, outgoing, socket)) + } + #[cfg(target_env = "p3")] + { + socket.connect(sockaddr_to_wasi(addr)).map_err(to_io_err)?; + Ok(Self::new(socket)) + } } /// Returns the local socket address of this socket. pub fn local_addr(&self) -> io::Result { - self.socket - .local_address() - .map_err(to_io_err) - .map(sockaddr_from_wasi) + #[cfg(target_env = "p2")] + let addr = self.socket.local_address(); + #[cfg(target_env = "p3")] + let addr = self.socket.get_local_address(); + addr.map_err(to_io_err).map(sockaddr_from_wasi) } /// Returns the socket address of the remote peer of this UDP association. pub fn peer_addr(&self) -> io::Result { - self.socket - .remote_address() - .map_err(to_io_err) - .map(sockaddr_from_wasi) + #[cfg(target_env = "p2")] + let addr = self.socket.remote_address(); + #[cfg(target_env = "p3")] + let addr = self.socket.get_remote_address(); + addr.map_err(to_io_err).map(sockaddr_from_wasi) } /// Sends a datagram to the remote peer. pub async fn send(&self, buf: &[u8]) -> io::Result { - self.outgoing.send_to(buf, None).await + #[cfg(target_env = "p2")] + return self.outgoing.send_to(buf, None).await; + #[cfg(target_env = "p3")] + { + self.socket + .send(buf.to_vec(), None) + .await + .map_err(to_io_err)?; + Ok(buf.len()) + } } /// Receives a single datagram from the remote peer. On success, returns the @@ -208,12 +289,24 @@ impl UdpStream { /// /// If `buf` is shorter than the datagram, the excess bytes are discarded. pub async fn recv(&self, buf: &mut [u8]) -> io::Result { - self.incoming.recv_from(buf).await.map(|(len, _addr)| len) + #[cfg(target_env = "p2")] + return self.incoming.recv_from(buf).await.map(|(len, _addr)| len); + #[cfg(target_env = "p3")] + { + let (datagram, _remote_address) = self.socket.receive().await.map_err(to_io_err)?; + let len = datagram.len().min(buf.len()); + buf[..len].copy_from_slice(&datagram[..len]); + Ok(len) + } } /// Returns the unicast hop limit ("time to live") of this socket. pub fn unicast_hop_limit(&self) -> io::Result { - self.socket.unicast_hop_limit().map_err(to_io_err) + #[cfg(target_env = "p2")] + let result = self.socket.unicast_hop_limit(); + #[cfg(target_env = "p3")] + let result = self.socket.get_unicast_hop_limit(); + result.map_err(to_io_err) } /// Sets the unicast hop limit ("time to live") of this socket. @@ -223,7 +316,11 @@ impl UdpStream { /// Returns the size of the receive buffer of this socket. pub fn receive_buffer_size(&self) -> io::Result { - self.socket.receive_buffer_size().map_err(to_io_err) + #[cfg(target_env = "p2")] + let result = self.socket.receive_buffer_size(); + #[cfg(target_env = "p3")] + let result = self.socket.get_receive_buffer_size(); + result.map_err(to_io_err) } /// Sets the size of the receive buffer of this socket. This is a hint: the @@ -236,7 +333,11 @@ impl UdpStream { /// Returns the size of the send buffer of this socket. pub fn send_buffer_size(&self) -> io::Result { - self.socket.send_buffer_size().map_err(to_io_err) + #[cfg(target_env = "p2")] + let result = self.socket.send_buffer_size(); + #[cfg(target_env = "p3")] + let result = self.socket.get_send_buffer_size(); + result.map_err(to_io_err) } /// Sets the size of the send buffer of this socket. This is a hint: the @@ -246,31 +347,38 @@ impl UdpStream { } } -async fn bind_socket(addr: SocketAddr) -> io::Result { +async fn bind_socket(addr: SocketAddr) -> io::Result { let family = match addr { SocketAddr::V4(_) => IpAddressFamily::Ipv4, SocketAddr::V6(_) => IpAddressFamily::Ipv6, }; let socket = create_udp_socket(family).map_err(to_io_err)?; - let network = instance_network(); let local_address = sockaddr_to_wasi(addr); - socket - .start_bind(&network, local_address) - .map_err(to_io_err)?; - let pollable = AsyncPollable::new(socket.subscribe()); - pollable.wait_for().await; - socket.finish_bind().map_err(to_io_err)?; + #[cfg(target_env = "p2")] + { + let network = instance_network(); + socket + .start_bind(&network, local_address) + .map_err(to_io_err)?; + let pollable = AsyncPollable::new(socket.subscribe()); + pollable.wait_for().await; + socket.finish_bind().map_err(to_io_err)?; + } + #[cfg(target_env = "p3")] + socket.bind(local_address).map_err(to_io_err)?; Ok(socket) } +#[cfg(target_env = "p2")] #[derive(Debug)] struct AsyncIncomingDatagramStream { subscription: OnceLock, stream: IncomingDatagramStream, } +#[cfg(target_env = "p2")] impl AsyncIncomingDatagramStream { fn new(stream: IncomingDatagramStream) -> Self { Self { @@ -310,12 +418,14 @@ impl AsyncIncomingDatagramStream { } } +#[cfg(target_env = "p2")] #[derive(Debug)] struct AsyncOutgoingDatagramStream { subscription: OnceLock, stream: OutgoingDatagramStream, } +#[cfg(target_env = "p2")] impl AsyncOutgoingDatagramStream { fn new(stream: OutgoingDatagramStream) -> Self { Self { diff --git a/test-programs/tests/tcp_echo_server.rs b/test-programs/tests/tcp_echo_server.rs index bb007dd..a27b607 100644 --- a/test-programs/tests/tcp_echo_server.rs +++ b/test-programs/tests/tcp_echo_server.rs @@ -1,20 +1,22 @@ use anyhow::{Context, Result}; use std::process::Command; -#[test_log::test] -fn tcp_echo_server() -> Result<()> { +fn run(component: &str, p3: bool) -> Result<()> { use std::io::{Read, Write}; use std::net::{Shutdown, TcpStream}; use test_programs::get_listening_address; - println!("testing {}", test_programs::TCP_ECHO_SERVER); + println!("testing {component}"); // Run the component in wasmtime // -Sinherit-network required for sockets to work - let mut wasmtime_process = Command::new("wasmtime") - .arg("run") - .arg("-Sinherit-network") - .arg(test_programs::TCP_ECHO_SERVER) + let mut command = Command::new("wasmtime"); + command.arg("run").arg("-Sinherit-network"); + if p3 { + command.arg("-Sp3"); + } + let mut wasmtime_process = command + .arg(component) .stdout(std::process::Stdio::piped()) .spawn()?; @@ -83,3 +85,16 @@ fn tcp_echo_server() -> Result<()> { Ok(()) } + +#[test_log::test] +fn tcp_echo_server_p2() -> Result<()> { + run(test_programs::TCP_ECHO_SERVER, false) +} + +#[test_log::test] +fn tcp_echo_server_p3() -> Result<()> { + if test_programs::NIGHTLY_TOOLCHAIN { + run(test_programs::TCP_ECHO_SERVER_P3, true)?; + } + Ok(()) +} diff --git a/test-programs/tests/tcp_stream_client.rs b/test-programs/tests/tcp_stream_client.rs index f3a87d6..23a3d59 100644 --- a/test-programs/tests/tcp_stream_client.rs +++ b/test-programs/tests/tcp_stream_client.rs @@ -2,19 +2,21 @@ use anyhow::{Context, Result}; use std::net::{Shutdown, TcpListener}; use std::process::{Command, Stdio}; -#[test_log::test] -fn tcp_stream_client() -> Result<()> { +fn run(component: &str, p3: bool) -> Result<()> { use std::io::{Read, Write}; - let server = TcpListener::bind("127.0.0.1:8082").context("binding temporary test server")?; + let server = TcpListener::bind("127.0.0.1:0").context("binding temporary test server")?; let addr = server .local_addr() .context("getting local listener address")?; - let child = Command::new("wasmtime") - .arg("run") - .arg("-Sinherit-network") - .arg(test_programs::TCP_STREAM_CLIENT) + let mut command = Command::new("wasmtime"); + command.arg("run").arg("-Sinherit-network"); + if p3 { + command.arg("-Sp3"); + } + let child = command + .arg(component) .arg(addr.to_string()) .stdout(Stdio::piped()) .spawn() @@ -51,3 +53,16 @@ fn tcp_stream_client() -> Result<()> { Ok(()) } + +#[test_log::test] +fn tcp_stream_client_p2() -> Result<()> { + run(test_programs::TCP_STREAM_CLIENT, false) +} + +#[test_log::test] +fn tcp_stream_client_p3() -> Result<()> { + if test_programs::NIGHTLY_TOOLCHAIN { + run(test_programs::TCP_STREAM_CLIENT_P3, true)?; + } + Ok(()) +} diff --git a/test-programs/tests/udp_echo_server.rs b/test-programs/tests/udp_echo_server.rs index 8ad8c8a..dd689b5 100644 --- a/test-programs/tests/udp_echo_server.rs +++ b/test-programs/tests/udp_echo_server.rs @@ -1,19 +1,21 @@ use anyhow::{Context, Result}; use std::process::Command; -#[test_log::test] -fn udp_echo_server() -> Result<()> { +fn run(component: &str, p3: bool) -> Result<()> { use std::net::{SocketAddr, UdpSocket}; use std::time::Duration; - println!("testing {}", test_programs::UDP_ECHO_SERVER); + println!("testing {component}"); // Run the component in wasmtime // -Sinherit-network required for sockets to work - let mut wasmtime_process = Command::new("wasmtime") - .arg("run") - .arg("-Sinherit-network") - .arg(test_programs::UDP_ECHO_SERVER) + let mut command = Command::new("wasmtime"); + command.arg("run").arg("-Sinherit-network"); + if p3 { + command.arg("-Sp3"); + } + let mut wasmtime_process = command + .arg(component) .stdout(std::process::Stdio::piped()) .spawn()?; @@ -64,3 +66,16 @@ fn udp_echo_server() -> Result<()> { Ok(()) } + +#[test_log::test] +fn udp_echo_server_p2() -> Result<()> { + run(test_programs::UDP_ECHO_SERVER, false) +} + +#[test_log::test] +fn udp_echo_server_p3() -> Result<()> { + if test_programs::NIGHTLY_TOOLCHAIN { + run(test_programs::UDP_ECHO_SERVER_P3, true)?; + } + Ok(()) +} diff --git a/test-programs/tests/udp_stream_client.rs b/test-programs/tests/udp_stream_client.rs index ebc4027..e365557 100644 --- a/test-programs/tests/udp_stream_client.rs +++ b/test-programs/tests/udp_stream_client.rs @@ -3,8 +3,7 @@ use std::net::UdpSocket; use std::process::{Command, Stdio}; use std::time::Duration; -#[test_log::test] -fn udp_stream_client() -> Result<()> { +fn run(component: &str, p3: bool) -> Result<()> { // Port 0: the host picks a free port, which the component is told about // by argument, so this test can't collide with anything else running. let server = UdpSocket::bind("127.0.0.1:0").context("binding temporary test server")?; @@ -17,10 +16,13 @@ fn udp_stream_client() -> Result<()> { .local_addr() .context("getting local server address")?; - let child = Command::new("wasmtime") - .arg("run") - .arg("-Sinherit-network") - .arg(test_programs::UDP_STREAM_CLIENT) + let mut command = Command::new("wasmtime"); + command.arg("run").arg("-Sinherit-network"); + if p3 { + command.arg("-Sp3"); + } + let child = command + .arg(component) .arg(addr.to_string()) .stderr(Stdio::piped()) .spawn() @@ -54,3 +56,16 @@ fn udp_stream_client() -> Result<()> { Ok(()) } + +#[test_log::test] +fn udp_stream_client_p2() -> Result<()> { + run(test_programs::UDP_STREAM_CLIENT, false) +} + +#[test_log::test] +fn udp_stream_client_p3() -> Result<()> { + if test_programs::NIGHTLY_TOOLCHAIN { + run(test_programs::UDP_STREAM_CLIENT_P3, true)?; + } + Ok(()) +} From eae3955c1b49fd6a45004b8bca92eed2f2d9ad1b Mon Sep 17 00:00:00 2001 From: Adam Bratschi-Kaye Date: Wed, 30 Sep 2026 20:09:15 +0000 Subject: [PATCH 2/4] cleanup --- src/net/udp.rs | 116 ++++++++++++++--------- test-programs/tests/tcp_echo_server.rs | 6 +- test-programs/tests/tcp_stream_client.rs | 6 +- test-programs/tests/udp_echo_server.rs | 6 +- test-programs/tests/udp_stream_client.rs | 6 +- 5 files changed, 79 insertions(+), 61 deletions(-) diff --git a/src/net/udp.rs b/src/net/udp.rs index 7728f30..b5a1ea7 100644 --- a/src/net/udp.rs +++ b/src/net/udp.rs @@ -19,6 +19,68 @@ use crate::io; #[cfg(target_env = "p2")] use crate::runtime::AsyncPollable; +#[cfg(target_env = "p2")] +mod getters { + use super::*; + + pub(super) fn local_address(socket: &WasiUdpSocket) -> io::Result { + socket + .local_address() + .map_err(to_io_err) + .map(sockaddr_from_wasi) + } + + pub(super) fn remote_address(socket: &WasiUdpSocket) -> io::Result { + socket + .remote_address() + .map_err(to_io_err) + .map(sockaddr_from_wasi) + } + + pub(super) fn unicast_hop_limit(socket: &WasiUdpSocket) -> io::Result { + socket.unicast_hop_limit().map_err(to_io_err) + } + + pub(super) fn receive_buffer_size(socket: &WasiUdpSocket) -> io::Result { + socket.receive_buffer_size().map_err(to_io_err) + } + + pub(super) fn send_buffer_size(socket: &WasiUdpSocket) -> io::Result { + socket.send_buffer_size().map_err(to_io_err) + } +} + +#[cfg(target_env = "p3")] +mod getters { + use super::*; + + pub(super) fn local_address(socket: &WasiUdpSocket) -> io::Result { + socket + .get_local_address() + .map_err(to_io_err) + .map(sockaddr_from_wasi) + } + + pub(super) fn remote_address(socket: &WasiUdpSocket) -> io::Result { + socket + .get_remote_address() + .map_err(to_io_err) + .map(sockaddr_from_wasi) + } + + pub(super) fn unicast_hop_limit(socket: &WasiUdpSocket) -> io::Result { + socket.get_unicast_hop_limit().map_err(to_io_err) + } + + pub(super) fn receive_buffer_size(socket: &WasiUdpSocket) -> io::Result { + socket.get_receive_buffer_size().map_err(to_io_err) + } + + pub(super) fn send_buffer_size(socket: &WasiUdpSocket) -> io::Result { + socket.get_send_buffer_size().map_err(to_io_err) + } +} + /// A UDP socket, bound to a local address. /// /// A `UdpSocket` is not associated with any remote address, so datagrams can be @@ -59,11 +121,7 @@ impl UdpSocket { /// Returns the local socket address of this socket. pub fn local_addr(&self) -> io::Result { - #[cfg(target_env = "p2")] - let addr = self.socket.local_address(); - #[cfg(target_env = "p3")] - let addr = self.socket.get_local_address(); - addr.map_err(to_io_err).map(sockaddr_from_wasi) + getters::local_address(&self.socket) } /// Sends a datagram to the given address. @@ -132,11 +190,7 @@ impl UdpSocket { /// Returns the unicast hop limit ("time to live") of this socket. pub fn unicast_hop_limit(&self) -> io::Result { - #[cfg(target_env = "p2")] - let result = self.socket.unicast_hop_limit(); - #[cfg(target_env = "p3")] - let result = self.socket.get_unicast_hop_limit(); - result.map_err(to_io_err) + getters::unicast_hop_limit(&self.socket) } /// Sets the unicast hop limit ("time to live") of this socket. @@ -146,11 +200,7 @@ impl UdpSocket { /// Returns the size of the receive buffer of this socket. pub fn receive_buffer_size(&self) -> io::Result { - #[cfg(target_env = "p2")] - let result = self.socket.receive_buffer_size(); - #[cfg(target_env = "p3")] - let result = self.socket.get_receive_buffer_size(); - result.map_err(to_io_err) + getters::receive_buffer_size(&self.socket) } /// Sets the size of the receive buffer of this socket. This is a hint: the @@ -163,11 +213,7 @@ impl UdpSocket { /// Returns the size of the send buffer of this socket. pub fn send_buffer_size(&self) -> io::Result { - #[cfg(target_env = "p2")] - let result = self.socket.send_buffer_size(); - #[cfg(target_env = "p3")] - let result = self.socket.get_send_buffer_size(); - result.map_err(to_io_err) + getters::send_buffer_size(&self.socket) } /// Sets the size of the send buffer of this socket. This is a hint: the @@ -254,20 +300,12 @@ impl UdpStream { /// Returns the local socket address of this socket. pub fn local_addr(&self) -> io::Result { - #[cfg(target_env = "p2")] - let addr = self.socket.local_address(); - #[cfg(target_env = "p3")] - let addr = self.socket.get_local_address(); - addr.map_err(to_io_err).map(sockaddr_from_wasi) + getters::local_address(&self.socket) } /// Returns the socket address of the remote peer of this UDP association. pub fn peer_addr(&self) -> io::Result { - #[cfg(target_env = "p2")] - let addr = self.socket.remote_address(); - #[cfg(target_env = "p3")] - let addr = self.socket.get_remote_address(); - addr.map_err(to_io_err).map(sockaddr_from_wasi) + getters::remote_address(&self.socket) } /// Sends a datagram to the remote peer. @@ -302,11 +340,7 @@ impl UdpStream { /// Returns the unicast hop limit ("time to live") of this socket. pub fn unicast_hop_limit(&self) -> io::Result { - #[cfg(target_env = "p2")] - let result = self.socket.unicast_hop_limit(); - #[cfg(target_env = "p3")] - let result = self.socket.get_unicast_hop_limit(); - result.map_err(to_io_err) + getters::unicast_hop_limit(&self.socket) } /// Sets the unicast hop limit ("time to live") of this socket. @@ -316,11 +350,7 @@ impl UdpStream { /// Returns the size of the receive buffer of this socket. pub fn receive_buffer_size(&self) -> io::Result { - #[cfg(target_env = "p2")] - let result = self.socket.receive_buffer_size(); - #[cfg(target_env = "p3")] - let result = self.socket.get_receive_buffer_size(); - result.map_err(to_io_err) + getters::receive_buffer_size(&self.socket) } /// Sets the size of the receive buffer of this socket. This is a hint: the @@ -333,11 +363,7 @@ impl UdpStream { /// Returns the size of the send buffer of this socket. pub fn send_buffer_size(&self) -> io::Result { - #[cfg(target_env = "p2")] - let result = self.socket.send_buffer_size(); - #[cfg(target_env = "p3")] - let result = self.socket.get_send_buffer_size(); - result.map_err(to_io_err) + getters::send_buffer_size(&self.socket) } /// Sets the size of the send buffer of this socket. This is a hint: the diff --git a/test-programs/tests/tcp_echo_server.rs b/test-programs/tests/tcp_echo_server.rs index a27b607..6e88c1b 100644 --- a/test-programs/tests/tcp_echo_server.rs +++ b/test-programs/tests/tcp_echo_server.rs @@ -91,10 +91,8 @@ fn tcp_echo_server_p2() -> Result<()> { run(test_programs::TCP_ECHO_SERVER, false) } +#[cfg(wstd_nightly)] #[test_log::test] fn tcp_echo_server_p3() -> Result<()> { - if test_programs::NIGHTLY_TOOLCHAIN { - run(test_programs::TCP_ECHO_SERVER_P3, true)?; - } - Ok(()) + run(test_programs::TCP_ECHO_SERVER_P3, true) } diff --git a/test-programs/tests/tcp_stream_client.rs b/test-programs/tests/tcp_stream_client.rs index 23a3d59..3725ab9 100644 --- a/test-programs/tests/tcp_stream_client.rs +++ b/test-programs/tests/tcp_stream_client.rs @@ -59,10 +59,8 @@ fn tcp_stream_client_p2() -> Result<()> { run(test_programs::TCP_STREAM_CLIENT, false) } +#[cfg(wstd_nightly)] #[test_log::test] fn tcp_stream_client_p3() -> Result<()> { - if test_programs::NIGHTLY_TOOLCHAIN { - run(test_programs::TCP_STREAM_CLIENT_P3, true)?; - } - Ok(()) + run(test_programs::TCP_STREAM_CLIENT_P3, true) } diff --git a/test-programs/tests/udp_echo_server.rs b/test-programs/tests/udp_echo_server.rs index dd689b5..8c638cd 100644 --- a/test-programs/tests/udp_echo_server.rs +++ b/test-programs/tests/udp_echo_server.rs @@ -72,10 +72,8 @@ fn udp_echo_server_p2() -> Result<()> { run(test_programs::UDP_ECHO_SERVER, false) } +#[cfg(wstd_nightly)] #[test_log::test] fn udp_echo_server_p3() -> Result<()> { - if test_programs::NIGHTLY_TOOLCHAIN { - run(test_programs::UDP_ECHO_SERVER_P3, true)?; - } - Ok(()) + run(test_programs::UDP_ECHO_SERVER_P3, true) } diff --git a/test-programs/tests/udp_stream_client.rs b/test-programs/tests/udp_stream_client.rs index e365557..565854f 100644 --- a/test-programs/tests/udp_stream_client.rs +++ b/test-programs/tests/udp_stream_client.rs @@ -62,10 +62,8 @@ fn udp_stream_client_p2() -> Result<()> { run(test_programs::UDP_STREAM_CLIENT, false) } +#[cfg(wstd_nightly)] #[test_log::test] fn udp_stream_client_p3() -> Result<()> { - if test_programs::NIGHTLY_TOOLCHAIN { - run(test_programs::UDP_STREAM_CLIENT_P3, true)?; - } - Ok(()) + run(test_programs::UDP_STREAM_CLIENT_P3, true) } From 0c4125fbf6ed62614abd96a011978a90ed440fcd Mon Sep 17 00:00:00 2001 From: Adam Bratschi-Kaye Date: Fri, 2 Oct 2026 13:24:59 +0000 Subject: [PATCH 3/4] Switch to sys modules --- src/net/mod.rs | 41 ++- .../sys/p2.rs} | 56 +--- src/net/tcp_listener/sys/p3.rs | 71 ++++++ .../{tcp_stream.rs => tcp_stream/sys/p2.rs} | 48 +--- src/net/tcp_stream/sys/p3.rs | 146 +++++++++++ src/net/{udp.rs => udp/sys/p2.rs} | 240 ++++-------------- src/net/udp/sys/p3.rs | 232 +++++++++++++++++ 7 files changed, 561 insertions(+), 273 deletions(-) rename src/net/{tcp_listener.rs => tcp_listener/sys/p2.rs} (57%) create mode 100644 src/net/tcp_listener/sys/p3.rs rename src/net/{tcp_stream.rs => tcp_stream/sys/p2.rs} (76%) create mode 100644 src/net/tcp_stream/sys/p3.rs rename src/net/{udp.rs => udp/sys/p2.rs} (63%) create mode 100644 src/net/udp/sys/p3.rs diff --git a/src/net/mod.rs b/src/net/mod.rs index f18287e..b0c5d96 100644 --- a/src/net/mod.rs +++ b/src/net/mod.rs @@ -12,9 +12,44 @@ use wasip3::sockets::types::{ ErrorCode, IpAddressFamily, IpSocketAddress, Ipv4SocketAddress, Ipv6SocketAddress, TcpSocket, }; -mod tcp_listener; -mod tcp_stream; -mod udp; +mod tcp_listener { + mod sys { + #[cfg(target_env = "p2")] + pub(super) mod p2; + #[cfg(target_env = "p3")] + pub(super) mod p3; + } + #[cfg(target_env = "p2")] + pub use sys::p2::*; + #[cfg(target_env = "p3")] + pub use sys::p3::*; +} + +mod tcp_stream { + mod sys { + #[cfg(target_env = "p2")] + pub(super) mod p2; + #[cfg(target_env = "p3")] + pub(super) mod p3; + } + #[cfg(target_env = "p2")] + pub use sys::p2::*; + #[cfg(target_env = "p3")] + pub use sys::p3::*; +} + +mod udp { + mod sys { + #[cfg(target_env = "p2")] + pub(super) mod p2; + #[cfg(target_env = "p3")] + pub(super) mod p3; + } + #[cfg(target_env = "p2")] + pub use sys::p2::*; + #[cfg(target_env = "p3")] + pub use sys::p3::*; +} pub use tcp_listener::*; pub use tcp_stream::*; diff --git a/src/net/tcp_listener.rs b/src/net/tcp_listener/sys/p2.rs similarity index 57% rename from src/net/tcp_listener.rs rename to src/net/tcp_listener/sys/p2.rs index d4cfdae..1c8670a 100644 --- a/src/net/tcp_listener.rs +++ b/src/net/tcp_listener/sys/p2.rs @@ -1,17 +1,10 @@ -#[cfg(target_env = "p2")] use wasip2::sockets::tcp::{IpAddressFamily, TcpSocket}; -#[cfg(target_env = "p3")] -use wasip3::{ - sockets::types::{IpAddressFamily, TcpSocket}, - wit_bindgen::StreamReader, -}; use crate::io; use crate::iter::AsyncIterator; use std::net::SocketAddr; -use super::{TcpStream, create_tcp_socket, sockaddr_from_wasi, sockaddr_to_wasi, to_io_err}; -#[cfg(target_env = "p2")] +use crate::net::{TcpStream, create_tcp_socket, sockaddr_from_wasi, sockaddr_to_wasi, to_io_err}; use crate::runtime::AsyncPollable; /// A TCP socket server, listening for connections. @@ -20,8 +13,6 @@ pub struct TcpListener { // Field order matters: must drop this child before parent below #[cfg(target_env = "p2")] pollable: AsyncPollable, - #[cfg(target_env = "p3")] - connections: StreamReader, socket: TcpSocket, } @@ -40,30 +31,18 @@ impl TcpListener { let socket = create_tcp_socket(family).map_err(to_io_err)?; let local_address = sockaddr_to_wasi(addr); - #[cfg(target_env = "p2")] - { - let network = wasip2::sockets::instance_network::instance_network(); - socket - .start_bind(&network, local_address) - .map_err(to_io_err)?; - let pollable = AsyncPollable::new(socket.subscribe()); - pollable.wait_for().await; - socket.finish_bind().map_err(to_io_err)?; + let network = wasip2::sockets::instance_network::instance_network(); + socket + .start_bind(&network, local_address) + .map_err(to_io_err)?; + let pollable = AsyncPollable::new(socket.subscribe()); + pollable.wait_for().await; + socket.finish_bind().map_err(to_io_err)?; - socket.start_listen().map_err(to_io_err)?; - pollable.wait_for().await; - socket.finish_listen().map_err(to_io_err)?; - Ok(Self { pollable, socket }) - } - #[cfg(target_env = "p3")] - { - socket.bind(local_address).map_err(to_io_err)?; - let connections = socket.listen().map_err(to_io_err)?; - Ok(Self { - connections, - socket, - }) - } + socket.start_listen().map_err(to_io_err)?; + pollable.wait_for().await; + socket.finish_listen().map_err(to_io_err)?; + Ok(Self { pollable, socket }) } /// Returns the local socket address of this listener. @@ -90,7 +69,6 @@ pub struct Incoming<'a> { impl<'a> AsyncIterator for Incoming<'a> { type Item = io::Result; - #[cfg(target_env = "p2")] async fn next(&mut self) -> Option { self.listener.pollable.wait_for().await; let (socket, input, output) = match self.listener.socket.accept().map_err(to_io_err) { @@ -99,14 +77,4 @@ impl<'a> AsyncIterator for Incoming<'a> { }; Some(Ok(TcpStream::new(input, output, socket))) } - - #[cfg(target_env = "p3")] - async fn next(&mut self) -> Option { - self.listener.connections.next().await.map(|socket| { - let (input, _receive_result) = socket.receive(); - let (output, receiver) = wasip3::wit_stream::new(); - let _send_result = socket.send(receiver); - Ok(TcpStream::new(input, output, socket)) - }) - } } diff --git a/src/net/tcp_listener/sys/p3.rs b/src/net/tcp_listener/sys/p3.rs new file mode 100644 index 0000000..9b5c7e0 --- /dev/null +++ b/src/net/tcp_listener/sys/p3.rs @@ -0,0 +1,71 @@ +use wasip3::{ + sockets::types::{IpAddressFamily, TcpSocket}, + wit_bindgen::StreamReader, +}; + +use crate::io; +use crate::iter::AsyncIterator; +use std::net::SocketAddr; + +use crate::net::{TcpStream, create_tcp_socket, sockaddr_from_wasi, sockaddr_to_wasi, to_io_err}; + +/// A TCP socket server, listening for connections. +#[derive(Debug)] +pub struct TcpListener { + connections: StreamReader, + socket: TcpSocket, +} + +impl TcpListener { + /// Creates a new TcpListener which will be bound to the specified address. + /// + /// The returned listener is ready for accepting connections. + pub async fn bind(addr: &str) -> io::Result { + let addr: SocketAddr = addr + .parse() + .map_err(|_| io::Error::other("failed to parse string to socket addr"))?; + let family = match addr { + SocketAddr::V4(_) => IpAddressFamily::Ipv4, + SocketAddr::V6(_) => IpAddressFamily::Ipv6, + }; + let socket = create_tcp_socket(family).map_err(to_io_err)?; + let local_address = sockaddr_to_wasi(addr); + + socket.bind(local_address).map_err(to_io_err)?; + let connections = socket.listen().map_err(to_io_err)?; + Ok(Self { + connections, + socket, + }) + } + + /// Returns the local socket address of this listener. + pub fn local_addr(&self) -> io::Result { + let addr = self.socket.get_local_address(); + addr.map_err(to_io_err).map(sockaddr_from_wasi) + } + + /// Returns an iterator over the connections being received on this listener. + pub fn incoming(&mut self) -> Incoming<'_> { + Incoming { listener: self } + } +} + +/// An iterator that infinitely accepts connections on a TcpListener. +#[derive(Debug)] +pub struct Incoming<'a> { + listener: &'a mut TcpListener, +} + +impl<'a> AsyncIterator for Incoming<'a> { + type Item = io::Result; + + async fn next(&mut self) -> Option { + self.listener.connections.next().await.map(|socket| { + let (input, _receive_result) = socket.receive(); + let (output, receiver) = wasip3::wit_stream::new(); + let _send_result = socket.send(receiver); + Ok(TcpStream::new(input, output, socket)) + }) + } +} diff --git a/src/net/tcp_stream.rs b/src/net/tcp_stream/sys/p2.rs similarity index 76% rename from src/net/tcp_stream.rs rename to src/net/tcp_stream/sys/p2.rs index 04dad77..79ffdd8 100644 --- a/src/net/tcp_stream.rs +++ b/src/net/tcp_stream/sys/p2.rs @@ -1,7 +1,6 @@ use std::io::ErrorKind; use std::net::{SocketAddr, ToSocketAddrs}; -#[cfg(target_env = "p2")] use wasip2::{ io::streams::{InputStream, OutputStream}, sockets::{ @@ -11,16 +10,8 @@ use wasip2::{ }, }; -#[cfg(target_env = "p3")] -use wasip3::sockets::types::{IpAddressFamily, IpSocketAddress, Ipv4SocketAddress, TcpSocket}; -#[cfg(target_env = "p3")] -type InputStream = wasip3::wit_bindgen::StreamReader; -#[cfg(target_env = "p3")] -type OutputStream = wasip3::wit_bindgen::StreamWriter; - -use super::{create_tcp_socket, to_io_err}; use crate::io::{self, AsyncInputStream, AsyncOutputStream}; -#[cfg(target_env = "p2")] +use crate::net::{create_tcp_socket, to_io_err}; use crate::runtime::AsyncPollable; /// A TCP stream between a local and a remote socket. @@ -79,33 +70,19 @@ impl TcpStream { } SocketAddr::V6(_) => todo!("IPv6 not yet supported in `wstd::net::TcpStream`"), }; - #[cfg(target_env = "p2")] - { - let network = instance_network(); - socket - .start_connect(&network, remote_address) - .map_err(to_io_err)?; - let pollable = AsyncPollable::new(socket.subscribe()); - pollable.wait_for().await; - let (input, output) = socket.finish_connect().map_err(to_io_err)?; - Ok(TcpStream::new(input, output, socket)) - } - #[cfg(target_env = "p3")] - { - socket.connect(remote_address).await.map_err(to_io_err)?; - let (input, _receive_result) = socket.receive(); - let (output, receiver) = wasip3::wit_stream::new(); - let _send_result = socket.send(receiver); - Ok(TcpStream::new(input, output, socket)) - } + let network = instance_network(); + socket + .start_connect(&network, remote_address) + .map_err(to_io_err)?; + let pollable = AsyncPollable::new(socket.subscribe()); + pollable.wait_for().await; + let (input, output) = socket.finish_connect().map_err(to_io_err)?; + Ok(TcpStream::new(input, output, socket)) } /// Returns the socket address of the remote peer of this TCP connection. pub fn peer_addr(&self) -> io::Result { - #[cfg(target_env = "p2")] let addr = self.socket.remote_address().map_err(to_io_err)?; - #[cfg(target_env = "p3")] - let addr = self.socket.get_remote_address().map_err(to_io_err)?; Ok(format!("{addr:?}")) } @@ -113,19 +90,16 @@ impl TcpStream { ( ReadHalf { stream: &mut self.input, - #[cfg(target_env = "p2")] socket: &self.socket, }, WriteHalf { stream: &mut self.output, - #[cfg(target_env = "p2")] socket: &self.socket, }, ) } } -#[cfg(target_env = "p2")] impl Drop for TcpStream { fn drop(&mut self) { let _ = self @@ -160,11 +134,9 @@ impl io::AsyncWrite for TcpStream { pub struct ReadHalf<'a> { stream: &'a mut AsyncInputStream, - #[cfg(target_env = "p2")] socket: &'a TcpSocket, } -#[cfg(target_env = "p2")] impl<'a> Drop for ReadHalf<'a> { fn drop(&mut self) { let _ = self @@ -185,7 +157,6 @@ impl<'a> io::AsyncRead for ReadHalf<'a> { pub struct WriteHalf<'a> { stream: &'a mut AsyncOutputStream, - #[cfg(target_env = "p2")] socket: &'a TcpSocket, } @@ -203,7 +174,6 @@ impl<'a> io::AsyncWrite for WriteHalf<'a> { } } -#[cfg(target_env = "p2")] impl<'a> Drop for WriteHalf<'a> { fn drop(&mut self) { let _ = self diff --git a/src/net/tcp_stream/sys/p3.rs b/src/net/tcp_stream/sys/p3.rs new file mode 100644 index 0000000..71c3a9b --- /dev/null +++ b/src/net/tcp_stream/sys/p3.rs @@ -0,0 +1,146 @@ +use std::io::ErrorKind; +use std::net::{SocketAddr, ToSocketAddrs}; + +use wasip3::sockets::types::{IpAddressFamily, IpSocketAddress, Ipv4SocketAddress, TcpSocket}; +type InputStream = wasip3::wit_bindgen::StreamReader; +type OutputStream = wasip3::wit_bindgen::StreamWriter; + +use crate::io::{self, AsyncInputStream, AsyncOutputStream}; +use crate::net::{create_tcp_socket, to_io_err}; + +/// A TCP stream between a local and a remote socket. +pub struct TcpStream { + input: AsyncInputStream, + output: AsyncOutputStream, + socket: TcpSocket, +} + +impl TcpStream { + pub(crate) fn new(input: InputStream, output: OutputStream, socket: TcpSocket) -> Self { + TcpStream { + input: AsyncInputStream::new(input), + output: AsyncOutputStream::new(output), + socket, + } + } + + /// Opens a TCP connection to a remote host. + /// + /// `addr` is an address of the remote host. Anything which implements the + /// [`ToSocketAddrs`] trait can be supplied as the address. If `addr` + /// yields multiple addresses, connect will be attempted with each of the + /// addresses until a connection is successful. If none of the addresses + /// result in a successful connection, the error returned from the last + /// connection attempt (the last address) is returned. + pub async fn connect(addr: impl ToSocketAddrs) -> io::Result { + let addrs = addr.to_socket_addrs()?; + let mut last_err = None; + for addr in addrs { + match TcpStream::connect_addr(addr).await { + Ok(stream) => return Ok(stream), + Err(e) => last_err = Some(e), + } + } + + Err(last_err.unwrap_or_else(|| { + io::Error::new(ErrorKind::InvalidInput, "could not resolve to any address") + })) + } + + /// Establishes a connection to the specified `addr`. + pub async fn connect_addr(addr: SocketAddr) -> io::Result { + let family = match addr { + SocketAddr::V4(_) => IpAddressFamily::Ipv4, + SocketAddr::V6(_) => IpAddressFamily::Ipv6, + }; + let socket = create_tcp_socket(family).map_err(to_io_err)?; + + let remote_address = match addr { + SocketAddr::V4(addr) => { + let ip = addr.ip().octets(); + let address = (ip[0], ip[1], ip[2], ip[3]); + let port = addr.port(); + IpSocketAddress::Ipv4(Ipv4SocketAddress { port, address }) + } + SocketAddr::V6(_) => todo!("IPv6 not yet supported in `wstd::net::TcpStream`"), + }; + socket.connect(remote_address).await.map_err(to_io_err)?; + let (input, _receive_result) = socket.receive(); + let (output, receiver) = wasip3::wit_stream::new(); + let _send_result = socket.send(receiver); + Ok(TcpStream::new(input, output, socket)) + } + + /// Returns the socket address of the remote peer of this TCP connection. + pub fn peer_addr(&self) -> io::Result { + let addr = self.socket.get_remote_address().map_err(to_io_err)?; + Ok(format!("{addr:?}")) + } + + pub fn split(&mut self) -> (ReadHalf<'_>, WriteHalf<'_>) { + ( + ReadHalf { + stream: &mut self.input, + }, + WriteHalf { + stream: &mut self.output, + }, + ) + } +} + +impl io::AsyncRead for TcpStream { + async fn read(&mut self, buf: &mut [u8]) -> io::Result { + self.input.read(buf).await + } + + fn as_async_input_stream(&mut self) -> Option<&mut AsyncInputStream> { + Some(&mut self.input) + } +} + +impl io::AsyncWrite for TcpStream { + async fn write(&mut self, buf: &[u8]) -> io::Result { + self.output.write(buf).await + } + + async fn flush(&mut self) -> io::Result<()> { + self.output.flush().await + } + + fn as_async_output_stream(&mut self) -> Option<&mut AsyncOutputStream> { + Some(&mut self.output) + } +} + +pub struct ReadHalf<'a> { + stream: &'a mut AsyncInputStream, +} + +impl<'a> io::AsyncRead for ReadHalf<'a> { + async fn read(&mut self, buf: &mut [u8]) -> io::Result { + self.stream.read(buf).await + } + + fn as_async_input_stream(&mut self) -> Option<&mut AsyncInputStream> { + self.stream.as_async_input_stream() + } +} + +pub struct WriteHalf<'a> { + stream: &'a mut AsyncOutputStream, +} + +impl<'a> io::AsyncWrite for WriteHalf<'a> { + async fn write(&mut self, buf: &[u8]) -> io::Result { + self.stream.write(buf).await + } + + async fn flush(&mut self) -> io::Result<()> { + self.stream.flush().await + } + + fn as_async_output_stream(&mut self) -> Option<&mut AsyncOutputStream> { + self.stream.as_async_output_stream() + } +} diff --git a/src/net/udp.rs b/src/net/udp/sys/p2.rs similarity index 63% rename from src/net/udp.rs rename to src/net/udp/sys/p2.rs index b5a1ea7..810a3ad 100644 --- a/src/net/udp.rs +++ b/src/net/udp/sys/p2.rs @@ -1,9 +1,7 @@ use std::io::ErrorKind; use std::net::{SocketAddr, ToSocketAddrs}; -#[cfg(target_env = "p2")] use std::sync::OnceLock; -#[cfg(target_env = "p2")] use wasip2::sockets::{ instance_network::instance_network, udp::{ @@ -11,76 +9,11 @@ use wasip2::sockets::{ OutgoingDatagramStream, UdpSocket as WasiUdpSocket, }, }; -#[cfg(target_env = "p3")] -use wasip3::sockets::types::{IpAddressFamily, UdpSocket as WasiUdpSocket}; -use super::{create_udp_socket, sockaddr_from_wasi, sockaddr_to_wasi, to_io_err}; use crate::io; -#[cfg(target_env = "p2")] +use crate::net::{create_udp_socket, sockaddr_from_wasi, sockaddr_to_wasi, to_io_err}; use crate::runtime::AsyncPollable; -#[cfg(target_env = "p2")] -mod getters { - use super::*; - - pub(super) fn local_address(socket: &WasiUdpSocket) -> io::Result { - socket - .local_address() - .map_err(to_io_err) - .map(sockaddr_from_wasi) - } - - pub(super) fn remote_address(socket: &WasiUdpSocket) -> io::Result { - socket - .remote_address() - .map_err(to_io_err) - .map(sockaddr_from_wasi) - } - - pub(super) fn unicast_hop_limit(socket: &WasiUdpSocket) -> io::Result { - socket.unicast_hop_limit().map_err(to_io_err) - } - - pub(super) fn receive_buffer_size(socket: &WasiUdpSocket) -> io::Result { - socket.receive_buffer_size().map_err(to_io_err) - } - - pub(super) fn send_buffer_size(socket: &WasiUdpSocket) -> io::Result { - socket.send_buffer_size().map_err(to_io_err) - } -} - -#[cfg(target_env = "p3")] -mod getters { - use super::*; - - pub(super) fn local_address(socket: &WasiUdpSocket) -> io::Result { - socket - .get_local_address() - .map_err(to_io_err) - .map(sockaddr_from_wasi) - } - - pub(super) fn remote_address(socket: &WasiUdpSocket) -> io::Result { - socket - .get_remote_address() - .map_err(to_io_err) - .map(sockaddr_from_wasi) - } - - pub(super) fn unicast_hop_limit(socket: &WasiUdpSocket) -> io::Result { - socket.get_unicast_hop_limit().map_err(to_io_err) - } - - pub(super) fn receive_buffer_size(socket: &WasiUdpSocket) -> io::Result { - socket.get_receive_buffer_size().map_err(to_io_err) - } - - pub(super) fn send_buffer_size(socket: &WasiUdpSocket) -> io::Result { - socket.get_send_buffer_size().map_err(to_io_err) - } -} - /// A UDP socket, bound to a local address. /// /// A `UdpSocket` is not associated with any remote address, so datagrams can be @@ -89,9 +22,7 @@ mod getters { /// single remote address instead, giving a [`UdpStream`]. #[derive(Debug)] pub struct UdpSocket { - #[cfg(target_env = "p2")] incoming: AsyncIncomingDatagramStream, - #[cfg(target_env = "p2")] outgoing: AsyncOutgoingDatagramStream, socket: WasiUdpSocket, } @@ -104,41 +35,30 @@ impl UdpSocket { .map_err(|_| io::Error::other("failed to parse string to socket addr"))?; let socket = bind_socket(addr).await?; - #[cfg(target_env = "p2")] - { - // Datagram streams without a remote address may send to, and receive - // from, any address. - let (incoming, outgoing) = socket.stream(None).map_err(to_io_err)?; - Ok(Self { - incoming: AsyncIncomingDatagramStream::new(incoming), - outgoing: AsyncOutgoingDatagramStream::new(outgoing), - socket, - }) - } - #[cfg(target_env = "p3")] - Ok(Self { socket }) + // Datagram streams without a remote address may send to, and receive + // from, any address. + let (incoming, outgoing) = socket.stream(None).map_err(to_io_err)?; + Ok(Self { + incoming: AsyncIncomingDatagramStream::new(incoming), + outgoing: AsyncOutgoingDatagramStream::new(outgoing), + socket, + }) } /// Returns the local socket address of this socket. pub fn local_addr(&self) -> io::Result { - getters::local_address(&self.socket) + self.socket + .local_address() + .map_err(to_io_err) + .map(sockaddr_from_wasi) } /// Sends a datagram to the given address. pub async fn send_to(&self, buf: &[u8], addr: SocketAddr) -> io::Result { - #[cfg(target_env = "p2")] return self .outgoing .send_to(buf, Some(sockaddr_to_wasi(addr))) .await; - #[cfg(target_env = "p3")] - { - self.socket - .send(buf.to_vec(), Some(sockaddr_to_wasi(addr))) - .await - .map_err(to_io_err)?; - Ok(buf.len()) - } } /// Receives a single datagram. On success, returns the number of bytes @@ -146,15 +66,7 @@ impl UdpSocket { /// /// If `buf` is shorter than the datagram, the excess bytes are discarded. pub async fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { - #[cfg(target_env = "p2")] return self.incoming.recv_from(buf).await; - #[cfg(target_env = "p3")] - { - let (datagram, remote_address) = self.socket.receive().await.map_err(to_io_err)?; - let len = datagram.len().min(buf.len()); - buf[..len].copy_from_slice(&datagram[..len]); - Ok((len, sockaddr_from_wasi(remote_address))) - } } /// Associates this socket with a remote address, giving a [`UdpStream`] @@ -163,34 +75,24 @@ impl UdpSocket { /// This only changes the local socket configuration, and does not generate /// any network traffic. pub fn connect(self, addr: SocketAddr) -> io::Result { - #[cfg(target_env = "p2")] - { - // WASI may trap if streams from a previous call to `stream` are still - // live, so drop the unconnected streams before creating connected ones. - let Self { - incoming, - outgoing, - socket, - } = self; - drop((incoming, outgoing)); - - let (incoming, outgoing) = socket - .stream(Some(sockaddr_to_wasi(addr))) - .map_err(to_io_err)?; - Ok(UdpStream::new(incoming, outgoing, socket)) - } - #[cfg(target_env = "p3")] - { - self.socket - .connect(sockaddr_to_wasi(addr)) - .map_err(to_io_err)?; - Ok(UdpStream::new(self.socket)) - } + // WASI may trap if streams from a previous call to `stream` are still + // live, so drop the unconnected streams before creating connected ones. + let Self { + incoming, + outgoing, + socket, + } = self; + drop((incoming, outgoing)); + + let (incoming, outgoing) = socket + .stream(Some(sockaddr_to_wasi(addr))) + .map_err(to_io_err)?; + Ok(UdpStream::new(incoming, outgoing, socket)) } /// Returns the unicast hop limit ("time to live") of this socket. pub fn unicast_hop_limit(&self) -> io::Result { - getters::unicast_hop_limit(&self.socket) + self.socket.unicast_hop_limit().map_err(to_io_err) } /// Sets the unicast hop limit ("time to live") of this socket. @@ -200,7 +102,7 @@ impl UdpSocket { /// Returns the size of the receive buffer of this socket. pub fn receive_buffer_size(&self) -> io::Result { - getters::receive_buffer_size(&self.socket) + self.socket.receive_buffer_size().map_err(to_io_err) } /// Sets the size of the receive buffer of this socket. This is a hint: the @@ -213,7 +115,7 @@ impl UdpSocket { /// Returns the size of the send buffer of this socket. pub fn send_buffer_size(&self) -> io::Result { - getters::send_buffer_size(&self.socket) + self.socket.send_buffer_size().map_err(to_io_err) } /// Sets the size of the send buffer of this socket. This is a hint: the @@ -230,15 +132,12 @@ impl UdpSocket { /// any other address are not received. #[derive(Debug)] pub struct UdpStream { - #[cfg(target_env = "p2")] incoming: AsyncIncomingDatagramStream, - #[cfg(target_env = "p2")] outgoing: AsyncOutgoingDatagramStream, socket: WasiUdpSocket, } impl UdpStream { - #[cfg(target_env = "p2")] fn new( incoming: IncomingDatagramStream, outgoing: OutgoingDatagramStream, @@ -251,11 +150,6 @@ impl UdpStream { } } - #[cfg(target_env = "p3")] - fn new(socket: WasiUdpSocket) -> Self { - Self { socket } - } - /// Associates a UDP socket with a remote host. pub async fn connect(addr: impl ToSocketAddrs) -> io::Result { let addrs = addr.to_socket_addrs()?; @@ -284,42 +178,31 @@ impl UdpStream { }; let socket = bind_socket(local_addr).await?; - #[cfg(target_env = "p2")] - { - let (incoming, outgoing) = socket - .stream(Some(sockaddr_to_wasi(addr))) - .map_err(to_io_err)?; - Ok(Self::new(incoming, outgoing, socket)) - } - #[cfg(target_env = "p3")] - { - socket.connect(sockaddr_to_wasi(addr)).map_err(to_io_err)?; - Ok(Self::new(socket)) - } + let (incoming, outgoing) = socket + .stream(Some(sockaddr_to_wasi(addr))) + .map_err(to_io_err)?; + Ok(Self::new(incoming, outgoing, socket)) } /// Returns the local socket address of this socket. pub fn local_addr(&self) -> io::Result { - getters::local_address(&self.socket) + self.socket + .local_address() + .map_err(to_io_err) + .map(sockaddr_from_wasi) } /// Returns the socket address of the remote peer of this UDP association. pub fn peer_addr(&self) -> io::Result { - getters::remote_address(&self.socket) + self.socket + .remote_address() + .map_err(to_io_err) + .map(sockaddr_from_wasi) } /// Sends a datagram to the remote peer. pub async fn send(&self, buf: &[u8]) -> io::Result { - #[cfg(target_env = "p2")] - return self.outgoing.send_to(buf, None).await; - #[cfg(target_env = "p3")] - { - self.socket - .send(buf.to_vec(), None) - .await - .map_err(to_io_err)?; - Ok(buf.len()) - } + self.outgoing.send_to(buf, None).await } /// Receives a single datagram from the remote peer. On success, returns the @@ -327,20 +210,12 @@ impl UdpStream { /// /// If `buf` is shorter than the datagram, the excess bytes are discarded. pub async fn recv(&self, buf: &mut [u8]) -> io::Result { - #[cfg(target_env = "p2")] - return self.incoming.recv_from(buf).await.map(|(len, _addr)| len); - #[cfg(target_env = "p3")] - { - let (datagram, _remote_address) = self.socket.receive().await.map_err(to_io_err)?; - let len = datagram.len().min(buf.len()); - buf[..len].copy_from_slice(&datagram[..len]); - Ok(len) - } + self.incoming.recv_from(buf).await.map(|(len, _addr)| len) } /// Returns the unicast hop limit ("time to live") of this socket. pub fn unicast_hop_limit(&self) -> io::Result { - getters::unicast_hop_limit(&self.socket) + self.socket.unicast_hop_limit().map_err(to_io_err) } /// Sets the unicast hop limit ("time to live") of this socket. @@ -350,7 +225,7 @@ impl UdpStream { /// Returns the size of the receive buffer of this socket. pub fn receive_buffer_size(&self) -> io::Result { - getters::receive_buffer_size(&self.socket) + self.socket.receive_buffer_size().map_err(to_io_err) } /// Sets the size of the receive buffer of this socket. This is a hint: the @@ -363,7 +238,7 @@ impl UdpStream { /// Returns the size of the send buffer of this socket. pub fn send_buffer_size(&self) -> io::Result { - getters::send_buffer_size(&self.socket) + self.socket.send_buffer_size().map_err(to_io_err) } /// Sets the size of the send buffer of this socket. This is a hint: the @@ -381,30 +256,23 @@ async fn bind_socket(addr: SocketAddr) -> io::Result { let socket = create_udp_socket(family).map_err(to_io_err)?; let local_address = sockaddr_to_wasi(addr); - #[cfg(target_env = "p2")] - { - let network = instance_network(); - socket - .start_bind(&network, local_address) - .map_err(to_io_err)?; - let pollable = AsyncPollable::new(socket.subscribe()); - pollable.wait_for().await; - socket.finish_bind().map_err(to_io_err)?; - } - #[cfg(target_env = "p3")] - socket.bind(local_address).map_err(to_io_err)?; + let network = instance_network(); + socket + .start_bind(&network, local_address) + .map_err(to_io_err)?; + let pollable = AsyncPollable::new(socket.subscribe()); + pollable.wait_for().await; + socket.finish_bind().map_err(to_io_err)?; Ok(socket) } -#[cfg(target_env = "p2")] #[derive(Debug)] struct AsyncIncomingDatagramStream { subscription: OnceLock, stream: IncomingDatagramStream, } -#[cfg(target_env = "p2")] impl AsyncIncomingDatagramStream { fn new(stream: IncomingDatagramStream) -> Self { Self { @@ -444,14 +312,12 @@ impl AsyncIncomingDatagramStream { } } -#[cfg(target_env = "p2")] #[derive(Debug)] struct AsyncOutgoingDatagramStream { subscription: OnceLock, stream: OutgoingDatagramStream, } -#[cfg(target_env = "p2")] impl AsyncOutgoingDatagramStream { fn new(stream: OutgoingDatagramStream) -> Self { Self { diff --git a/src/net/udp/sys/p3.rs b/src/net/udp/sys/p3.rs new file mode 100644 index 0000000..de4fffa --- /dev/null +++ b/src/net/udp/sys/p3.rs @@ -0,0 +1,232 @@ +use std::io::ErrorKind; +use std::net::{SocketAddr, ToSocketAddrs}; +use wasip3::sockets::types::{IpAddressFamily, UdpSocket as WasiUdpSocket}; + +use crate::io; +use crate::net::{create_udp_socket, sockaddr_from_wasi, sockaddr_to_wasi, to_io_err}; + +/// A UDP socket, bound to a local address. +/// +/// A `UdpSocket` is not associated with any remote address, so datagrams can be +/// sent to, and received from, any address, using [`UdpSocket::send_to`] and +/// [`UdpSocket::recv_from`]. Use [`UdpSocket::connect`] to associate it with a +/// single remote address instead, giving a [`UdpStream`]. +#[derive(Debug)] +pub struct UdpSocket { + socket: WasiUdpSocket, +} + +impl UdpSocket { + /// Creates a new UdpSocket bound to the specified local address. + pub async fn bind(addr: &str) -> io::Result { + let addr: SocketAddr = addr + .parse() + .map_err(|_| io::Error::other("failed to parse string to socket addr"))?; + let socket = bind_socket(addr).await?; + + Ok(Self { socket }) + } + + /// Returns the local socket address of this socket. + pub fn local_addr(&self) -> io::Result { + self.socket + .get_local_address() + .map_err(to_io_err) + .map(sockaddr_from_wasi) + } + + /// Sends a datagram to the given address. + pub async fn send_to(&self, buf: &[u8], addr: SocketAddr) -> io::Result { + self.socket + .send(buf.to_vec(), Some(sockaddr_to_wasi(addr))) + .await + .map_err(to_io_err)?; + Ok(buf.len()) + } + + /// Receives a single datagram. On success, returns the number of bytes + /// received and the address the datagram was sent from. + /// + /// If `buf` is shorter than the datagram, the excess bytes are discarded. + pub async fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + let (datagram, remote_address) = self.socket.receive().await.map_err(to_io_err)?; + let len = datagram.len().min(buf.len()); + buf[..len].copy_from_slice(&datagram[..len]); + Ok((len, sockaddr_from_wasi(remote_address))) + } + + /// Associates this socket with a remote address, giving a [`UdpStream`] + /// which sends to, and receives from, only that address. + /// + /// This only changes the local socket configuration, and does not generate + /// any network traffic. + pub fn connect(self, addr: SocketAddr) -> io::Result { + self.socket + .connect(sockaddr_to_wasi(addr)) + .map_err(to_io_err)?; + Ok(UdpStream::new(self.socket)) + } + + /// Returns the unicast hop limit ("time to live") of this socket. + pub fn unicast_hop_limit(&self) -> io::Result { + self.socket.get_unicast_hop_limit().map_err(to_io_err) + } + + /// Sets the unicast hop limit ("time to live") of this socket. + pub fn set_unicast_hop_limit(&self, value: u8) -> io::Result<()> { + self.socket.set_unicast_hop_limit(value).map_err(to_io_err) + } + + /// Returns the size of the receive buffer of this socket. + pub fn receive_buffer_size(&self) -> io::Result { + self.socket.get_receive_buffer_size().map_err(to_io_err) + } + + /// Sets the size of the receive buffer of this socket. This is a hint: the + /// size reported by [`UdpSocket::receive_buffer_size`] may differ. + pub fn set_receive_buffer_size(&self, value: u64) -> io::Result<()> { + self.socket + .set_receive_buffer_size(value) + .map_err(to_io_err) + } + + /// Returns the size of the send buffer of this socket. + pub fn send_buffer_size(&self) -> io::Result { + self.socket.get_send_buffer_size().map_err(to_io_err) + } + + /// Sets the size of the send buffer of this socket. This is a hint: the + /// size reported by [`UdpSocket::send_buffer_size`] may differ. + pub fn set_send_buffer_size(&self, value: u64) -> io::Result<()> { + self.socket.set_send_buffer_size(value).map_err(to_io_err) + } +} + +/// A UDP socket associated with a remote address. +/// +/// A `UdpStream` sends to, and receives from, only the address it was connected +/// to, using [`UdpStream::send`] and [`UdpStream::recv`]. Datagrams sent from +/// any other address are not received. +#[derive(Debug)] +pub struct UdpStream { + socket: WasiUdpSocket, +} + +impl UdpStream { + fn new(socket: WasiUdpSocket) -> Self { + Self { socket } + } + + /// Associates a UDP socket with a remote host. + pub async fn connect(addr: impl ToSocketAddrs) -> io::Result { + let addrs = addr.to_socket_addrs()?; + let mut last_err = None; + for addr in addrs { + match UdpStream::connect_addr(addr).await { + Ok(stream) => return Ok(stream), + Err(e) => last_err = Some(e), + } + } + + Err(last_err.unwrap_or_else(|| { + io::Error::new(ErrorKind::InvalidInput, "could not resolve to any address") + })) + } + + /// Establishes an association with the specified `addr`. + pub async fn connect_addr(addr: SocketAddr) -> io::Result { + // Unlike in POSIX, WASI requires a UDP socket be explicitly bound + // before it can be associated with a remote address. Bind to the + // unspecified address of the same family, and let the host choose a + // port. + let local_addr = match addr { + SocketAddr::V4(_) => SocketAddr::from((std::net::Ipv4Addr::UNSPECIFIED, 0)), + SocketAddr::V6(_) => SocketAddr::from((std::net::Ipv6Addr::UNSPECIFIED, 0)), + }; + let socket = bind_socket(local_addr).await?; + socket.connect(sockaddr_to_wasi(addr)).map_err(to_io_err)?; + Ok(Self::new(socket)) + } + + /// Returns the local socket address of this socket. + pub fn local_addr(&self) -> io::Result { + self.socket + .get_local_address() + .map_err(to_io_err) + .map(sockaddr_from_wasi) + } + + /// Returns the socket address of the remote peer of this UDP association. + pub fn peer_addr(&self) -> io::Result { + self.socket + .get_remote_address() + .map_err(to_io_err) + .map(sockaddr_from_wasi) + } + + /// Sends a datagram to the remote peer. + pub async fn send(&self, buf: &[u8]) -> io::Result { + self.socket + .send(buf.to_vec(), None) + .await + .map_err(to_io_err)?; + Ok(buf.len()) + } + + /// Receives a single datagram from the remote peer. On success, returns the + /// number of bytes received. + /// + /// If `buf` is shorter than the datagram, the excess bytes are discarded. + pub async fn recv(&self, buf: &mut [u8]) -> io::Result { + let (datagram, _remote_address) = self.socket.receive().await.map_err(to_io_err)?; + let len = datagram.len().min(buf.len()); + buf[..len].copy_from_slice(&datagram[..len]); + Ok(len) + } + + /// Returns the unicast hop limit ("time to live") of this socket. + pub fn unicast_hop_limit(&self) -> io::Result { + self.socket.get_unicast_hop_limit().map_err(to_io_err) + } + + /// Sets the unicast hop limit ("time to live") of this socket. + pub fn set_unicast_hop_limit(&self, value: u8) -> io::Result<()> { + self.socket.set_unicast_hop_limit(value).map_err(to_io_err) + } + + /// Returns the size of the receive buffer of this socket. + pub fn receive_buffer_size(&self) -> io::Result { + self.socket.get_receive_buffer_size().map_err(to_io_err) + } + + /// Sets the size of the receive buffer of this socket. This is a hint: the + /// size reported by [`UdpStream::receive_buffer_size`] may differ. + pub fn set_receive_buffer_size(&self, value: u64) -> io::Result<()> { + self.socket + .set_receive_buffer_size(value) + .map_err(to_io_err) + } + + /// Returns the size of the send buffer of this socket. + pub fn send_buffer_size(&self) -> io::Result { + self.socket.get_send_buffer_size().map_err(to_io_err) + } + + /// Sets the size of the send buffer of this socket. This is a hint: the + /// size reported by [`UdpStream::send_buffer_size`] may differ. + pub fn set_send_buffer_size(&self, value: u64) -> io::Result<()> { + self.socket.set_send_buffer_size(value).map_err(to_io_err) + } +} + +async fn bind_socket(addr: SocketAddr) -> io::Result { + let family = match addr { + SocketAddr::V4(_) => IpAddressFamily::Ipv4, + SocketAddr::V6(_) => IpAddressFamily::Ipv6, + }; + let socket = create_udp_socket(family).map_err(to_io_err)?; + let local_address = sockaddr_to_wasi(addr); + socket.bind(local_address).map_err(to_io_err)?; + + Ok(socket) +} From ad68f7220e371af42757ea2abee41b6864d26aaf Mon Sep 17 00:00:00 2001 From: Adam Bratschi-Kaye Date: Fri, 2 Oct 2026 13:51:04 +0000 Subject: [PATCH 4/4] cleanup old p2 changes --- src/net/tcp_listener/sys/p2.rs | 4 ---- src/net/udp/sys/p2.rs | 7 +++---- 2 files changed, 3 insertions(+), 8 deletions(-) diff --git a/src/net/tcp_listener/sys/p2.rs b/src/net/tcp_listener/sys/p2.rs index 1c8670a..06a8804 100644 --- a/src/net/tcp_listener/sys/p2.rs +++ b/src/net/tcp_listener/sys/p2.rs @@ -11,7 +11,6 @@ use crate::runtime::AsyncPollable; #[derive(Debug)] pub struct TcpListener { // Field order matters: must drop this child before parent below - #[cfg(target_env = "p2")] pollable: AsyncPollable, socket: TcpSocket, } @@ -47,10 +46,7 @@ impl TcpListener { /// Returns the local socket address of this listener. pub fn local_addr(&self) -> io::Result { - #[cfg(target_env = "p2")] let addr = self.socket.local_address(); - #[cfg(target_env = "p3")] - let addr = self.socket.get_local_address(); addr.map_err(to_io_err).map(sockaddr_from_wasi) } diff --git a/src/net/udp/sys/p2.rs b/src/net/udp/sys/p2.rs index 810a3ad..6de0ef0 100644 --- a/src/net/udp/sys/p2.rs +++ b/src/net/udp/sys/p2.rs @@ -55,10 +55,9 @@ impl UdpSocket { /// Sends a datagram to the given address. pub async fn send_to(&self, buf: &[u8], addr: SocketAddr) -> io::Result { - return self - .outgoing + self.outgoing .send_to(buf, Some(sockaddr_to_wasi(addr))) - .await; + .await } /// Receives a single datagram. On success, returns the number of bytes @@ -66,7 +65,7 @@ impl UdpSocket { /// /// If `buf` is shorter than the datagram, the excess bytes are discarded. pub async fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { - return self.incoming.recv_from(buf).await; + self.incoming.recv_from(buf).await } /// Associates this socket with a remote address, giving a [`UdpStream`]