wasmer 7.3.0

High-performance WebAssembly runtime
#![cfg(all(
    feature = "experimental-async",
    any(not(target_arch = "wasm32"), feature = "js")
))]

use std::collections::BTreeMap;
use std::future::Future;
use std::sync::atomic::AtomicU32;
use std::sync::{Arc, RwLock};

use anyhow::Result;
use futures::channel::oneshot;
use futures::future::{AbortHandle, Abortable};
#[cfg(not(target_arch = "wasm32"))]
use futures::task::LocalSpawnExt;
#[cfg(target_arch = "wasm32")]
use wasm_bindgen_test::wasm_bindgen_test;
use wasmer::{
    AsStoreAsync, AsyncFunctionEnvMut, Function, FunctionEnv, FunctionEnvMut, FunctionType,
    Instance, Memory, Module, RuntimeError, Store, Type, Value, imports,
};

const SWITCHING_WAT: &[u8] = include_bytes!(concat!(
    env!("CARGO_MANIFEST_DIR"),
    "/../../tests/examples/simple-greenthread.wat"
));
const SWITCHING_LOGS: &[&str] = &[
    "[gr1] main  -> test1",
    "[gr2] test1 -> test2",
    "[gr1] test1 <- test2",
    "[gr2] test1 -> test2",
    "[gr1] test1 <- test2",
    "[main] main <- test1",
];
const REGRESSION_WAT: &[u8] = include_bytes!(concat!(
    env!("CARGO_MANIFEST_DIR"),
    "/../../tests/examples/simple-greenthread2.wat"
));
const REGRESSION_LOGS: &[&str] = &[
    "[main] switching to side",
    "[side] switching to main",
    "[main] switching to side",
    "[side] switching to main",
    "[main] returned",
];

#[derive(Clone)]
struct TestSpawner {
    #[cfg(not(target_arch = "wasm32"))]
    inner: futures::executor::LocalSpawner,
}

impl TestSpawner {
    fn spawn(&self, future: impl Future<Output = ()> + 'static) {
        #[cfg(not(target_arch = "wasm32"))]
        self.inner.spawn_local(future).unwrap();

        #[cfg(target_arch = "wasm32")]
        wasm_bindgen_futures::spawn_local(future);
    }
}

struct SpawnedTask {
    abort: AbortHandle,
    done: oneshot::Receiver<()>,
}

struct GreenEnv {
    logs: Vec<String>,
    memory: Option<Memory>,
    greenthreads: Arc<RwLock<BTreeMap<u32, Greenthread>>>,
    current_greenthread_id: Arc<RwLock<u32>>,
    next_free_id: AtomicU32,
    entrypoint: Option<Function>,
    spawner: Option<TestSpawner>,
    spawned_tasks: Vec<SpawnedTask>,
}

// Required for carrying the spawner around. Safe because we don't do threads.
// It worked before with a thread-local! spawner, so this is equivalent.
// The thread-local version does not work with multiple tests
unsafe impl Send for GreenEnv {}
unsafe impl Sync for GreenEnv {}

impl GreenEnv {
    fn new() -> Self {
        Self {
            logs: Vec::new(),
            memory: None,
            greenthreads: Arc::new(RwLock::new(BTreeMap::new())),
            current_greenthread_id: Arc::new(RwLock::new(0)),
            next_free_id: AtomicU32::new(1),
            entrypoint: None,
            spawner: None,
            spawned_tasks: Vec::new(),
        }
    }
}

struct Greenthread {
    entrypoint: Option<u32>,
    resumer: Option<oneshot::Sender<()>>,
}

impl Clone for Greenthread {
    fn clone(&self) -> Self {
        if self.resumer.is_some() {
            panic!("Cannot clone a greenthread with a resumer");
        }
        Self {
            entrypoint: self.entrypoint,
            resumer: None,
        }
    }
}

async fn greenthread_new(
    env: AsyncFunctionEnvMut<GreenEnv>,
    entrypoint_data: u32,
) -> core::result::Result<u32, RuntimeError> {
    let async_store = env.as_store_async();
    let mut env_write = env.write().await;
    let data = env_write.data_mut();
    let new_greenthread_id = data
        .next_free_id
        .fetch_add(1, std::sync::atomic::Ordering::SeqCst);

    let function = data.entrypoint.clone().expect("entrypoint set");
    let (sender, receiver) = oneshot::channel::<()>();

    let new_greenthread = Greenthread {
        entrypoint: Some(entrypoint_data),
        resumer: Some(sender),
    };

    data.greenthreads
        .write()
        .unwrap()
        .insert(new_greenthread_id, new_greenthread);

    let spawner = data.spawner.as_ref().expect("spawner set").clone();
    let (abort, registration) = AbortHandle::new_pair();
    let (done_tx, done) = oneshot::channel();
    data.spawned_tasks.push(SpawnedTask { abort, done });
    spawner.spawn(async move {
        let result = Abortable::new(
            async move {
                receiver.await.unwrap();
                function
                    .call_async(&async_store, vec![Value::I32(entrypoint_data as i32)])
                    .await
            },
            registration,
        )
        .await;
        if let Ok(result) = result {
            panic!("Greenthread function returned {result:?}");
        }
        let _ = done_tx.send(());
    });

    Ok(new_greenthread_id)
}

async fn greenthread_switch(env: AsyncFunctionEnvMut<GreenEnv>, next_greenthread_id: u32) {
    let (receiver, current_id_arc, current_greenthread_id) = {
        let mut write_lock = env.write().await;
        let data = write_lock.data_mut();

        let current_greenthread_id = {
            let mut current = data.current_greenthread_id.write().unwrap();
            let old = *current;
            *current = next_greenthread_id;
            old
        };

        if current_greenthread_id == next_greenthread_id {
            panic!("Switching to self is not allowed");
        }

        let (sender, receiver) = oneshot::channel::<()>();

        {
            let mut greenthreads = data.greenthreads.write().unwrap();
            let this_one = greenthreads.get_mut(&current_greenthread_id).unwrap();
            if this_one.resumer.is_some() {
                panic!("Switching from a greenthread that is already switched out");
            }
            this_one.resumer = Some(sender);
        }

        {
            let mut greenthreads = data.greenthreads.write().unwrap();
            let next_one = greenthreads.get_mut(&next_greenthread_id).unwrap();
            let Some(resumer) = next_one.resumer.take() else {
                panic!("Switching to greenthread that has no resumer");
            };
            resumer.send(()).unwrap();
        }
        let current_id_arc = data.current_greenthread_id.clone();

        (receiver, current_id_arc, current_greenthread_id)
    };

    let _ = receiver.await;

    *current_id_arc.write().unwrap() = current_greenthread_id;
}

async fn run_greenthread_test(wat: &[u8], spawner: TestSpawner) -> Result<Vec<String>> {
    let mut store = Store::default();
    let wasm = wat::parse_bytes(wat)?;
    let module = Module::new(&store.engine(), wasm)?;

    let env = FunctionEnv::new(&mut store, GreenEnv::new());

    // log(ptr, len)
    let log_fn = Function::new_with_env(
        &mut store,
        &env,
        FunctionType::new(vec![Type::I32, Type::I32], vec![]),
        |mut env: FunctionEnvMut<GreenEnv>, params: &[Value]| {
            let ptr = params[0].unwrap_i32() as u32;
            let len = params[1].unwrap_i32() as u32;
            let (data, storemut) = env.data_and_store_mut();
            let memory = data.memory.as_ref().expect("memory set");
            let view = memory.view(&storemut);
            let mut bytes = Vec::with_capacity(len as usize);
            for i in ptr..ptr + len {
                bytes.push(view.read_u8(i as u64).expect("in bounds"));
            }
            let s = String::from_utf8_lossy(&bytes).to_string();
            data.logs.push(s.trim_matches('\0').to_string());
            Ok(vec![])
        },
    );

    let greenthread_new = Function::new_typed_with_env_async(&mut store, &env, greenthread_new);

    let greenthread_switch =
        Function::new_typed_with_env_async(&mut store, &env, greenthread_switch);

    let import_object = imports! {
        "test" => {
            "log" => log_fn,
            "greenthread_new" => greenthread_new,
            "greenthread_switch" => greenthread_switch,
        }
    };

    let instance = Instance::new(&mut store, &module, &import_object)?;

    let entrypoint = instance.exports.get_function("entrypoint")?.clone();
    env.as_mut(&mut store).entrypoint = Some(entrypoint);

    let memory = instance.exports.get_memory("memory")?.clone();
    env.as_mut(&mut store).memory = Some(memory);

    let main_fn = instance.exports.get_function("_main")?;

    let main_greenthread = Greenthread {
        entrypoint: None,
        resumer: None,
    };

    env.as_mut(&mut store)
        .greenthreads
        .write()
        .unwrap()
        .insert(0, main_greenthread);

    env.as_mut(&mut store).spawner = Some(spawner);

    let store_async = store.into_async();

    main_fn.call_async(&store_async, vec![]).await?;

    let spawned_tasks = {
        let mut store = store_async.write_lock().await;
        std::mem::take(&mut env.as_mut(&mut store).spawned_tasks)
    };
    for task in &spawned_tasks {
        task.abort.abort();
    }
    for task in spawned_tasks {
        let _ = task.done.await;
    }

    let store = store_async.read_lock().await;
    Ok(env.as_ref(&store).logs.clone())
}

#[cfg(not(target_arch = "wasm32"))]
fn run_greenthread_test_native(wat: &[u8]) -> Result<Vec<String>> {
    let mut local_pool = futures::executor::LocalPool::new();
    let spawner = TestSpawner {
        inner: local_pool.spawner(),
    };
    local_pool.run_until(run_greenthread_test(wat, spawner))
}

fn assert_logs(logs: &[String], expected: &[&str]) {
    assert_eq!(logs.len(), expected.len());
    for (index, expected) in expected.iter().enumerate() {
        assert_eq!(
            logs[index], *expected,
            "Log entry mismatch at index {index}: {logs:?}"
        );
    }
}

#[cfg(not(target_arch = "wasm32"))]
#[test]
#[cfg_attr(
    feature = "v8-default",
    ignore = "async functions are not supported by the default v8 backend"
)]
fn green_threads_switch_and_log_in_expected_order() -> Result<()> {
    let logs = run_greenthread_test_native(SWITCHING_WAT)?;
    assert_logs(&logs, SWITCHING_LOGS);

    Ok(())
}

#[cfg(target_arch = "wasm32")]
#[wasm_bindgen_test]
async fn green_threads_switch_and_log_in_expected_order() {
    let logs = run_greenthread_test(SWITCHING_WAT, TestSpawner {})
        .await
        .unwrap();
    assert_logs(&logs, SWITCHING_LOGS);
}

#[cfg(not(target_arch = "wasm32"))]
#[test]
#[cfg_attr(
    feature = "v8-default",
    ignore = "async functions are not supported by the default v8 backend"
)]
fn green_threads_switch_main_crashed() -> Result<()> {
    let logs = run_greenthread_test_native(REGRESSION_WAT)?;
    assert_logs(&logs, REGRESSION_LOGS);

    Ok(())
}

#[cfg(target_arch = "wasm32")]
#[wasm_bindgen_test]
async fn green_threads_switch_main_crashed() {
    let logs = run_greenthread_test(REGRESSION_WAT, TestSpawner {})
        .await
        .unwrap();
    assert_logs(&logs, REGRESSION_LOGS);
}