diff --git a/src/passes/Asyncify.cpp b/src/passes/Asyncify.cpp index feea9adf7c2..989e7b29d31 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,53 @@ 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); + // Each operand gets its own copy: the IR is a tree, so an Expression* + // cannot be shared between the condition and the grow. + auto deficit = makeBinary(Abstract::getBinary(pointerType, Abstract::Sub), + makeRefIndex(heapType, count), + makeTableSize(table)); + auto full = makeBinary(Abstract::getBinary(pointerType, Abstract::GtU), + makeRefIndex(heapType, count), + makeTableSize(table)); + return makeIf(full, + makeDrop(makeTableGrow( + table, + LiteralUtils::makeZero(Type(heapType, Nullable), wasm), + deficit))); + } + Expression* makeGetStackPos() { return makeLoad(pointerType.getByteSize(), false, @@ -1443,6 +1514,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 +1673,121 @@ 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); - } + // Next slot for a local: a table slot for reference components, a byte + // range in linear memory for everything else + struct SpillCursors { + AsyncifyBuilder& builder; + RefCountMap refOffsets; + Index memOffset = 0; + + SpillCursors(AsyncifyBuilder& builder) : builder(builder) {} + + // Table slot for the next reference + Expression* nextRef(HeapType heapType) { + return builder.makeRefIndex(heapType, refOffsets[heapType]++); + } + + // Byte offset of the next non-reference component + Index takeMem(Index size) { + auto offset = memOffset; + memOffset += size; + return offset; + } + }; + + // Reads one component back from its spill slot + Expression* + makeComponentLoading(Type type, Index tempIndex, SpillCursors& cursors) { + if (type.isRef()) { + auto heapType = type.getHeapType(); + // The table is nullable, so a non-nullable local needs a cast back. + Expression* load = builder->makeTableGet(builder->getRefTable(heapType), + cursors.nextRef(heapType), + Type(heapType, Nullable)); + if (!type.isNullable()) { + load = builder->makeRefAs(RefAsNonNull, load); } - }; + return load; + } + auto size = getByteSize(type); + assert(size % STACK_ALIGN == 0); + // TODO: higher alignment? + return builder->makeLoad( + size, + true, + cursors.takeMem(size), + STACK_ALIGN, + builder->makeLocalGet(tempIndex, builder->pointerType), + type, + asyncifyMemory); + } + + // Writes one component to its spill slot + Expression* makeComponentSaving(Type type, + Expression* value, + Index tempIndex, + SpillCursors& cursors) { + if (type.isRef()) { + return builder->makeTableSet(builder->getRefTable(type.getHeapType()), + cursors.nextRef(type.getHeapType()), + value); + } + auto size = getByteSize(type); + assert(size % STACK_ALIGN == 0); + // TODO: higher alignment? + return builder->makeStore( + size, + cursors.takeMem(size), + STACK_ALIGN, + builder->makeLocalGet(tempIndex, builder->pointerType), + value, + type, + asyncifyMemory); + } + + 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; + SpillCursors cursors(*builder); for (Index i = 0; i < numLocals; i++) { if (!relevantLiveLocals.contains(i)) { continue; @@ -1635,18 +1795,7 @@ 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; + loads.push_back(makeComponentLoading(type, tempIndex, cursors)); } Expression* load; if (loads.size() == 1) { @@ -1666,13 +1815,18 @@ 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())); - Index offset = 0; + // 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)); + } + SpillCursors cursors(*builder); for (Index i = 0; i < numLocals; i++) { if (!relevantLiveLocals.contains(i)) { continue; @@ -1680,26 +1834,19 @@ 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; + block->list.push_back( + makeComponentSaving(type, localGet, tempIndex, cursors)); ++j; } } - block->list.push_back(builder->makeIncStackPos(offset)); + block->list.push_back(builder->makeIncStackPos(cursors.memOffset)); + for (const auto& [heapType, total] : cursors.refOffsets) { + block->list.push_back(builder->makeIncRefPos(heapType, int32_t(total))); + } block->finalize(); return block; } @@ -1717,10 +1864,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 +2034,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 +2106,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 +2158,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 +2224,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)