diff --git a/crates/guest-rust/src/lib.rs b/crates/guest-rust/src/lib.rs index 3efde07ad..26f9813de 100644 --- a/crates/guest-rust/src/lib.rs +++ b/crates/guest-rust/src/lib.rs @@ -908,8 +908,6 @@ pub mod resource; #[cfg(feature = "inter-task-wakeup")] pub use rt::async_support::UnitStreamOps; -#[cfg(feature = "async-spawn")] -pub use rt::async_support::spawn_local; #[cfg(feature = "async")] pub use rt::async_support::{ AbiBuffer, FutureOps, FutureRead, FutureReader, FutureWrite, FutureWriteCancel, @@ -918,3 +916,5 @@ pub use rt::async_support::{ StreamRead, StreamReader, StreamResult, StreamWrite, StreamWriter, backpressure_dec, backpressure_inc, block_on, yield_async, yield_blocking, }; +#[cfg(feature = "async-spawn")] +pub use rt::async_support::{Task, spawn_local}; diff --git a/crates/guest-rust/src/rt/async_support.rs b/crates/guest-rust/src/rt/async_support.rs index bfcc091c8..d09274e7c 100644 --- a/crates/guest-rust/src/rt/async_support.rs +++ b/crates/guest-rust/src/rt/async_support.rs @@ -97,7 +97,7 @@ type BoxFuture<'a> = Pin + 'a>>; #[cfg(feature = "async-spawn")] mod spawn; #[cfg(feature = "async-spawn")] -pub use spawn::spawn_local; +pub use spawn::{Task, spawn_local}; #[cfg(not(feature = "async-spawn"))] mod spawn_disabled; #[cfg(not(feature = "async-spawn"))] diff --git a/crates/guest-rust/src/rt/async_support/spawn.rs b/crates/guest-rust/src/rt/async_support/spawn.rs index 2dff5b378..fbdc7ccb0 100644 --- a/crates/guest-rust/src/rt/async_support/spawn.rs +++ b/crates/guest-rust/src/rt/async_support/spawn.rs @@ -7,7 +7,10 @@ use crate::rt::async_support::BoxFuture; use alloc::boxed::Box; use alloc::vec::Vec; use core::future::Future; +use core::pin::Pin; use core::task::{Context, Poll}; +use futures::channel::oneshot; +use futures::future::{AbortHandle, Abortable, Aborted}; use futures::stream::{FuturesUnordered, StreamExt}; /// Any newly-deferred work queued by calls to the `spawn` function while @@ -94,10 +97,10 @@ impl<'a> Tasks<'a> { /// computations executing within a [`block_on`] call, however, the spawned /// tasks will be executed within that scope. This notably means that for /// [`block_on`] spawned tasks will prevent the [`block_on`] function from -/// returning, even if a value is available to return. -/// -/// * There is no handle returned to the spawned task meaning that it cannot be -/// cancelled or monitored. +/// returning, even if a value is available to return. If `spawn_local` is +/// called within a component-model async task which is then terminated (e.g. +/// by the host) before the future resolves, awating the `Task` will return +/// `None`. /// /// * The task spawned here is executed *concurrently*, not in *parallel*. This /// means that while one future is being polled no other future can be polled @@ -108,8 +111,79 @@ impl<'a> Tasks<'a> { /// exported async function has produced a value this can be used to continue to /// execute some more code before the component model async task exits. /// +/// # Cancellation +/// +/// Dropping the resulting [`Task`] will cancel the spawned future. [`Task::detach`] will +/// allow the future to continue running in the background and [`Task::cancel`] will +/// explicitly wait for the cancelation to complete. +/// /// [`block_on`]: crate::block_on /// [#1305]: https://github.com/bytecodealliance/wit-bindgen/issues/1305 -pub fn spawn_local(future: impl Future + 'static) { - unsafe { SPAWNED.push(Box::pin(future)) } +pub fn spawn_local(future: impl Future + 'static) -> Task { + let (sender, receiver) = oneshot::channel(); + let (abort, registration) = AbortHandle::new_pair(); + unsafe { + SPAWNED.push(Box::pin(async move { + let _ = sender.send(Abortable::new(future, registration).await); + })); + } + Task { + receiver, + abort, + cancel_on_drop: true, + } +} + +/// A handle to a spawned task which can be awaited for its result. +/// +/// Dropping this handle cancels the task. To drop the handle without cancelling +/// the task, call [`detach`](Self::detach). Awaiting the handle returns `None` +/// if the task was cancelled or otherwise terminated without producing a +/// result. +#[must_use = "dropping the handle cancels the spawned task"] +pub struct Task { + receiver: oneshot::Receiver>, + abort: AbortHandle, + cancel_on_drop: bool, +} + +impl Task { + /// Cancels the spawned task and waits for cancellation to complete. + /// + /// This returns the task's output if it completed before it could be + /// cancelled, or `None` if it was cancelled or otherwise terminated. + pub async fn cancel(mut self) -> Option { + self.abort.abort(); + self.cancel_on_drop = false; + match (&mut self.receiver).await { + Ok(Ok(result)) => Some(result), + Ok(Err(_)) => None, + Err(_) => None, + } + } + + /// Detaches the spawned task, allowing it to continue in the background. + pub fn detach(mut self) { + self.cancel_on_drop = false; + } +} + +impl Future for Task { + type Output = Option; + + fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + match Pin::new(&mut self.receiver).poll(cx) { + Poll::Ready(Ok(Ok(result))) => Poll::Ready(Some(result)), + Poll::Ready(Ok(Err(_)) | Err(_)) => Poll::Ready(None), + Poll::Pending => Poll::Pending, + } + } +} + +impl Drop for Task { + fn drop(&mut self) { + if self.cancel_on_drop { + self.abort.abort(); + } + } } diff --git a/tests/runtime/moonbit/nested-future-stream/test.rs b/tests/runtime/moonbit/nested-future-stream/test.rs index 7dec15953..055e4554b 100644 --- a/tests/runtime/moonbit/nested-future-stream/test.rs +++ b/tests/runtime/moonbit/nested-future-stream/test.rs @@ -56,7 +56,8 @@ impl Guest for Component { StreamResult::Cancelled => unreachable!(), } } - }); + }) + .detach(); outer_reader } @@ -81,7 +82,8 @@ impl Guest for Component { StreamResult::Cancelled => unreachable!(), } } - }); + }) + .detach(); output_reader } @@ -89,7 +91,8 @@ impl Guest for Component { let (mut writer, reader) = wit_stream::new(); wit_bindgen::spawn_local(async move { assert!(writer.write_all(vec![1, 2]).await.is_empty()); - }); + }) + .detach(); reader } @@ -97,7 +100,8 @@ impl Guest for Component { let (mut writer, reader) = wit_stream::new(); wit_bindgen::spawn_local(async move { assert!(writer.write_one(42).await.is_none()); - }); + }) + .detach(); reader } diff --git a/tests/runtime/moonbit/stream-write-cancel/test.rs b/tests/runtime/moonbit/stream-write-cancel/test.rs index 53a6ea7cd..dfed24799 100644 --- a/tests/runtime/moonbit/stream-write-cancel/test.rs +++ b/tests/runtime/moonbit/stream-write-cancel/test.rs @@ -84,7 +84,8 @@ impl Guest for Component { SECOND_WRITE_STARTED.store(true, Ordering::SeqCst); assert!(writer.write_one(holder::Leaf::new()).await.is_some()); assert!(writer.write_one(holder::Leaf::new()).await.is_some()); - }); + }) + .detach(); holder::hold(reader).await; } diff --git a/tests/runtime/ping-pong/test.rs b/tests/runtime/ping-pong/test.rs index d6b3e97ea..ff4610333 100644 --- a/tests/runtime/ping-pong/test.rs +++ b/tests/runtime/ping-pong/test.rs @@ -12,7 +12,8 @@ impl crate::exports::my::test::i::Guest for Component { let (tx, rx) = wit_future::new(|| unreachable!()); wit_bindgen::spawn_local(async move { tx.write(msg).await.unwrap(); - }); + }) + .detach(); rx } diff --git a/tests/runtime/rust-spawn-and-await/runner.rs b/tests/runtime/rust-spawn-and-await/runner.rs new file mode 100644 index 000000000..4fee1bb17 --- /dev/null +++ b/tests/runtime/rust-spawn-and-await/runner.rs @@ -0,0 +1,62 @@ +//@ wasmtime-flags = '-Wcomponent-model-async' + +include!(env!("BINDINGS")); + +use crate::test::rust_spawn_and_await::i::{ + await_resolve, await_task, cancel_task, resolve, start, +}; +use futures::task::noop_waker_ref; +use std::future::Future; +use std::pin::Pin; +use std::task::Context; + +struct Component; + +export!(Component); + +impl Guest for Component { + async fn run() { + // Awaiting a `Task` works. + let _cm_task = start_task(); + resolve().await; + let result = await_task().await; + assert_eq!(result, Some(42)); + + // Cancelling a `Task` before it completes returns `None`. + let _cm_task = start_task(); + let result = cancel_task().await; + resolve().await; + assert_eq!(result, None); + + // Cancelling a `Task` after it completes returns the result anyway. + let _cm_task = start_task(); + resolve().await; + await_resolve().await; + let result = cancel_task().await; + assert_eq!(result, Some(42)); + + // Check that awaiting a `Task` returns None after the CM-async task has + // been terminated. + let cm_task = start_task(); + drop(cm_task); + assert_eq!(await_task().await, None); + resolve().await; + + // Check that cancelling a `Task` returns None after the CM-async task + // has been terminated. + let cm_task = start_task(); + drop(cm_task); + assert_eq!(cancel_task().await, None); + resolve().await; + } +} + +fn start_task() -> Pin>> { + let mut task = Box::pin(start()); + assert!( + task.as_mut() + .poll(&mut Context::from_waker(noop_waker_ref())) + .is_pending() + ); + task +} diff --git a/tests/runtime/rust-spawn-and-await/test.rs b/tests/runtime/rust-spawn-and-await/test.rs new file mode 100644 index 000000000..6c2e44270 --- /dev/null +++ b/tests/runtime/rust-spawn-and-await/test.rs @@ -0,0 +1,57 @@ +include!(env!("BINDINGS")); + +use futures::channel::oneshot; +use std::cell::RefCell; +use wit_bindgen::{Task, spawn_local}; + +struct Component; + +export!(Component); + +std::thread_local! { + static TASK: RefCell>> = const { RefCell::new(None) }; + // Send through this channel to resolve the `Task`. + static RESOLVE_CHANNEL: RefCell>> = const { RefCell::new(None) }; + // Side channel to check that the `Task` has resolved without explicitly awaiting it. + static ACK_CHANNEL: RefCell>> = const { RefCell::new(None) }; +} + +impl crate::exports::test::rust_spawn_and_await::i::Guest for Component { + async fn start() { + let (tx, rx) = oneshot::channel(); + let (ack_tx, ack_rx) = oneshot::channel(); + let task = spawn_local(async { + rx.await.unwrap(); + let _ = ack_tx.send(()); + 42 + }); + TASK.with(|slot| assert!(slot.replace(Some(task)).is_none())); + RESOLVE_CHANNEL.with(|slot| assert!(slot.replace(Some(tx)).is_none())); + ACK_CHANNEL.with(|slot| slot.replace(Some(ack_rx))); + std::future::pending::<()>().await; + } + + async fn await_task() -> Option { + let task = TASK.with(|slot| slot.borrow_mut().take().unwrap()); + task.await + } + + async fn cancel_task() -> Option { + let task = TASK.with(|slot| slot.borrow_mut().take().unwrap()); + task.cancel().await + } + + async fn resolve() { + let channel = RESOLVE_CHANNEL.with(|slot| slot.borrow_mut().take().unwrap()); + // Ignore error when trying to resolve the `Task` because some tests + // cancel it before it completes. + let _ = channel.send(()); + } + + async fn await_resolve() { + let channel = ACK_CHANNEL.with(|slot| slot.borrow_mut().take().unwrap()); + // Ignore error when trying to resolve the `Task` because some tests + // cancel it before it completes. + channel.await.unwrap(); + } +} diff --git a/tests/runtime/rust-spawn-and-await/test.wit b/tests/runtime/rust-spawn-and-await/test.wit new file mode 100644 index 000000000..dce5f9b00 --- /dev/null +++ b/tests/runtime/rust-spawn-and-await/test.wit @@ -0,0 +1,20 @@ +//@ async = true +package test:rust-spawn-and-await; + +interface i { + start: async func(); + await-task: async func() -> option; + cancel-task: async func() -> option; + resolve: async func(); + await-resolve: async func(); +} + +world test { + export i; +} + +world runner { + import i; + + export run: async func(); +} diff --git a/tests/runtime/yield-loop-receives-events/middle.rs b/tests/runtime/yield-loop-receives-events/middle.rs index 4979c580a..63921a31f 100644 --- a/tests/runtime/yield-loop-receives-events/middle.rs +++ b/tests/runtime/yield-loop-receives-events/middle.rs @@ -16,7 +16,8 @@ impl crate::exports::test::common::i_runner::Guest for Component { unsafe { HIT = true; } - }); + }) + .detach(); // This is an "infinite loop" but it's also effectively a yield which // should enable not only making progress on sibling rust-level tasks