From 4d3c936439cd6fb454ee34744096b006999c5f27 Mon Sep 17 00:00:00 2001 From: NellowTCS Date: Fri, 25 Sep 2026 23:28:33 -0600 Subject: [PATCH] spill reference locals to growable tables --- src/passes/Asyncify.cpp | 312 ++++++++++++++---- test/lit/passes/asyncify-reference-types.wast | 43 +++ 2 files changed, 287 insertions(+), 68 deletions(-) create mode 100644 test/lit/passes/asyncify-reference-types.wast diff --git a/src/passes/Asyncify.cpp b/src/passes/Asyncify.cpp index feea9adf7c2..a36dc73e2aa 100644 --- a/src/passes/Asyncify.cpp +++ b/src/passes/Asyncify.cpp @@ -342,6 +342,9 @@ // of their original range. // +#include +#include + #include "asmjs/shared-constants.h" #include "cfg/liveness-traversal.h" #include "ir/effects.h" @@ -387,6 +390,18 @@ enum class DataOffset { BStackPos = 0, BStackEnd = 4, BStackEnd64 = 8 }; const auto STACK_ALIGN = 4; +// Reference spill tables start at the acyclic per-function sum and grow on +// unwind (for recursion); overflow traps on the out-of-bounds table.set. +const uint32_t REF_TABLE_MAX = 1 << 26; + +static Name refTableName(HeapType ht) { + return Name(std::string("__asyncify_ref_table_") + ht.toString()); +} + +static Name refCursorName(HeapType ht) { + return Name(std::string("__asyncify_ref_pos_") + ht.toString()); +} + // A helper class for managing fake global names. Creates the globals and // provides mappings for using them. // Fake globals are used to stash and then use return values from calls. We need @@ -875,6 +890,15 @@ static bool doesCall(Expression* curr) { return curr->is() || curr->is(); } +// Interned types (Type, HeapType) carry no ordering operators, but their +// canonical ids give a stable, deterministic order for map iteration. +struct HeapTypeIdLess { + bool operator()(HeapType a, HeapType b) const { + return a.getID() < b.getID(); + } +}; +using RefCountMap = std::map; + class AsyncifyBuilder : public Builder { public: Module& wasm; @@ -885,6 +909,51 @@ class AsyncifyBuilder : public Builder { : Builder(wasm), wasm(wasm), pointerType(pointerType), asyncifyMemory(asyncifyMemory) {} + Name getRefTable(HeapType heapType) { return refTableName(heapType); } + + Expression* makeGetRefPos(HeapType heapType) { + return makeGlobalGet(refCursorName(heapType), pointerType); + } + + // The table index for the `offset`th spilled slot of this heap type within + // the current frame, i.e. the cursor global advanced by the slot offset. + Expression* makeRefIndex(HeapType heapType, Index offset) { + return makeBinary(Abstract::getBinary(pointerType, Abstract::Add), + makeGetRefPos(heapType), + makeConst(Literal::makeFromInt64(offset, pointerType))); + } + + Expression* makeIncRefPos(HeapType heapType, int32_t by) { + if (by == 0) { + return makeNop(); + } + return makeGlobalSet( + refCursorName(heapType), + makeBinary(Abstract::getBinary(pointerType, Abstract::Add), + makeGetRefPos(heapType), + makeConst(Literal::makeFromInt64(by, pointerType)))); + } + + // Ensure the table has room for `count` more slots at the cursor, so that + // recursive callers (live on the stack more than once) still fit. + Expression* makeEnsureRefCapacity(HeapType heapType, uint32_t count) { + if (count == 0) { + return makeNop(); + } + auto table = getRefTable(heapType); + auto need = makeRefIndex(heapType, count); + auto deficit = makeBinary(Abstract::getBinary(pointerType, Abstract::Sub), + need, + makeTableSize(table)); + return makeIf(makeBinary(Abstract::getBinary(pointerType, Abstract::GtU), + need, + makeTableSize(table)), + makeDrop(makeTableGrow( + table, + LiteralUtils::makeZero(Type(heapType, Nullable), wasm), + deficit))); + } + Expression* makeGetStackPos() { return makeLoad(pointerType.getByteSize(), false, @@ -1443,6 +1512,40 @@ struct AsyncifyAssertUnwindCorrectness : Pass { }; // Instrument local saving/restoring. +// TODO: look more precisely inside basic blocks +std::set computeRelevantLiveLocals(Function* func, Module* module) { + struct RelevantLiveLocalsWalker + : public LivenessWalker> { + // Basic blocks that have a possible unwind/rewind in them. + std::set relevantBasicBlocks; + + void visitCall(Call* curr) { + if (!currBasicBlock) { + return; + } + // Blocks where we might unwind/rewind, each marked by a call to + // ASYNCIFY_CHECK_CALL_INDEX right before the real call. + if (curr->target == ASYNCIFY_CHECK_CALL_INDEX) { + relevantBasicBlocks.insert(currBasicBlock); + } + } + }; + + RelevantLiveLocalsWalker walker; + walker.setFunction(func); + walker.walkFunctionInModule(func, module); + std::set relevant; + for (auto* block : walker.liveBlocks) { + if (walker.relevantBasicBlocks.contains(block)) { + for (auto local : block->contents.start) { + relevant.insert(local); + } + } + } + return relevant; +} + struct AsyncifyLocals : public WalkerPass> { bool isFunctionParallel() override { return true; } @@ -1568,66 +1671,50 @@ struct AsyncifyLocals : public WalkerPass> { std::set relevantLiveLocals; void findRelevantLiveLocals(Function* func) { - struct RelevantLiveLocalsWalker - : public LivenessWalker> { - // Basic blocks that have a possible unwind/rewind in them. - std::set relevantBasicBlocks; + relevantLiveLocals = computeRelevantLiveLocals(func, getModule()); + } - void visitCall(Call* curr) { - if (!currBasicBlock) { - return; - } - // Note blocks where we might unwind/rewind, all of which have a - // possible call to ASYNCIFY_CHECK_CALL_INDEX emitted right before the - // actual call. - // Note that each relevant original call was turned into a sequence of - // instructions, one of which is an if and then a call to this special - // intrinsic. We rely on the fact that if a local was live at the - // original call, it also would be in all that sequence of instructions, - // and in particular at the call we look for here (which is right before - // the call, and so anything that has its final use at the call is still - // live here). - if (curr->target == ASYNCIFY_CHECK_CALL_INDEX) { - relevantBasicBlocks.insert(currBasicBlock); - } - } - }; + struct SpillLayout { + Index memTotal = 0; + RefCountMap refTotals; + }; - RelevantLiveLocalsWalker walker; - walker.setFunction(func); - walker.walkFunctionInModule(func, getModule()); - // The relevant live locals are ones that are alive at an unwind/rewind - // location. TODO look more precisely inside basic blocks, as one might stop - // being live in the middle - for (auto* block : walker.liveBlocks) { - if (walker.relevantBasicBlocks.contains(block)) { - for (auto local : block->contents.start) { - relevantLiveLocals.insert(local); + SpillLayout computeSpillLayout() { + SpillLayout layout; + auto* func = getFunction(); + for (Index i = 0; i < func->getNumLocals(); i++) { + if (!relevantLiveLocals.contains(i)) { + continue; + } + for (const auto& type : func->getLocalType(i)) { + if (type.isRef()) { + layout.refTotals[type.getHeapType()] += 1; + } else { + layout.memTotal += getByteSize(type); } } } + return layout; } Expression* makeLocalLoading() { if (relevantLiveLocals.empty()) { return builder->makeNop(); } + auto layout = computeSpillLayout(); auto* func = getFunction(); auto numLocals = func->getNumLocals(); - Index total = 0; - for (Index i = 0; i < numLocals; i++) { - if (!relevantLiveLocals.contains(i)) { - continue; - } - total += getByteSize(func->getLocalType(i)); - } auto* block = builder->makeBlock(); - block->list.push_back(builder->makeIncStackPos(-total)); + + block->list.push_back(builder->makeIncStackPos(-layout.memTotal)); + for (const auto& [heapType, count] : layout.refTotals) { + block->list.push_back(builder->makeIncRefPos(heapType, -int32_t(count))); + } auto tempIndex = builder->addVar(func, builder->pointerType); block->list.push_back( builder->makeLocalSet(tempIndex, builder->makeGetStackPos())); Index offset = 0; + RefCountMap refOffsets; for (Index i = 0; i < numLocals; i++) { if (!relevantLiveLocals.contains(i)) { continue; @@ -1635,18 +1722,33 @@ struct AsyncifyLocals : public WalkerPass> { auto localType = func->getLocalType(i); SmallVector loads; for (const auto& type : localType) { - auto size = getByteSize(type); - assert(size % STACK_ALIGN == 0); - // TODO: higher alignment? - loads.push_back(builder->makeLoad( - size, - true, - offset, - STACK_ALIGN, - builder->makeLocalGet(tempIndex, builder->pointerType), - type, - asyncifyMemory)); - offset += size; + if (type.isRef()) { + auto heapType = type.getHeapType(); + auto& refOffset = refOffsets[heapType]; + // The table is nullable, so a non-nullable local needs a cast back. + Expression* load = + builder->makeTableGet(builder->getRefTable(heapType), + builder->makeRefIndex(heapType, refOffset), + Type(heapType, Nullable)); + if (!type.isNullable()) { + load = builder->makeRefAs(RefAsNonNull, load); + } + loads.push_back(load); + refOffset += 1; + } else { + auto size = getByteSize(type); + assert(size % STACK_ALIGN == 0); + // TODO: higher alignment? + loads.push_back(builder->makeLoad( + size, + true, + offset, + STACK_ALIGN, + builder->makeLocalGet(tempIndex, builder->pointerType), + type, + asyncifyMemory)); + offset += size; + } } Expression* load; if (loads.size() == 1) { @@ -1666,13 +1768,19 @@ struct AsyncifyLocals : public WalkerPass> { if (relevantLiveLocals.empty()) { return builder->makeNop(); } + auto layout = computeSpillLayout(); auto* func = getFunction(); auto numLocals = func->getNumLocals(); auto* block = builder->makeBlock(); auto tempIndex = builder->addVar(func, builder->pointerType); block->list.push_back( builder->makeLocalSet(tempIndex, builder->makeGetStackPos())); + // Ensure room for this frame's reference spill before writing any of it + for (const auto& [heapType, total] : layout.refTotals) { + block->list.push_back(builder->makeEnsureRefCapacity(heapType, total)); + } Index offset = 0; + RefCountMap refOffsets; for (Index i = 0; i < numLocals; i++) { if (!relevantLiveLocals.contains(i)) { continue; @@ -1680,26 +1788,39 @@ struct AsyncifyLocals : public WalkerPass> { auto localType = func->getLocalType(i); size_t j = 0; for (const auto& type : localType) { - auto size = getByteSize(type); Expression* localGet = builder->makeLocalGet(i, localType); if (localType.size() > 1) { localGet = builder->makeTupleExtract(localGet, j); } - assert(size % STACK_ALIGN == 0); - // TODO: higher alignment? - block->list.push_back(builder->makeStore( - size, - offset, - STACK_ALIGN, - builder->makeLocalGet(tempIndex, builder->pointerType), - localGet, - type, - asyncifyMemory)); - offset += size; + if (type.isRef()) { + auto heapType = type.getHeapType(); + auto& refOffset = refOffsets[heapType]; + block->list.push_back( + builder->makeTableSet(builder->getRefTable(heapType), + builder->makeRefIndex(heapType, refOffset), + localGet)); + refOffset += 1; + } else { + auto size = getByteSize(type); + assert(size % STACK_ALIGN == 0); + // TODO: higher alignment? + block->list.push_back(builder->makeStore( + size, + offset, + STACK_ALIGN, + builder->makeLocalGet(tempIndex, builder->pointerType), + localGet, + type, + asyncifyMemory)); + offset += size; + } ++j; } } block->list.push_back(builder->makeIncStackPos(offset)); + for (const auto& [heapType, total] : refOffsets) { + block->list.push_back(builder->makeIncRefPos(heapType, int32_t(total))); + } block->finalize(); return block; } @@ -1717,10 +1838,12 @@ struct AsyncifyLocals : public WalkerPass> { builder->makeIncStackPos(4)); } + // Only called for non-reference types unsigned getByteSize(Type type) { if (!type.hasByteSize()) { - Fatal() << "Asyncify does not yet support non-number types, like " - "references (see " + Fatal() << "Asyncify cannot spill a value of type " << type + << " into the linear-memory spill stack, as it has no linear " + "memory representation (see " "https://github.com/WebAssembly/binaryen/issues/3739)"; } return type.getByteSize(); @@ -1885,6 +2008,7 @@ struct Asyncify : public Pass { runner.setValidateGlobally(false); runner.run(); } + setupRefSpill(module, instrumentedFuncs); if (asserts) { // Add asserts in non-instrumented code. Note we do not use an // instrumented pass runner here as we do want to run on all functions. @@ -1956,6 +2080,44 @@ struct Asyncify : public Pass { } } + // Reference locals have no linear-memory representation, so they are spilled + // into pass-created tables + void setupRefSpill(Module* module, + const PassUtils::FuncSet& instrumentedFuncs) { + Builder builder(*module); + RefCountMap counts; + for (auto* func : instrumentedFuncs) { + if (!func->body) { + continue; + } + for (auto local : computeRelevantLiveLocals(func, module)) { + for (const auto& type : func->getLocalType(local)) { + if (type.isRef()) { + counts[type.getHeapType()] += 1; + } + } + } + } + for (const auto& [heapType, count] : counts) { + if (count == 0) { + continue; + } + + auto table = Builder::makeTable(refTableName(heapType), + Type(heapType, Nullable), + count, // initial size + REF_TABLE_MAX, // max size (growable) + pointerType, // address type + nullptr); // init: default null + module->addTable(std::move(table)); + module->addGlobal(builder.makeGlobal(refCursorName(heapType), + pointerType, + builder.makeConst(pointerType), + Builder::Mutable)); + refHeapTypes.insert(heapType); + } + } + void addFunctions(Module* module) { Builder builder(*module); auto makeFunction = [&](Name name, bool setData, State state) { @@ -1970,6 +2132,19 @@ struct Asyncify : public Pass { body->list.push_back(builder.makeGlobalSet( ASYNCIFY_DATA, builder.makeLocalGet(0, pointerType))); } + // Asyncify only supports one pause at a time. + if (name == ASYNCIFY_START_UNWIND) { + for (auto heapType : refHeapTypes) { + body->list.push_back(builder.makeIf( + builder.makeBinary( + Abstract::getBinary(pointerType, Abstract::Ne), + builder.makeGlobalGet(refCursorName(heapType), pointerType), + builder.makeConst(pointerType)), + builder.makeUnreachable())); + body->list.push_back(builder.makeGlobalSet( + refCursorName(heapType), builder.makeConst(pointerType))); + } + } // Verify the data is valid. auto* stackPos = builder.makeLoad(pointerType.getByteSize(), @@ -2023,6 +2198,7 @@ struct Asyncify : public Pass { Type pointerType; Name asyncifyMemory; + std::set refHeapTypes; }; Pass* createAsyncifyPass() { return new Asyncify(); } diff --git a/test/lit/passes/asyncify-reference-types.wast b/test/lit/passes/asyncify-reference-types.wast new file mode 100644 index 00000000000..18db15ff78f --- /dev/null +++ b/test/lit/passes/asyncify-reference-types.wast @@ -0,0 +1,43 @@ +;; Asyncify spills reference-typed locals (externref/funcref) into growable +;; per-heap-type tables. + +;; RUN: foreach %s %t wasm-opt --enable-reference-types --asyncify -S -o - | filecheck %s + +(module + (import "env" "unwind" (func $unwind)) + (memory 1) + (func (export "run") (param $r externref) (result i32) + (local $x i32) + (call $unwind) + (drop (ref.is_null (local.get $r))) + (local.get $x) + ) +) + +;; CHECK: (global $__asyncify_ref_pos_extern (mut i32) (i32.const 0)) +;; CHECK: (table $__asyncify_ref_table_extern 1 67108864 externref) + +;; On rewind, the externref local is restored from the table at the cursor. +;; CHECK: (global.set $__asyncify_ref_pos_extern +;; CHECK: (local.set $r +;; CHECK: (table.get $__asyncify_ref_table_extern +;; CHECK: (i32.add +;; CHECK: (global.get $__asyncify_ref_pos_extern) + +;; On unwind, the table is grown if the cursor has reached the current size, and +;; the local is spilled to the table. Growth is what makes recursion safe. +;; CHECK: (table.size $__asyncify_ref_table_extern) +;; CHECK: (table.grow $__asyncify_ref_table_extern +;; CHECK: (table.set $__asyncify_ref_table_extern +;; CHECK: (global.get $__asyncify_ref_pos_extern) + +;; The cursor is reset at the start of each unwind. +;; A nonzero cursor at the start of an unwind means two pauses overlap, which +;; asyncify does not support, so it traps before resetting the cursor. +;; CHECK: (func $asyncify_start_unwind +;; CHECK: (i32.ne +;; CHECK: (global.get $__asyncify_ref_pos_extern) +;; CHECK: (i32.const 0) +;; CHECK: (unreachable) +;; CHECK: (global.set $__asyncify_ref_pos_extern +;; CHECK: (i32.const 0)