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
12 changes: 12 additions & 0 deletions crates/wasi-http/src/ctx.rs
Original file line number Diff line number Diff line change
Expand Up @@ -125,13 +125,17 @@ const DEFAULT_FIELD_SIZE_LIMIT: usize = 128 * 1024;
#[derive(Debug, Clone)]
pub struct WasiHttpCtx {
pub(crate) field_size_limit: usize,
#[cfg(feature = "p3")]
pub(crate) spawned_task_shutdown_grace_period: std::time::Duration,
}

impl WasiHttpCtx {
/// Create a new context.
pub fn new() -> Self {
Self {
field_size_limit: DEFAULT_FIELD_SIZE_LIMIT,
#[cfg(feature = "p3")]
spawned_task_shutdown_grace_period: std::time::Duration::from_secs(10),
}
}

Expand All @@ -145,6 +149,14 @@ impl WasiHttpCtx {
pub fn set_field_size_limit(&mut self, limit: usize) {
self.field_size_limit = limit;
}

/// Set how long spawned HTTP I/O tasks may continue after their Store is
/// dropped. The default is ten seconds. A zero duration cancels pending I/O
/// immediately when the Store is dropped.
#[cfg(feature = "p3")]
pub fn set_spawned_task_shutdown_grace_period(&mut self, timeout: std::time::Duration) {
self.spawned_task_shutdown_grace_period = timeout;
}
}

impl Default for WasiHttpCtx {
Expand Down
148 changes: 122 additions & 26 deletions crates/wasi-http/src/p3/host/handler.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,30 +6,96 @@ use crate::p3::{HttpError, HttpResult};
use crate::{Error, WasiHttp, WasiHttpCtxView};
use core::task::{Context, Poll, Waker};
use http_body_util::BodyExt as _;
use std::pin::Pin;
use std::sync::Arc;
use tokio::sync::oneshot;
use tokio::task::{self, JoinHandle};
use tokio::task::AbortHandle;
use tracing::debug;
use wasmtime::AsContextMut as _;
use wasmtime::component::{Accessor, HasData, Resource};
use wasmtime::error::Context as _;
use wasmtime_wasi::runtime::AbortOnDropJoinHandle;

/// A wrapper around [`JoinHandle`], which will [`JoinHandle::abort`] the task
/// A wrapper around [`AbortHandle`], which will [`AbortHandle::abort`] the task
/// when dropped
struct AbortOnDropJoinHandle(JoinHandle<()>);
struct AbortOnDropHandle(AbortHandle);

impl Drop for AbortOnDropJoinHandle {
impl Drop for AbortOnDropHandle {
fn drop(&mut self) {
self.0.abort();
}
}

/// Own an I/O task and allow it to finish during the Store's shutdown grace period.
struct DelayedAbortOnDropHandle {
inner: Option<AbortOnDropJoinHandle<()>>,
tx: Option<oneshot::Sender<AbortOnDropJoinHandle<()>>>,
}

impl DelayedAbortOnDropHandle {
fn new(
handle: wasmtime_wasi::runtime::AbortOnDropJoinHandle<()>,
timeout: std::time::Duration,
) -> Self {
let tx = if timeout.is_zero() {
None
} else {
let (tx, rx) = oneshot::channel::<AbortOnDropJoinHandle<()>>();
wasmtime_wasi::runtime::with_ambient_tokio_runtime(|| {
tokio::spawn(async move {
let Ok(handle) = rx.await else { return };
if !handle.is_finished() {
// We don't care if the task completes in the deadline
// or not.
let _ = tokio::time::timeout(timeout, handle).await;
}
});
});
Some(tx)
};
Self {
inner: Some(handle),
tx,
}
}
}

impl Future for DelayedAbortOnDropHandle {
type Output = Result<(), tokio::task::JoinError>;

fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let Some(task) = self.inner.as_mut() else {
return Poll::Ready(Ok(()));
};
let result = Pin::new(&mut **task).poll(cx);
if result.is_ready() {
// The task has finished, so dropping this wrapper needs no timer.
drop(self.inner.take());
}
result
}
}

impl Drop for DelayedAbortOnDropHandle {
fn drop(&mut self) {
let Some(inner) = self.inner.take() else {
return;
};
// Try sending the handle down the channel to be timed out.
// We took ownership of inner so it will be dropped (and cancelled) if
// the channel doesn't exist or sending fails.
if let Some(tx) = self.tx.take() {
let _ = tx.send(inner);
};
}
}

const DROPPED_FUTURE_ERROR: &str =
"Future indicating transmission result dropped without being resolved.";

async fn io_task_result(
rx: oneshot::Receiver<(
Option<Arc<AbortOnDropJoinHandle>>,
Option<Arc<AbortOnDropHandle>>,
oneshot::Receiver<Result<(), Error>>,
)>,
) -> Result<(), Error> {
Expand All @@ -44,7 +110,7 @@ async fn io_task_result(
fn send_dummy_io(
result: Result<(), Error>,
io_result_tx: oneshot::Sender<(
Option<Arc<AbortOnDropJoinHandle>>,
Option<Arc<AbortOnDropHandle>>,
oneshot::Receiver<Result<(), Error>>,
)>,
) {
Expand All @@ -58,7 +124,7 @@ fn send_dummy_io_err<T, D>(
mut getter: impl FnMut(&mut T) -> WasiHttpCtxView<'_>,
e: Error,
io_result_tx: oneshot::Sender<(
Option<Arc<AbortOnDropJoinHandle>>,
Option<Arc<AbortOnDropHandle>>,
oneshot::Receiver<Result<(), Error>>,
)>,
) -> HttpError
Expand Down Expand Up @@ -90,13 +156,13 @@ where
D: HasData,
T: 'static,
{
// A handle to the I/O task, if spawned, will be sent on this channel
// An abort handle to the I/O task, if spawned, will be sent on this channel
// and kept as part of request body state
let (io_task_tx, io_task_rx) = oneshot::channel();
let (request_body_io_tx, request_body_io_rx) = oneshot::channel();

// A handle to the I/O task, if spawned, will be sent on this channel
// An abort handle to the I/O task, if spawned, will be sent on this channel
// along with the result receiver
let (io_result_tx, io_result_rx) = oneshot::channel();
let (transmission_fut_io_tx, transmission_fut_io_rx) = oneshot::channel();

// Response processing result will be sent on this channel
let (res_result_tx, res_result_rx) = oneshot::channel();
Expand All @@ -108,11 +174,12 @@ where
.context("failed to delete request from table")
.map_err(HttpError::trap)?;
let (req, options) =
req.into_http_with_getter(&mut store, io_task_result(io_result_rx), getter)?;
req.into_http_with_getter(&mut store, io_task_result(transmission_fut_io_rx), getter)?;
HttpResult::Ok(getter(store.data_mut()).hooks.send_request(
// Attach a reference to the io task to the body so that it
// isn't cancelled if the body is dropped.
req.map(|body| body.with_state(io_task_rx).boxed_unsync()),
// Attach a reference to the io task to the body so that the task
// can be canceled if the body is dropped and all other references
// are dropped.
req.map(|body| body.with_state(request_body_io_rx).boxed_unsync()),
options.as_deref().copied(),
Box::new(async {
// Forward the response processing result to `WasiHttpCtx` implementation
Expand All @@ -127,19 +194,19 @@ where
Ok(fut) => fut,
Err(e) => match e.downcast() {
Ok(err_code) => {
send_dummy_io(Err(err_code.clone().into()), io_result_tx);
send_dummy_io(Err(err_code.clone().into()), transmission_fut_io_tx);
return Err(err_code.into());
}
Err(e) => {
let e = Error::InternalError(Some(format!("{e}")));
return Err(send_dummy_io_err(store, getter, e, io_result_tx));
return Err(send_dummy_io_err(store, getter, e, transmission_fut_io_tx));
}
},
};
let (res, io) = match Box::into_pin(fut).await {
Ok(r) => r,
Err(e) => {
return Err(send_dummy_io_err(store, getter, e, io_result_tx));
return Err(send_dummy_io_err(store, getter, e, transmission_fut_io_tx));
}
};
let (
Expand All @@ -152,24 +219,53 @@ where
let mut io = Box::into_pin(io);
let body = match io.as_mut().poll(&mut Context::from_waker(Waker::noop())) {
Poll::Ready(Ok(())) => {
send_dummy_io(Ok(()), io_result_tx);
send_dummy_io(Ok(()), transmission_fut_io_tx);
body
}
Poll::Ready(Err(e)) => {
return Err(send_dummy_io_err(store, getter, e, io_result_tx));
return Err(send_dummy_io_err(store, getter, e, transmission_fut_io_tx));
}
Poll::Pending => {
// I/O driver still needs to be polled, spawn a task and send handles to it
let (tx, rx) = oneshot::channel();
let io = Arc::new(AbortOnDropJoinHandle(task::spawn(async move {
let shutdown_timeout = store.with(|mut store| {
getter(store.data_mut())
.ctx
.spawned_task_shutdown_grace_period
});
// `task` is a tokio task which will be owned by the `Store` so
// that the task is cancelled if the `Store` is dropped. The
// outgoing body, transmission future, and incoming response body
// will all hold reference references to an abort handle on the task
// so that it is aborted when all three are dropped. But they will
// not retain owneship of the task, so they cannot keep it alive
// after the `Store` has dropped.
let task = wasmtime_wasi::runtime::spawn(async move {
let res = io.await;
debug!(?res, "`send_request` I/O future finished");
_ = tx.send(res);
})));
_ = io_result_tx.send((Some(Arc::clone(&io)), rx));
_ = io_task_tx.send(Arc::clone(&io));
// Attach a reference to the io task to the body so that it
// isn't cancelled if the body is dropped.
});
// `task` will be aborted when there are no more references to `io`.
let io = Arc::new(AbortOnDropHandle(task.abort_handle()));
let task = DelayedAbortOnDropHandle::new(task, shutdown_timeout);
// Pass ownership of `task` to `store`.
store
.spawn(async move |_| {
match task.await {
Ok(()) => {}
Err(e) if e.is_cancelled() => {}
Err(e) => std::panic::resume_unwind(e.into_panic()),
}
Ok(())
})
.map_err(HttpError::trap)?;
// Send one copy of `io` to the transmission future.
_ = transmission_fut_io_tx.send((Some(Arc::clone(&io)), rx));
// Send one copy of `io` to the request body.
_ = request_body_io_tx.send(Arc::clone(&io));
// Attach a reference to the io task to the response body so that
// the `task` can be cancelled if the body is dropped and no other
// references remain.
body.with_state(io).boxed_unsync()
}
};
Expand Down
Loading
Loading