Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 4 additions & 4 deletions examples/tcp_echo_server.rs
Original file line number Diff line number Diff line change
@@ -1,15 +1,15 @@
#![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;
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 <PORT>` to create a TCP client");

let mut incoming = listener.incoming();
while let Some(stream) = incoming.next().await {
Expand Down
4 changes: 2 additions & 2 deletions examples/tcp_stream_client.rs
Original file line number Diff line number Diff line change
@@ -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;
Expand Down
8 changes: 4 additions & 4 deletions examples/udp_echo_server.rs
Original file line number Diff line number Diff line change
@@ -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 <PORT>` to create a UDP client");

let mut buf = vec![0; 65535];
loop {
Expand Down
4 changes: 2 additions & 2 deletions examples/udp_stream_client.rs
Original file line number Diff line number Diff line change
@@ -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};
Expand Down
1 change: 0 additions & 1 deletion src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down
47 changes: 41 additions & 6 deletions src/net/mod.rs
Original file line number Diff line number Diff line change
@@ -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;
Expand All @@ -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(
Expand All @@ -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();
Expand All @@ -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, wasip3::sockets::types::ErrorCode> {
TcpSocket::create(family)
}

#[cfg(target_env = "p3")]
fn create_udp_socket(
family: IpAddressFamily,
) -> Result<wasip3::sockets::types::UdpSocket, wasip3::sockets::types::ErrorCode> {
wasip3::sockets::types::UdpSocket::create(family)
}
70 changes: 51 additions & 19 deletions src/net/tcp_listener.rs
Original file line number Diff line number Diff line change
@@ -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<TcpSocket>,
socket: TcpSocket,
}

Expand All @@ -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<std::net::SocketAddr> {
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.
Expand All @@ -69,6 +90,7 @@ pub struct Incoming<'a> {
impl<'a> AsyncIterator for Incoming<'a> {
type Item = io::Result<TcpStream>;

#[cfg(target_env = "p2")]
async fn next(&mut self) -> Option<Self::Item> {
self.listener.pollable.wait_for().await;
let (socket, input, output) = match self.listener.socket.accept().map_err(to_io_err) {
Expand All @@ -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::Item> {
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))
Comment on lines +106 to +109

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We could save _receive_result and _send_result in the TcpStream and add a graceful shutdown method which checks them. But that should probably be a separate follow up.

})
}
}
60 changes: 45 additions & 15 deletions src/net/tcp_stream.rs
Original file line number Diff line number Diff line change
@@ -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<u8>;
#[cfg(target_env = "p3")]
type OutputStream = wasip3::wit_bindgen::StreamWriter<u8>;

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.
Expand Down Expand Up @@ -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) => {
Expand All @@ -70,36 +79,53 @@ 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<String> {
#[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:?}"))
}

pub fn split(&mut self) -> (ReadHalf<'_>, WriteHalf<'_>) {
(
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
Expand Down Expand Up @@ -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
Expand All @@ -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,
}

Expand All @@ -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
Expand Down
Loading
Loading