From d426cae51d025188bd12d8ab1905b7e10f1e7261 Mon Sep 17 00:00:00 2001 From: Adam Bratschi-Kaye Date: Tue, 22 Sep 2026 01:16:49 +0000 Subject: [PATCH] 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 7d20a2f..ec04b5a 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(()) +}