diff --git a/crates/test-programs/src/bin/p2_tcp_streams.rs b/crates/test-programs/src/bin/p2_tcp_streams.rs index 814c2f548bb3..a577dbcee2d9 100644 --- a/crates/test-programs/src/bin/p2_tcp_streams.rs +++ b/crates/test-programs/src/bin/p2_tcp_streams.rs @@ -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}; @@ -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| { @@ -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); @@ -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); } diff --git a/crates/wasi/src/p2/tcp.rs b/crates/wasi/src/p2/tcp.rs index 3589b473ba07..6a70a816cdda 100644 --- a/crates/wasi/src/p2/tcp.rs +++ b/crates/wasi/src/p2/tcp.rs @@ -149,7 +149,12 @@ impl From for StreamError { enum WriteState { Ready(TcpSendStream, usize), Writing(MaybeSpawned>), - Closing(MaybeSpawned>), + // 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>, + Option, + ), Closed(WriteError), } @@ -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()), } } @@ -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, @@ -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() { @@ -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!(), + } } } @@ -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>); + + 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; + } +}