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
50 changes: 50 additions & 0 deletions crates/test-programs/src/bin/p2_tcp_streams.rs
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
use test_programs::sockets::supports_ipv6;
use test_programs::wasi::clocks::monotonic_clock;
use test_programs::wasi::io::streams::{InputStream, OutputStream, StreamError};
use test_programs::wasi::sockets::network::{IpAddress, IpAddressFamily, IpSocketAddress, Network};
use test_programs::wasi::sockets::tcp::{ShutdownType, TcpSocket};
Expand Down Expand Up @@ -110,6 +111,51 @@ fn test_tcp_shutdown_should_not_lose_data(net: &Network, family: IpAddressFamily
});
}

/// Drop must cancel an unflushed write so the guest can read its peer.
fn test_tcp_drop_should_not_wait_for_peer(net: &Network, family: IpAddressFamily, shutdown: bool) {
setup(net, family, |server, client| {
client.socket.set_send_buffer_size(1024).unwrap();
let writable = client.output.subscribe();
let chunk = vec![0x5a; 64 * 1024];
let mut sent = 0;
let mut last_write = 0;

// Fill the connection without reading the peer. Wait out transient
// background work so check-write returning zero means backpressure.
loop {
let permit = client.output.check_write().unwrap() as usize;
if permit == 0 {
let timeout = monotonic_clock::subscribe_duration(1_000_000_000);
if writable.block_until(&timeout).is_err() {
break;
}
continue;
}
let len = permit.min(chunk.len());
client.output.write(&chunk[..len]).unwrap();
sent += len;
last_write = len;
assert!(sent < 64 * 1024 * 1024, "connection never backed up");
}
assert!(sent > 0);
assert_eq!(client.output.check_write().unwrap(), 0);
drop(writable);

// Cover cancelling both a write and a shutdown waiting for that write.
if shutdown {
client.socket.shutdown(ShutdownType::Send).unwrap();
}
drop(client);

// Dropping an unflushed stream may lose the unfinished write, but all
// earlier writes must have reached the OS and remain intact.
let received = server.input.blocking_read_to_end().unwrap();
assert!(received.len() >= sent - last_write);
assert!(received.len() <= sent);
assert!(received.iter().all(|byte| *byte == 0x5a));
});
}

// Once a stream is writable it should in theory always be writable...
fn test_tcp_check_write_should_not_be_rate_limited(net: &Network, family: IpAddressFamily) {
setup(net, family, |_server, client| {
Expand Down Expand Up @@ -156,6 +202,8 @@ fn main() {
test_tcp_input_stream_should_be_closed_by_local_shutdown(&net, IpAddressFamily::Ipv4);
test_tcp_output_stream_should_be_closed_by_local_shutdown(&net, IpAddressFamily::Ipv4);
test_tcp_shutdown_should_not_lose_data(&net, IpAddressFamily::Ipv4);
test_tcp_drop_should_not_wait_for_peer(&net, IpAddressFamily::Ipv4, false);
test_tcp_drop_should_not_wait_for_peer(&net, IpAddressFamily::Ipv4, true);
test_tcp_check_write_should_not_be_rate_limited(&net, IpAddressFamily::Ipv4);
test_tcp_nonblocking_write_loop(&net, IpAddressFamily::Ipv4);

Expand All @@ -164,6 +212,8 @@ fn main() {
test_tcp_input_stream_should_be_closed_by_local_shutdown(&net, IpAddressFamily::Ipv6);
test_tcp_output_stream_should_be_closed_by_local_shutdown(&net, IpAddressFamily::Ipv6);
test_tcp_shutdown_should_not_lose_data(&net, IpAddressFamily::Ipv6);
test_tcp_drop_should_not_wait_for_peer(&net, IpAddressFamily::Ipv6, false);
test_tcp_drop_should_not_wait_for_peer(&net, IpAddressFamily::Ipv6, true);
test_tcp_check_write_should_not_be_rate_limited(&net, IpAddressFamily::Ipv6);
test_tcp_nonblocking_write_loop(&net, IpAddressFamily::Ipv6);
}
Expand Down
124 changes: 114 additions & 10 deletions crates/wasi/src/p2/tcp.rs
Original file line number Diff line number Diff line change
Expand Up @@ -149,7 +149,12 @@ impl From<WriteError> for StreamError {
enum WriteState {
Ready(TcpSendStream, usize),
Writing(MaybeSpawned<Result<TcpSendStream, WriteError>>),
Closing(MaybeSpawned<Result<(), WriteError>>),
// Keep the inner write's abort handle so cancellation can stop the write
// and join it through the shutdown task instead of only aborting its owner.
Closing(
MaybeSpawned<Result<(), WriteError>>,
Option<tokio::task::AbortHandle>,
),
Closed(WriteError),
}

Expand Down Expand Up @@ -222,7 +227,7 @@ impl WriteState {
// always be delivered to the OS as soon as possible. There's nothing
// for `flush` to do here that will speed up that process.
match self {
WriteState::Ready(..) | WriteState::Writing(_) | WriteState::Closing(_) => Ok(()),
WriteState::Ready(..) | WriteState::Writing(_) | WriteState::Closing(..) => Ok(()),
WriteState::Closed(e) => Err(e.clone().into()),
}
}
Expand All @@ -234,10 +239,25 @@ impl WriteState {

// Schedule the shutdown after the current write has finished:
WriteState::Writing(write) => {
WriteState::Closing(MaybeSpawned::poll_or_spawn(async move {
_ = write.into_future().await?;
let abort = match &write {
MaybeSpawned::Pending(task) => Some(task.abort_handle()),
MaybeSpawned::Ready(_) => None,
};
let close = MaybeSpawned::poll_or_spawn(async move {
let result = match write {
MaybeSpawned::Ready(result) => result,
// Await the raw Tokio handle so a cancellation request
// can resolve normally instead of panicking.
MaybeSpawned::Pending(mut task) => match (&mut *task).await {
Ok(result) => result,
Err(e) if e.is_cancelled() => return Ok(()),
Err(e) => std::panic::resume_unwind(e.into_panic()),
},
};
_ = result?;
Ok(())
}))
});
WriteState::Closing(close, abort)
}

s => s,
Expand All @@ -259,9 +279,9 @@ impl WriteState {
Err(err) => WriteState::Closed(err),
};
}
WriteState::Closing(close) => {
WriteState::Closing(close, _) => {
ready!(close.poll_ready(cx));
let WriteState::Closing(close) = self.take() else {
let WriteState::Closing(close, _) = self.take() else {
unreachable!()
};
*self = match close.unwrap_ready() {
Expand Down Expand Up @@ -306,9 +326,39 @@ impl OutputStream for TcpWriter {
}

async fn cancel(&mut self) {
// Wait for background writes to finish in order to prevent silently
// dropping data that (from the guest's perspective) was already written.
self.ready().await
let state = {
let mut state = self.0.lock().unwrap();
// The socket also owns this writer. Keep an idle send stream alive
// until the socket is shut down or dropped.
if matches!(*state, WriteState::Ready(..) | WriteState::Closed(..)) {
return;
}
state.take()
};

// Abort pending work and wait for cleanup, rather than waiting for the
// peer to read. As allowed by wasi:io/streams, an unflushed write may
// lose its remaining bytes when the output stream is dropped.
match state {
WriteState::Writing(write) => {
let result = match write {
MaybeSpawned::Pending(task) => task.cancel().await,
MaybeSpawned::Ready(result) => Some(result),
};
if let Some(Ok(stream)) = result {
*self.0.lock().unwrap() = WriteState::Ready(stream, 0);
}
}
WriteState::Closing(close, abort) => {
if let Some(abort) = abort {
abort.abort();
}
// The shutdown task joins the aborted write and releases its
// resources before finishing. Joining it waits for both tasks.
let _ = close.into_future().await;
}
_ => unreachable!(),
}
}
}

Expand All @@ -318,3 +368,57 @@ impl Pollable for TcpWriter {
poll_fn(|cx| self.0.lock().unwrap().poll_ready(cx).map(|_| ())).await;
}
}

#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
use tokio::sync::oneshot;

struct NotifyOnDrop(Option<oneshot::Sender<()>>);

impl Drop for NotifyOnDrop {
fn drop(&mut self) {
let _ = self.0.take().unwrap().send(());
}
}

async fn cancel_pending_write(shutdown: bool) {
let (resume, paused) = oneshot::channel::<()>();
let (finished, mut cleanup) = oneshot::channel();
let guard = NotifyOnDrop(Some(finished));
let write = MaybeSpawned::poll_or_spawn(async move {
let _guard = guard;
paused.await.unwrap();
Err(WriteError::Closed)
});
let mut writer = TcpWriter(Arc::new(Mutex::new(WriteState::Writing(write))));
let socket_writer = writer.clone();
if shutdown {
writer.0.lock().unwrap().shutdown();
}

// Cancellation must finish even though the peer never permits the
// write to complete, and must wait for the worker's cleanup.
tokio::time::timeout(Duration::from_secs(5), writer.cancel())
.await
.unwrap();
cleanup.try_recv().unwrap();
assert!(resume.send(()).is_err());
assert!(matches!(
*socket_writer.0.lock().unwrap(),
WriteState::Closed(WriteError::Closed)
));
writer.cancel().await;
}

#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn cancel_aborts_pending_write_and_waits_for_cleanup() {
cancel_pending_write(false).await;
}

#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn cancel_aborts_pending_shutdown_and_waits_for_cleanup() {
cancel_pending_write(true).await;
}
}
Loading