#[cfg(feature = "p2")]
use crate::p2;
#[cfg(feature = "p2")]
use crate::p2::bindings::http::types as p2_types;
#[cfg(feature = "p3")]
use crate::p3;
use crate::{WasiBody, WasiHttpCtxView};
use futures::{
channel::oneshot,
future::{Either, FutureExt},
stream::{FuturesUnordered, Stream},
};
#[cfg(feature = "p3")]
use p3::bindings::http::types as p3_types;
use std::collections::VecDeque;
use std::collections::btree_map::{BTreeMap, Entry};
use std::error;
use std::fmt;
use std::future;
use std::mem;
use std::ops::DerefMut;
use std::pin::{Pin, pin};
use std::sync::{
Arc, Mutex,
atomic::{AtomicBool, AtomicUsize, Ordering::Relaxed},
};
use std::task::{Context, Poll};
use std::time::Instant;
use tokio::sync::Notify;
use wasmtime::component::{Accessor, GuestTaskId, Resource, TypedFuncCallConcurrent};
#[cfg(feature = "p2")]
use wasmtime::error::Context as _;
use wasmtime::{AsContextMut, Result, Store, StoreContextMut, format_err};
pub type Request = http::Request<WasiBody>;
pub type Response = http::Response<WasiBody>;
pub enum ProxyPre<T: 'static> {
#[cfg(feature = "p2")]
P2(p2::bindings::ProxyPre<T>),
#[cfg(feature = "p3")]
P3(p3::bindings::ServicePre<T>),
}
impl<T: 'static> ProxyPre<T> {
pub async fn instantiate_async(&self, store: impl AsContextMut<Data = T>) -> Result<Proxy>
where
T: Send,
{
Ok(match self {
#[cfg(feature = "p2")]
Self::P2(pre) => Proxy::P2(pre.instantiate_async(store).await?),
#[cfg(feature = "p3")]
Self::P3(pre) => Proxy::P3(pre.instantiate_async(store).await?),
})
}
}
pub enum Proxy {
#[cfg(feature = "p2")]
P2(p2::bindings::Proxy),
#[cfg(feature = "p3")]
P3(p3::bindings::Service),
}
struct Queue<T> {
queue: Mutex<VecDeque<T>>,
notify_push: Notify,
}
impl<T> Default for Queue<T> {
fn default() -> Self {
Self {
queue: Default::default(),
notify_push: Default::default(),
}
}
}
impl<T> Queue<T> {
fn is_empty(&self) -> bool {
self.queue.lock().unwrap().is_empty()
}
fn try_pop(&self) -> Option<T> {
self.queue.lock().unwrap().pop_front()
}
async fn pop(&self) -> T {
let mut notified = pin!(self.notify_push.notified());
loop {
notified.as_mut().enable();
if let Some(item) = self.try_pop() {
return item;
}
notified.as_mut().await;
notified.set(self.notify_push.notified());
}
}
}
#[derive(Clone, Copy, Eq, PartialEq, Debug)]
pub enum WorkerStatus {
Idle,
Requests,
PostReturn,
}
pub trait WorkerExpiration: 'static + Send + Sync {
fn poll(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
state: WorkerStatus,
start: Instant,
) -> Poll<()>;
}
pub trait WorkerState: 'static + Send + Sync {
type StoreData: Send;
type RequestData: Send + Sync;
fn should_accept_request(&self, concurrent_count: usize, total_count: usize) -> ShouldAccept;
fn on_request_start(
&self,
store: StoreContextMut<'_, Self::StoreData>,
data: Self::RequestData,
task: GuestTaskId,
) -> Pin<Box<dyn Future<Output = ()> + 'static + Send + Sync>>;
fn drop(&self, store: Store<Self::StoreData>, result: Result<(), wasmtime::Error>);
}
pub struct Instance<T: 'static, E: WorkerExpiration, S: WorkerState> {
pub store: Store<T>,
pub proxy: Proxy,
pub view: fn(&mut T) -> WasiHttpCtxView<'_>,
pub expiration: E,
pub state: S,
}
pub enum ShouldAccept {
Yes,
No,
Never,
}
pub trait HandlerState: 'static + Sync + Send + Sized {
type StoreData: Send;
type WorkerExpiration: WorkerExpiration;
type WorkerState: WorkerState<StoreData = Self::StoreData>;
fn instantiate(
&self,
) -> impl Future<
Output = Result<Instance<Self::StoreData, Self::WorkerExpiration, Self::WorkerState>>,
> + Send;
}
struct ProxyHandlerInner<S: HandlerState> {
state: S,
request_queue: Queue<WorkerRequest<S>>,
worker_count: AtomicUsize,
}
#[derive(Default)]
struct StartTimes(BTreeMap<Instant, usize>);
impl StartTimes {
fn add(&mut self, time: Instant) {
*self.0.entry(time).or_insert(0) += 1;
}
fn remove(&mut self, time: Instant) {
let Entry::Occupied(mut entry) = self.0.entry(time) else {
unreachable!()
};
match *entry.get() {
0 => unreachable!(),
1 => {
entry.remove();
}
_ => {
*entry.get_mut() -= 1;
}
}
}
fn most_recent(&self) -> Option<Instant> {
self.0.last_key_value().map(|(&k, _)| k)
}
}
type WorkerRequest<S> = (
<<S as HandlerState>::WorkerState as WorkerState>::RequestData,
Request,
oneshot::Sender<Result<Response, wasmtime::Error>>,
);
struct Worker<S>
where
S: HandlerState,
{
handler: ProxyHandler<S>,
available: bool,
}
impl<S> Worker<S>
where
S: HandlerState,
{
fn set_available(&mut self, available: bool) {
if available != self.available {
self.available = available;
if available {
self.handler.0.worker_count.fetch_add(1, Relaxed);
} else {
let count = self.handler.0.worker_count.fetch_sub(1, Relaxed);
assert!(count >= 1);
if count == 1 && !self.handler.0.request_queue.is_empty() {
self.handler.start_worker(None);
}
}
}
}
async fn run(self, request: Option<WorkerRequest<S>>) {
match self.handler.0.state.instantiate().await {
Ok(Instance {
store,
proxy,
view,
expiration,
state,
}) => {
self.run_(store, proxy, view, expiration, state, request)
.await
}
Err(error) => {
let error = Arc::new(error);
if let Some((request_data, request, tx)) = request {
_ = tx.send(Err(InstantiationError {
request_data,
request: Mutex::new(request),
error,
}
.into()));
} else {
for (request_data, request, tx) in mem::take(
self.handler
.0
.request_queue
.queue
.lock()
.unwrap()
.deref_mut(),
) {
_ = tx.send(Err(InstantiationError {
request_data,
request: Mutex::new(request),
error: error.clone(),
}
.into()));
}
}
}
}
}
async fn run_(
mut self,
store: Store<S::StoreData>,
proxy: Proxy,
view: fn(&mut S::StoreData) -> WasiHttpCtxView<'_>,
expiration: S::WorkerExpiration,
state: S::WorkerState,
request: Option<WorkerRequest<S>>,
) {
struct Dropper<S: HandlerState> {
state: S::WorkerState,
store: Option<Store<S::StoreData>>,
}
impl<S: HandlerState> Drop for Dropper<S> {
fn drop(&mut self) {
if let Some(store) = self.store.take() {
self.state
.drop(store, Err(wasmtime::format_err!("worker panicked")));
}
}
}
let mut dropper = Dropper::<S> {
state,
store: Some(store),
};
let proxy = &proxy;
let accept_concurrent = AtomicBool::new(true);
let status = Mutex::new((WorkerStatus::Idle, Instant::now()));
let mut expiration = pin!(expiration);
let function = async |accessor: &Accessor<_>| {
let mut reuse_count = 0;
let mut may_accept = true;
let mut futures = FuturesUnordered::new();
let mut start_times = StartTimes::default();
let accept_request = |(request_data, request, tx): WorkerRequest<S>,
futures: &mut FuturesUnordered<_>,
start_times: &mut StartTimes,
reuse_count: &mut usize| {
accept_concurrent.store(false, Relaxed);
*reuse_count += 1;
let prepared = accessor.with(|mut store| {
let prepared = Prepared::new(store.as_context_mut(), proxy, request, view, tx);
match prepared {
Ok(prepared) => {
let expiration = dropper.state.on_request_start(
store.as_context_mut(),
request_data,
prepared.task(),
);
Ok((prepared, expiration))
}
Err(e) => Err(e),
}
});
let start_time = Instant::now();
start_times.add(start_time);
*status.try_lock().unwrap() = (WorkerStatus::Requests, start_time);
futures.push(async move {
let (prepared, expiration) = prepared?;
let sent = prepared.run(accessor, expiration).await?;
wasmtime::error::Ok((sent, start_time))
});
};
if let Some(req) = request {
accept_request(req, &mut futures, &mut start_times, &mut reuse_count);
}
let mut futures = pin!(futures);
let handler = self.handler.clone();
let mut incoming_requests = pin!(futures::stream::unfold(
&handler.0.request_queue,
|queue| async move {
let pair = queue.pop().await;
Some((pair, queue))
}
));
let func = match proxy {
#[cfg(feature = "p3")]
Proxy::P3(guest) => *guest.wasi_http_handler().func_handle().func(),
#[cfg(feature = "p2")]
Proxy::P2(guest) => *guest.wasi_http_incoming_handler().func_handle().func(),
};
future::poll_fn(|cx| {
loop {
match futures.as_mut().poll_next(cx) {
Poll::Ready(Some(Ok((responded, start_time)))) => {
start_times.remove(start_time);
*status.try_lock().unwrap() =
if let Some(start_time) = start_times.most_recent() {
(WorkerStatus::Requests, start_time)
} else {
(WorkerStatus::PostReturn, Instant::now())
};
if responded {
} else {
may_accept = false;
}
}
Poll::Ready(Some(Err(error))) => {
break Poll::Ready(Err(error));
}
Poll::Ready(None) | Poll::Pending => {}
}
let is_ready = accessor.poll_ready_for_concurrent_call(func, cx).is_ready();
self.set_available(
may_accept
&& is_ready
&& match dropper
.state
.should_accept_request(futures.len(), reuse_count)
{
ShouldAccept::Yes => {
futures.is_empty() || accept_concurrent.load(Relaxed)
}
ShouldAccept::No => false,
ShouldAccept::Never => {
may_accept = false;
false
}
},
);
if self.available
&& let Poll::Ready(Some(req)) = incoming_requests.as_mut().poll_next(cx)
{
accept_request(req, &mut futures, &mut start_times, &mut reuse_count);
continue;
}
if !futures.is_empty() {
break Poll::Pending;
}
if accessor.poll_no_interesting_tasks(cx).is_pending() {
break Poll::Pending;
}
if !(may_accept && is_ready) {
break Poll::Ready(Ok(()));
}
{
let mut status = status.try_lock().unwrap();
if status.0 != WorkerStatus::Idle {
*status = (WorkerStatus::Idle, Instant::now());
}
}
break Poll::Pending;
}
})
.await
};
let result = {
let mut future = pin!(
dropper
.store
.as_mut()
.unwrap()
.run_concurrent(function)
.map(|v| v.flatten())
);
future::poll_fn(|cx| {
let poll = future.as_mut().poll(cx);
if poll.is_pending() {
let (status, start) = *status.try_lock().unwrap();
if let Poll::Ready(()) = expiration.as_mut().poll(cx, status, start) {
return Poll::Ready(match status {
WorkerStatus::Requests | WorkerStatus::PostReturn => {
Err(format_err!("guest timed out"))
}
WorkerStatus::Idle => Ok(()),
});
}
if !accept_concurrent.swap(true, Relaxed) {
return future.as_mut().poll(cx);
}
}
poll
})
.await
};
dropper.state.drop(dropper.store.take().unwrap(), result);
}
}
impl<S> Drop for Worker<S>
where
S: HandlerState,
{
fn drop(&mut self) {
self.set_available(false);
}
}
pub struct ProxyHandler<S: HandlerState>(Arc<ProxyHandlerInner<S>>);
impl<S: HandlerState> Clone for ProxyHandler<S> {
fn clone(&self) -> Self {
Self(self.0.clone())
}
}
pub struct InstantiationError<T> {
pub request_data: T,
pub request: Mutex<Request>,
pub error: Arc<wasmtime::Error>,
}
impl<T> fmt::Display for InstantiationError<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> Result<(), fmt::Error> {
write!(f, "instantiation error: {}", self.error)
}
}
impl<T> fmt::Debug for InstantiationError<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> Result<(), fmt::Error> {
write!(f, "instantiation error: {:?}", self.error)
}
}
impl<T> error::Error for InstantiationError<T> {}
pub struct ExpirationError;
impl fmt::Display for ExpirationError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> Result<(), fmt::Error> {
fmt::Debug::fmt(self, f)
}
}
impl fmt::Debug for ExpirationError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> Result<(), fmt::Error> {
write!(f, "guest timed out")
}
}
impl error::Error for ExpirationError {}
pub struct TrapOrPanicError;
impl fmt::Display for TrapOrPanicError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> Result<(), fmt::Error> {
fmt::Debug::fmt(self, f)
}
}
impl fmt::Debug for TrapOrPanicError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> Result<(), fmt::Error> {
write!(f, "worker trapped or panicked")
}
}
impl error::Error for TrapOrPanicError {}
impl<S> ProxyHandler<S>
where
S: HandlerState,
{
pub fn new(state: S) -> Self {
Self(Arc::new(ProxyHandlerInner {
state,
request_queue: Default::default(),
worker_count: AtomicUsize::from(0),
}))
}
pub async fn handle(
&self,
data: <S::WorkerState as WorkerState>::RequestData,
request: Request,
) -> Result<Response, wasmtime::Error> {
let (tx, rx) = oneshot::channel();
let req = (data, request, tx);
if self.0.worker_count.load(Relaxed) == 0 {
self.start_worker(Some(req));
} else {
let mut queue = self.0.request_queue.queue.lock().unwrap();
queue.push_back(req);
if self.0.worker_count.load(Relaxed) == 0 {
let req = queue.pop_back().unwrap();
drop(queue);
self.start_worker(Some(req));
} else {
drop(queue);
self.0.request_queue.notify_push.notify_one();
}
}
rx.await.map_err(|_| TrapOrPanicError)?
}
pub fn state(&self) -> &S {
&self.0.state
}
fn start_worker(&self, request: Option<WorkerRequest<S>>) {
tokio::spawn(
Worker {
handler: self.clone(),
available: false,
}
.run(request),
);
}
}
pub enum Prepared<'a, T: 'static> {
#[doc(hidden)]
#[cfg(feature = "p2")]
P2 {
guest: &'a p2::bindings::Proxy,
call: TypedFuncCallConcurrent<
T,
(
Resource<p2_types::IncomingRequest>,
Resource<p2_types::ResponseOutparam>,
),
(),
>,
tx: Arc<Mutex<Option<oneshot::Sender<Result<Response, wasmtime::Error>>>>>,
},
#[doc(hidden)]
#[cfg(feature = "p3")]
P3 {
guest: &'a p3::bindings::Service,
call: TypedFuncCallConcurrent<
T,
(Resource<p3_types::Request>,),
(Result<Resource<p3_types::Response>, p3_types::ErrorCode>,),
>,
tx: oneshot::Sender<Result<Response, wasmtime::Error>>,
request_io_result: Pin<Box<dyn Future<Output = Result<(), crate::Error>> + Send>>,
view: fn(&mut T) -> crate::WasiHttpCtxView,
},
}
impl<'a, T: Send> Prepared<'a, T> {
pub fn new(
mut store: StoreContextMut<'_, T>,
proxy: &'a Proxy,
request: Request,
view: fn(&mut T) -> WasiHttpCtxView<'_>,
tx: oneshot::Sender<Result<Response, wasmtime::Error>>,
) -> Result<Prepared<'a, T>> {
match proxy {
#[cfg(feature = "p3")]
Proxy::P3(guest) => {
let (request, body) = request.into_parts();
let request = http::Request::from_parts(request, body);
let hooks = view(store.data_mut()).hooks;
let (request, request_io_result) = p3::Request::from_http(hooks, request);
let request = view(store.data_mut()).table.push(request)?;
Ok(Prepared::P3 {
tx,
request_io_result: Box::pin(request_io_result),
guest,
view,
call: guest
.wasi_http_handler()
.func_handle()
.start_call_concurrent(store, (request,))?,
})
}
#[cfg(feature = "p2")]
Proxy::P2(guest) => {
let tx = Arc::new(Mutex::new(Some(tx)));
let request =
view(store.data_mut()).new_incoming_request(p2_types::Scheme::Http, request)?;
let out = view(store.data_mut()).new_response_outparam_from_callback({
let tx = tx.clone();
move |value| {
if let Some(tx) = tx.lock().unwrap().take() {
_ = tx.send(value.map_err(|e| e.into()));
}
}
})?;
Ok(Prepared::P2 {
guest,
tx,
call: guest
.wasi_http_incoming_handler()
.func_handle()
.start_call_concurrent(store, (request, out))?,
})
}
}
}
fn task(&self) -> GuestTaskId {
match self {
#[cfg(feature = "p3")]
Prepared::P3 { call, .. } => call.task(),
#[cfg(feature = "p2")]
Prepared::P2 { call, .. } => call.task(),
}
}
pub async fn run(
self,
accessor: &Accessor<T>,
expiration: impl Future<Output = ()>,
) -> Result<bool> {
let expiration = pin!(expiration);
match self {
#[cfg(feature = "p3")]
Prepared::P3 {
guest,
call,
tx,
request_io_result,
view,
} => {
let handle = pin!(async move {
let response = guest
.wasi_http_handler()
.func_handle()
.finish_call_concurrent(accessor, call)
.await?
.0?;
accessor.with(|mut store| {
let response = view(store.get()).table.delete(response)?;
response.into_http_with_getter(&mut store, request_io_result, view)
})
});
let (result, sent) = match futures::future::select(handle, expiration).await {
Either::Left((result, _)) => (result, true),
Either::Right(((), _)) => (Err(ExpirationError.into()), false),
};
_ = tx.send(result);
Ok(sent)
}
#[cfg(feature = "p2")]
Prepared::P2 { guest, call, tx } => {
let handle = pin!(
guest
.wasi_http_incoming_handler()
.func_handle()
.finish_call_concurrent(accessor, call)
);
const MESSAGE: &str = "guest never invoked `response-outparam::set` method";
struct Dropper(
Arc<Mutex<Option<oneshot::Sender<Result<Response, wasmtime::Error>>>>>,
);
impl Drop for Dropper {
fn drop(&mut self) {
if let Some(tx) = self.0.lock().unwrap().take() {
_ = tx.send(Err(format_err!("{MESSAGE}")));
}
}
}
let tx = Dropper(tx);
let (result, sent) = match futures::future::select(handle, expiration).await {
Either::Left((result, _)) => (result.context(MESSAGE), true),
Either::Right(((), _)) => (Err(ExpirationError.into()), false),
};
if let Some(tx) = tx.0.lock().unwrap().take() {
_ = tx.send(result.and_then(|()| Err(format_err!("{MESSAGE}"))));
}
Ok(sent)
}
}
}
}