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
1 change: 1 addition & 0 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

3 changes: 3 additions & 0 deletions crates/wasi/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,9 @@ env_logger = { workspace = true }
rustix = { workspace = true, features = ["event", "fs", "net", "time"] }
rustix-linux-procfs = "0.1.1"

[target.'cfg(target_vendor = "apple")'.dependencies]
libc = { workspace = true }

[target.'cfg(windows)'.dependencies]
rustix = { workspace = true, features = ["event", "net"] }

Expand Down
132 changes: 132 additions & 0 deletions crates/wasi/src/p2/tcp.rs
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,15 @@ pub struct TcpSocket {
writer: Option<TcpWriter>,
}

impl Drop for TcpSocket {
fn drop(&mut self) {
// Reset before the retained streams shut down. Shutting down the
// writer can make the peer observe a clean EOF, and on macOS shutting
// down the reader discards unread data that would otherwise cause a reset.
self.inner.abort_if_unread();
}
}

#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum AsyncOperation {
Bind,
Expand Down Expand Up @@ -318,3 +327,126 @@ impl Pollable for TcpWriter {
poll_fn(|cx| self.0.lock().unwrap().poll_ready(cx).map(|_| ())).await;
}
}

#[cfg(all(test, any(target_os = "linux", target_os = "macos")))]
mod tests {
use super::*;
use std::time::Duration;

#[derive(Clone, Copy)]
enum CloseMode {
Unread,
Read,
Empty,
ShutdownSend,
RetainWriter,
}

async fn close_with_unread_data(family: crate::sockets::SocketAddressFamily, mode: CloseMode) {
use crate::WasiCtxBuilder;
use tokio::io::{AsyncReadExt, AsyncWriteExt};

tokio::time::timeout(Duration::from_secs(5), async {
let address = match family {
crate::sockets::SocketAddressFamily::Ipv4 => "127.0.0.1:0",
crate::sockets::SocketAddressFamily::Ipv6 => "[::1]:0",
};
let listener = tokio::net::TcpListener::bind(address).await.unwrap();
let mut ctx = WasiCtxBuilder::new();
ctx.inherit_network().allow_tcp(true);
let ctx = ctx.build();
let mut inner = P3Socket::new(&ctx.sockets, family).unwrap();
inner.start_connect(listener.local_addr().unwrap()).unwrap();
poll_fn(|cx| inner.poll_finish_connect(cx)).await.unwrap();
let (mut peer, _) = listener.accept().await.unwrap();
let mut socket = TcpSocket::new(inner);
let (mut input, output) = socket.take_streams().unwrap();

if !matches!(mode, CloseMode::Empty) {
peer.write_all(b"unread").await.unwrap();
input.ready().await;
}
if matches!(mode, CloseMode::Read) {
let mut received = Vec::new();
while received.len() < 6 {
received.extend_from_slice(&input.read(6).unwrap());
if received.len() < 6 {
input.ready().await;
}
}
assert_eq!(received, b"unread");
}

if matches!(mode, CloseMode::ShutdownSend) {
socket.shutdown(Shutdown::Write).unwrap();
// An explicit send shutdown must still give the peer a clean EOF
// even with unread incoming data. Observe it before dropping the socket.
assert_eq!(peer.read(&mut [0]).await.unwrap(), 0);
}

let retained_writer = if matches!(mode, CloseMode::RetainWriter) {
// Keep the native descriptor open past socket drop.
socket.writer.clone()
} else {
None
};

// Match wasi-libc close: the socket retains both streams until
// after their guest resources have been dropped.
drop(input);
drop(output);
drop(socket);

if matches!(mode, CloseMode::ShutdownSend) {
return;
}
let result = peer.read(&mut [0]).await;
if matches!(mode, CloseMode::Unread | CloseMode::RetainWriter) {
assert_eq!(
result.unwrap_err().kind(),
std::io::ErrorKind::ConnectionReset
);
} else {
assert_eq!(result.unwrap(), 0);
}
drop(retained_writer);
})
.await
.unwrap();
}

#[tokio::test]
async fn close_resets_unread_data() {
use crate::sockets::SocketAddressFamily;
close_with_unread_data(SocketAddressFamily::Ipv4, CloseMode::Unread).await;
close_with_unread_data(SocketAddressFamily::Ipv6, CloseMode::Unread).await;
}

#[tokio::test]
async fn close_is_orderly_after_reading_data() {
use crate::sockets::SocketAddressFamily;
close_with_unread_data(SocketAddressFamily::Ipv4, CloseMode::Read).await;
close_with_unread_data(SocketAddressFamily::Ipv6, CloseMode::Read).await;
}

#[tokio::test]
async fn close_is_orderly_without_received_data() {
use crate::sockets::SocketAddressFamily;
close_with_unread_data(SocketAddressFamily::Ipv4, CloseMode::Empty).await;
close_with_unread_data(SocketAddressFamily::Ipv6, CloseMode::Empty).await;
}

#[tokio::test]
async fn send_shutdown_is_orderly_with_unread_data() {
use crate::sockets::SocketAddressFamily;
close_with_unread_data(SocketAddressFamily::Ipv4, CloseMode::ShutdownSend).await;
close_with_unread_data(SocketAddressFamily::Ipv6, CloseMode::ShutdownSend).await;
}

#[tokio::test]
async fn close_resets_unread_data_with_retained_writer() {
use crate::sockets::SocketAddressFamily;
close_with_unread_data(SocketAddressFamily::Ipv4, CloseMode::RetainWriter).await;
close_with_unread_data(SocketAddressFamily::Ipv6, CloseMode::RetainWriter).await;
}
}
43 changes: 43 additions & 0 deletions crates/wasi/src/sockets/tcp.rs
Original file line number Diff line number Diff line change
Expand Up @@ -398,6 +398,20 @@ impl TcpSocket {
}
}

/// Preserve native close behavior for a P2 socket with unread data.
pub(crate) fn abort_if_unread(&self) {
#[cfg(unix)]
if let TcpState::Connected { stream, .. } = &self.tcp_state {
if matches!(rustix::io::ioctl_fionread(&**stream), Ok(unread) if unread > 0) {
// Zero linger alone takes effect when the native descriptor closes.
// Disconnect now, before stream shutdowns can signal a clean EOF
// or discard unread data.
_ = sockopt::set_socket_linger(&**stream, Some(Duration::ZERO));
abort_connection(stream);
}
}
}

pub(crate) fn is_listening(&self) -> bool {
matches!(self.tcp_state, TcpState::Listening(_))
}
Expand Down Expand Up @@ -613,6 +627,35 @@ impl TcpListenStream {
}
}

#[cfg(any(target_os = "linux", target_os = "android"))]
fn abort_connection(stream: &tokio::net::TcpStream) {
_ = rustix::net::connect_unspec(stream);
}

#[cfg(target_vendor = "apple")]
fn abort_connection(stream: &tokio::net::TcpStream) {
use std::os::fd::AsRawFd;

// Unlike Linux, connecting to AF_UNSPEC does not disconnect TCP on macOS.
// SAFETY: The descriptor remains open for this call, and both connection
// identifiers are passed by value; no pointers are involved.
unsafe {
libc::disconnectx(
stream.as_raw_fd(),
libc::SAE_ASSOCID_ANY,
libc::SAE_CONNID_ANY,
);
}
}

#[cfg(all(
unix,
not(any(target_os = "linux", target_os = "android", target_vendor = "apple"))
))]
fn abort_connection(_stream: &tokio::net::TcpStream) {
// On other Unix platforms, zero linger takes effect at the last close.
}

pub(crate) struct TcpSendStream {
inner: Arc<tokio::net::TcpStream>,
}
Expand Down
Loading