use crate::wit::{AdapterKind, Instruction, NonstandardWitSection, WasmBindgenAux};
use anyhow::{anyhow, bail, Error};
pub const TASK_POLL_EXPORT: &str = "__wbg_jspi_task_poll";
use std::collections::HashMap;
use walrus::ir::{self, BinaryOp, MemArg, UnaryOp, Value};
use walrus::{
ConstExpr, ExportItem, FunctionBuilder, FunctionId, GlobalId, InstrSeqBuilder, MemoryId,
Module, RefType, TagId, ValType,
};
pub fn run(
module: &mut Module,
aux: &mut WasmBindgenAux,
wit: &mut NonstandardWitSection,
externref: bool,
) -> Result<(), Error> {
let mut jspi_exports = Vec::new();
for (id, export) in crate::sorted_iter(&aux.export_map) {
if !export.jspi {
continue;
}
if let Some(adapter) = wit.adapters.get(id) {
let AdapterKind::Local { instructions } = &adapter.kind else {
continue;
};
let export_id = instructions
.iter()
.find_map(|i| match i.instr {
Instruction::CallExport(e) => Some(e),
_ => None,
})
.ok_or_else(|| {
anyhow!("jspi export adapter never calls the underlying function")
})?;
jspi_exports.push(export_id);
}
}
let spawn_imports = module
.imports
.iter()
.filter(|i| i.name.contains("__wbindgen_jspi_spawn_"))
.filter_map(|i| match i.kind {
walrus::ImportKind::Function(f) => Some(f),
_ => None,
})
.collect::<Vec<_>>();
let task_poll_export = module
.exports
.iter()
.find(|e| e.name == TASK_POLL_EXPORT)
.map(|e| e.id());
if let Some(export_id) = task_poll_export {
if !jspi_exports.is_empty() && !spawn_imports.is_empty() {
jspi_exports.push(export_id);
} else {
module.exports.delete(export_id);
for func in spawn_imports {
module.replace_imported_func(func, |(body, _)| {
body.unreachable();
})?;
}
}
}
let suspending_imports = wit
.implements
.iter()
.filter(|(_, _, adapter)| aux.imports_with_suspending.contains(adapter))
.map(|(_, func, adapter)| (*func, aux.imports_with_catch.contains(adapter)))
.collect::<Vec<_>>();
let in_context_import = module
.imports
.iter()
.find(|i| i.name.contains("__wbindgen_jspi_in_context"))
.and_then(|i| match i.kind {
walrus::ImportKind::Function(f) => Some(f),
_ => None,
});
if jspi_exports.is_empty() && suspending_imports.is_empty() {
if let Some(func) = in_context_import {
module.replace_imported_func(func, |(body, _)| {
body.i32_const(0);
})?;
}
return Ok(());
}
if !externref {
bail!(
"JSPI support requires reference types: the resume value of a \
suspending import bypasses the JS shim's return conversion, so \
it must travel as an externref (enabled by default since \
Rust 1.82)"
);
}
if module.memories.iter().any(|m| m.shared) {
bail!("JSPI is not supported with threads/atomics");
}
let sp = aux.stack_pointer.ok_or_else(|| {
anyhow!(
"could not locate the `__stack_pointer` global in the Wasm module; \
JSPI requires it so that suspended fibers' shadow stacks can be \
saved and restored — ensure the linker retains the symbol name"
)
})?;
let memory = crate::wasm_conventions::get_memory(module)?;
let malloc = aux
.jspi_malloc
.ok_or_else(|| anyhow!("JSPI support requires `__wbindgen_malloc` to be present"))?;
let free = aux
.jspi_free
.ok_or_else(|| anyhow!("JSPI support requires `__wbindgen_free` to be present"))?;
let ptr_ty = module.globals.get(sp).ty;
let base = module
.globals
.add_local(ptr_ty, true, false, const_zero(ptr_ty));
module.globals.get_mut(base).name = Some("__jspi_stack_base".to_string());
let suspended = module
.globals
.add_local(ValType::I32, true, false, const_zero(ValType::I32));
module.globals.get_mut(suspended).name = Some("__jspi_suspended".to_string());
let top = module
.globals
.add_local(ptr_ty, true, false, const_zero(ptr_ty));
module.globals.get_mut(top).name = Some("__jspi_stack_top".to_string());
crate::wasm_conventions::get_or_insert_start_builder(module)
.func_body()
.global_get(sp)
.global_set(top);
let ctx = JspiContext {
sp,
base,
suspended,
top,
memory,
malloc,
free,
ptr_ty,
};
if let Some(func) = in_context_import {
module.replace_imported_func(func, |(body, _)| {
body.global_get(base);
match ptr_ty {
ValType::I64 => body.unop(UnaryOp::I64Eqz),
_ => body.unop(UnaryOp::I32Eqz),
};
body.unop(UnaryOp::I32Eqz);
})?;
}
for export_id in jspi_exports {
let inner = match module.exports.get(export_id).item {
ExportItem::Function(f) => f,
_ => bail!("jspi export is not a function"),
};
let wrapper = wrap_export(module, inner, ctx);
module.exports.get_mut(export_id).item = wrapper.into();
}
if suspending_imports.is_empty() {
return Ok(());
}
let restore = make_restore_helper(module, ctx);
let has_catch = suspending_imports.iter().any(|(_, c)| *c);
let rejection = if has_catch {
let addr = aux.jspi_rejected.ok_or_else(|| {
anyhow!(
"could not locate the `__wbindgen_jspi_rejected` flag static; \
it is defined by the wasm-bindgen runtime"
)
})?;
let js_tag = crate::transforms::catch_handler::get_or_import_js_tag(module);
Some(Rejection { js_tag, addr })
} else {
None
};
let mut wrappers = HashMap::new();
for (import, catch) in suspending_imports {
if wrappers.contains_key(&import) {
continue;
}
let rejection = if catch {
let ty = module.types.get(module.funcs.get(import).ty());
if ty.results() != [ValType::Ref(RefType::EXTERNREF)] {
bail!(
"unexpected ABI for a `catch` suspending import: expected \
an externref return, found {:?}",
ty.results()
);
}
rejection
} else {
None
};
let wrapper = wrap_suspending(module, import, ctx, restore, rejection);
wrappers.insert(import, wrapper);
}
rewrite_calls(module, &wrappers);
for (_, func, adapter) in wit.implements.iter_mut() {
if aux.imports_with_suspending.contains(adapter) {
if let Some(wrapper) = wrappers.get(func) {
*func = *wrapper;
}
}
}
Ok(())
}
#[derive(Clone, Copy)]
struct JspiContext {
sp: GlobalId,
base: GlobalId,
suspended: GlobalId,
top: GlobalId,
memory: MemoryId,
malloc: FunctionId,
free: FunctionId,
ptr_ty: ValType,
}
#[derive(Clone, Copy)]
struct Rejection {
js_tag: TagId,
addr: u64,
}
fn const_zero(ty: ValType) -> ConstExpr {
match ty {
ValType::I64 => ConstExpr::Value(Value::I64(0)),
_ => ConstExpr::Value(Value::I32(0)),
}
}
fn wrap_export(module: &mut Module, inner: FunctionId, ctx: JspiContext) -> FunctionId {
let ty = module.types.get(module.funcs.get(inner).ty());
let params = ty.params().to_vec();
let results = ty.results().to_vec();
let results_ty: ir::InstrSeqType = match results.len() {
0 => ir::InstrSeqType::Simple(None),
1 => ir::InstrSeqType::Simple(Some(results[0])),
_ => module.types.add(&[], &results).into(),
};
let exnref_ty: ir::InstrSeqType = ValType::Ref(RefType::EXNREF).into();
let mut builder = FunctionBuilder::new(&mut module.types, ¶ms, &results);
let param_locals: Vec<_> = params.iter().map(|ty| module.locals.add(*ty)).collect();
let prev = module.locals.add(ctx.ptr_ty);
let base = module.locals.add(ctx.ptr_ty);
let prev_suspended = module.locals.add(ValType::I32);
let exn = module.locals.add(ValType::Ref(RefType::EXNREF));
let sub = ptr_sub(ctx.ptr_ty);
let emit_exit = |seq: &mut InstrSeqBuilder| {
seq.global_get(ctx.suspended).if_else(
None,
|then| {
then.global_get(ctx.top).global_set(ctx.sp);
match ctx.ptr_ty {
ValType::I64 => then.i64_const(0),
_ => then.i32_const(0),
};
then.global_set(ctx.base);
},
|else_| {
else_.local_get(base).global_set(ctx.sp);
else_.local_get(prev).global_set(ctx.base);
},
);
seq.local_get(prev_suspended).global_set(ctx.suspended);
};
let try_seq = builder.dangling_instr_seq(results_ty).id();
let catch_seq = builder.dangling_instr_seq(exnref_ty).id();
let done_seq = builder.dangling_instr_seq(results_ty).id();
{
let mut seq = builder.instr_seq(try_seq);
for local in ¶m_locals {
seq.local_get(*local);
}
seq.call(inner);
}
{
let mut seq = builder.instr_seq(catch_seq);
seq.instr(ir::TryTable {
seq: try_seq,
catches: vec![ir::TryTableCatch::CatchAllRef { label: catch_seq }],
});
emit_exit(&mut seq);
seq.br(done_seq);
}
{
let mut seq = builder.instr_seq(done_seq);
seq.instr(ir::Block { seq: catch_seq });
seq.local_set(exn);
emit_exit(&mut seq);
seq.local_get(exn);
seq.instr(ir::ThrowRef {});
}
let mut body = builder.func_body();
body.global_get(ctx.base).local_set(prev);
body.global_get(ctx.sp).local_tee(base).global_set(ctx.base);
body.local_get(base);
push_align(&mut body, ctx.ptr_ty);
body.binop(sub).global_set(ctx.sp);
body.local_get(base);
push_align(&mut body, ctx.ptr_ty);
body.binop(sub).local_get(prev);
store_ptr(&mut body, ctx);
body.global_get(ctx.suspended).local_set(prev_suspended);
body.i32_const(0).global_set(ctx.suspended);
body.instr(ir::Block { seq: done_seq });
let wrapper = builder.finish(param_locals, &mut module.funcs);
let name = module.funcs.get(inner).name.clone();
module.funcs.get_mut(wrapper).name = name.map(|n| format!("{n} jspi wrapper"));
wrapper
}
fn store_ptr(body: &mut InstrSeqBuilder, ctx: JspiContext) {
let (kind, align) = match ctx.ptr_ty {
ValType::I64 => (ir::StoreKind::I64 { atomic: false }, 8),
_ => (ir::StoreKind::I32 { atomic: false }, 4),
};
body.store(ctx.memory, kind, MemArg { align, offset: 0 });
}
fn load_ptr(body: &mut InstrSeqBuilder, ctx: JspiContext) {
let (kind, align) = match ctx.ptr_ty {
ValType::I64 => (ir::LoadKind::I64 { atomic: false }, 8),
_ => (ir::LoadKind::I32 { atomic: false }, 4),
};
body.load(ctx.memory, kind, MemArg { align, offset: 0 });
}
fn make_restore_helper(module: &mut Module, ctx: JspiContext) -> FunctionId {
let params = [ctx.ptr_ty; 3];
let mut builder = FunctionBuilder::new(&mut module.types, ¶ms, &[]);
let base = module.locals.add(ctx.ptr_ty);
let len = module.locals.add(ctx.ptr_ty);
let buf = module.locals.add(ctx.ptr_ty);
let sub = ptr_sub(ctx.ptr_ty);
let mut body = builder.func_body();
body.local_get(base).local_get(len).binop(sub);
body.local_get(buf).local_get(len);
body.instr(ir::MemoryCopy {
src: ctx.memory,
dst: ctx.memory,
});
body.local_get(base).local_get(len).binop(sub);
body.global_set(ctx.sp);
body.local_get(base).global_set(ctx.base);
body.i32_const(1).global_set(ctx.suspended);
body.local_get(buf).local_get(len);
push_align(&mut body, ctx.ptr_ty);
body.call(ctx.free);
let helper = builder.finish(vec![base, len, buf], &mut module.funcs);
module.funcs.get_mut(helper).name = Some("__jspi_restore".to_string());
helper
}
fn ptr_sub(ty: ValType) -> BinaryOp {
match ty {
ValType::I64 => BinaryOp::I64Sub,
_ => BinaryOp::I32Sub,
}
}
fn push_align(body: &mut InstrSeqBuilder, ty: ValType) {
match ty {
ValType::I64 => body.i64_const(16),
_ => body.i32_const(16),
};
}
fn wrap_suspending(
module: &mut Module,
import: FunctionId,
ctx: JspiContext,
restore: FunctionId,
rejection: Option<Rejection>,
) -> FunctionId {
let ty = module.types.get(module.funcs.get(import).ty());
let params = ty.params().to_vec();
let results = ty.results().to_vec();
let results_ty: ir::InstrSeqType = match results.len() {
0 => ir::InstrSeqType::Simple(None),
1 => ir::InstrSeqType::Simple(Some(results[0])),
_ => module.types.add(&[], &results).into(),
};
let exnref_ty: ir::InstrSeqType = ValType::Ref(RefType::EXNREF).into();
let externref_ty: ir::InstrSeqType = ValType::Ref(RefType::EXTERNREF).into();
let mut builder = FunctionBuilder::new(&mut module.types, ¶ms, &results);
let param_locals: Vec<_> = params.iter().map(|ty| module.locals.add(*ty)).collect();
let result_locals: Vec<_> = results.iter().map(|ty| module.locals.add(*ty)).collect();
let base = module.locals.add(ctx.ptr_ty);
let len = module.locals.add(ctx.ptr_ty);
let buf = module.locals.add(ctx.ptr_ty);
let exn = module.locals.add(ValType::Ref(RefType::EXNREF));
let eqz = match ctx.ptr_ty {
ValType::I64 => UnaryOp::I64Eqz,
_ => UnaryOp::I32Eqz,
};
let sub = ptr_sub(ctx.ptr_ty);
let call_restore = |seq: &mut InstrSeqBuilder| {
seq.local_get(base).local_get(len).local_get(buf);
seq.call(restore);
};
let stash_results = |seq: &mut InstrSeqBuilder| {
for local in result_locals.iter().rev() {
seq.local_set(*local);
}
};
let unstash_results = |seq: &mut InstrSeqBuilder| {
for local in &result_locals {
seq.local_get(*local);
}
};
let store_rejected = |seq: &mut InstrSeqBuilder, rejection: Rejection, value: i32| {
match ctx.ptr_ty {
ValType::I64 => seq.i64_const(rejection.addr as i64),
_ => seq.i32_const(rejection.addr as i32),
};
seq.i32_const(value);
seq.store(
ctx.memory,
ir::StoreKind::I32 { atomic: false },
MemArg {
align: 4,
offset: 0,
},
);
};
let try_seq = builder.dangling_instr_seq(results_ty).id();
let rejected_seq = rejection.map(|_| builder.dangling_instr_seq(externref_ty).id());
let catch_all_seq = builder.dangling_instr_seq(exnref_ty).id();
let done_seq = builder.dangling_instr_seq(results_ty).id();
{
let mut seq = builder.instr_seq(try_seq);
for local in ¶m_locals {
seq.local_get(*local);
}
seq.call(import);
}
let innermost = rejected_seq.unwrap_or(catch_all_seq);
{
let mut catches = Vec::new();
if let (Some(rejection), Some(rejected_seq)) = (rejection, rejected_seq) {
catches.push(ir::TryTableCatch::Catch {
tag: rejection.js_tag,
label: rejected_seq,
});
}
catches.push(ir::TryTableCatch::CatchAllRef {
label: catch_all_seq,
});
let mut seq = builder.instr_seq(innermost);
seq.instr(ir::TryTable {
seq: try_seq,
catches,
});
stash_results(&mut seq);
call_restore(&mut seq);
if let Some(rejection) = rejection {
store_rejected(&mut seq, rejection, 0);
}
unstash_results(&mut seq);
seq.br(done_seq);
}
if let (Some(rejection), Some(rejected_seq)) = (rejection, rejected_seq) {
let mut seq = builder.instr_seq(catch_all_seq);
seq.instr(ir::Block { seq: rejected_seq });
stash_results(&mut seq);
call_restore(&mut seq);
store_rejected(&mut seq, rejection, 1);
unstash_results(&mut seq);
seq.br(done_seq);
}
{
let mut seq = builder.instr_seq(done_seq);
seq.instr(ir::Block { seq: catch_all_seq });
seq.local_set(exn);
call_restore(&mut seq);
seq.local_get(exn);
seq.instr(ir::ThrowRef {});
}
let mut body = builder.func_body();
body.global_get(ctx.base).local_tee(base).unop(eqz).if_else(
None,
|then| {
for local in ¶m_locals {
then.local_get(*local);
}
then.call(import);
then.instr(ir::Return {});
},
|_| {},
);
body.local_get(base).global_get(ctx.sp).binop(sub);
body.local_set(len);
body.local_get(len);
push_align(&mut body, ctx.ptr_ty);
body.call(ctx.malloc).local_set(buf);
body.local_get(buf).global_get(ctx.sp).local_get(len);
body.instr(ir::MemoryCopy {
src: ctx.memory,
dst: ctx.memory,
});
body.local_get(base).global_set(ctx.sp);
body.local_get(base);
push_align(&mut body, ctx.ptr_ty);
body.binop(sub);
load_ptr(&mut body, ctx);
body.global_set(ctx.base);
body.instr(ir::Block { seq: done_seq });
let wrapper = builder.finish(param_locals, &mut module.funcs);
let name = module.funcs.get(import).name.clone();
module.funcs.get_mut(wrapper).name = Some(match name {
Some(n) => format!("{n} suspending wrapper"),
None => "suspending wrapper".to_string(),
});
wrapper
}
fn rewrite_calls(module: &mut Module, wrappers: &HashMap<FunctionId, FunctionId>) {
let wrapper_ids: std::collections::HashSet<_> = wrappers.values().copied().collect();
for (func_id, func) in module.funcs.iter_local_mut() {
if wrapper_ids.contains(&func_id) {
continue;
}
let entry = func.entry_block();
ir::dfs_pre_order_mut(&mut CallRewriter { wrappers }, func, entry);
}
}
struct CallRewriter<'a> {
wrappers: &'a HashMap<FunctionId, FunctionId>,
}
impl ir::VisitorMut for CallRewriter<'_> {
fn start_instr_seq_mut(&mut self, seq: &mut ir::InstrSeq) {
for (instr, _) in seq.instrs.iter_mut() {
let func = match instr {
ir::Instr::Call(ir::Call { func }) => func,
ir::Instr::ReturnCall(ir::ReturnCall { func }) => func,
_ => continue,
};
if let Some(wrapper) = self.wrappers.get(func) {
*func = *wrapper;
}
}
}
}