#include "ir/manipulation.h"
#include "ir/module-utils.h"
#include "ir/names.h"
#include "ir/utils.h"
#include "pass.h"
#include "support/space.h"
#include "wasm-binary.h"
#include "wasm-builder.h"
#include "wasm.h"
namespace wasm {
namespace {
struct Range {
bool isZero;
size_t start;
size_t end;
};
using Replacement = std::function<Expression*(Function*)>;
using Replacements = std::unordered_map<Expression*, Replacement>;
using Referrers = std::vector<Expression*>;
using ReferrersMap = std::unordered_map<Name, Referrers>;
const size_t MEMORY_INIT_SIZE = 10;
const size_t MEMORY_FILL_SIZE = 9;
const size_t DATA_DROP_SIZE = 3;
Expression*
makeGtShiftedMemorySize(Builder& builder, Module& module, MemoryInit* curr) {
auto mem = module.getMemory(curr->memory);
return builder.makeBinary(
mem->is64() ? GtUInt64 : GtUInt32,
curr->dest,
builder.makeBinary(mem->is64() ? ShlInt64 : ShlInt32,
builder.makeMemorySize(mem->name),
builder.makeConstPtr(16, mem->indexType)));
}
}
struct MemoryPacking : public Pass {
bool requiresNonNullableLocalFixups() override { return false; }
void run(Module* module) override;
bool canOptimize(std::vector<std::unique_ptr<Memory>>& memories,
std::vector<std::unique_ptr<DataSegment>>& dataSegments);
void optimizeSegmentOps(Module* module);
void getSegmentReferrers(Module* module, ReferrersMap& referrers);
void dropUnusedSegments(Module* module,
std::vector<std::unique_ptr<DataSegment>>& segments,
ReferrersMap& referrers);
bool canSplit(const std::unique_ptr<DataSegment>& segment,
const Referrers& referrers);
void calculateRanges(const std::unique_ptr<DataSegment>& segment,
const Referrers& referrers,
std::vector<Range>& ranges);
void createSplitSegments(Builder& builder,
const DataSegment* segment,
std::vector<Range>& ranges,
std::vector<std::unique_ptr<DataSegment>>& packed,
size_t segmentsRemaining);
void createReplacements(Module* module,
const std::vector<Range>& ranges,
const std::vector<Name>& segments,
const Referrers& referrers,
Replacements& replacements);
void replaceSegmentOps(Module* module, Replacements& replacements);
};
void MemoryPacking::run(Module* module) {
if (!canOptimize(module->memories, module->dataSegments)) {
return;
}
bool canHaveSegmentReferrers =
module->features.hasBulkMemory() || module->features.hasGC();
auto& segments = module->dataSegments;
ReferrersMap referrers;
if (canHaveSegmentReferrers) {
optimizeSegmentOps(module);
getSegmentReferrers(module, referrers);
dropUnusedSegments(module, segments, referrers);
}
std::vector<std::unique_ptr<DataSegment>> packed;
Replacements replacements;
Builder builder(*module);
for (size_t index = 0; index < segments.size(); ++index) {
auto& segment = segments[index];
auto& currReferrers = referrers[segment->name];
std::vector<Range> ranges;
if (canSplit(segment, currReferrers)) {
calculateRanges(segment, currReferrers, ranges);
} else {
ranges.push_back({false, 0, segment->data.size()});
}
size_t segmentsRemaining = segments.size() - index;
size_t currSegmentsStart = packed.size();
createSplitSegments(
builder, segment.get(), ranges, packed, segmentsRemaining);
std::vector<Name> currSegmentNames;
for (size_t i = currSegmentsStart; i < packed.size(); ++i) {
currSegmentNames.push_back(packed[i]->name);
}
createReplacements(
module, ranges, currSegmentNames, currReferrers, replacements);
}
segments.swap(packed);
module->updateDataSegmentsMap();
if (canHaveSegmentReferrers) {
replaceSegmentOps(module, replacements);
}
}
bool MemoryPacking::canOptimize(
std::vector<std::unique_ptr<Memory>>& memories,
std::vector<std::unique_ptr<DataSegment>>& dataSegments) {
if (memories.empty() || memories.size() > 1) {
return false;
}
auto& memory = memories[0];
if (memory->imported() && !getPassOptions().zeroFilledMemory) {
return false;
}
if (dataSegments.size() <= 1) {
return true;
}
Address maxAddress = 0;
for (auto& segment : dataSegments) {
if (!segment->isPassive) {
auto* c = segment->offset->dynCast<Const>();
if (!c) {
return false;
}
maxAddress = std::max(
maxAddress, Address(c->value.getUnsigned() + segment->data.size()));
}
}
DisjointSpans space;
for (auto& segment : dataSegments) {
if (!segment->isPassive) {
auto* c = segment->offset->cast<Const>();
Address start = c->value.getUnsigned();
DisjointSpans::Span span{start, start + segment->data.size()};
if (space.addAndCheckOverlap(span)) {
std::cerr << "warning: active memory segments have overlap, which "
<< "prevents some optimizations.\n";
return false;
}
}
}
return true;
}
bool MemoryPacking::canSplit(const std::unique_ptr<DataSegment>& segment,
const Referrers& referrers) {
if (segment->name.is() && segment->name.startsWith("__llvm")) {
return false;
}
for (auto* referrer : referrers) {
if (auto* curr = referrer->dynCast<MemoryInit>()) {
if (segment->isPassive) {
if (!curr->offset->is<Const>() || !curr->size->is<Const>()) {
return false;
}
}
} else if (referrer->is<ArrayNewData>() || referrer->is<ArrayInitData>()) {
return false;
}
}
return segment->isPassive || segment->offset->is<Const>();
}
void MemoryPacking::calculateRanges(const std::unique_ptr<DataSegment>& segment,
const Referrers& referrers,
std::vector<Range>& ranges) {
auto& data = segment->data;
if (data.size() == 0) {
return;
}
size_t start = 0;
while (start < data.size()) {
size_t end = start;
while (end < data.size() && data[end] == 0) {
end++;
}
if (end > start) {
ranges.push_back({true, start, end});
start = end;
}
while (end < data.size() && data[end] != 0) {
end++;
}
if (end > start) {
ranges.push_back({false, start, end});
start = end;
}
}
size_t threshold = 0;
if (segment->isPassive) {
threshold += 2;
size_t edgeThreshold = 0;
for (auto* referrer : referrers) {
if (referrer->is<MemoryInit>()) {
threshold += MEMORY_FILL_SIZE + MEMORY_INIT_SIZE;
edgeThreshold += MEMORY_FILL_SIZE;
} else {
threshold += DATA_DROP_SIZE;
}
}
if (ranges.size() >= 2) {
auto last = ranges.end() - 1;
auto penultimate = ranges.end() - 2;
if (last->isZero && last->end - last->start <= edgeThreshold) {
penultimate->end = last->end;
ranges.erase(last);
}
}
if (ranges.size() >= 2) {
auto first = ranges.begin();
auto second = ranges.begin() + 1;
if (first->isZero && first->end - first->start <= edgeThreshold) {
second->start = first->start;
ranges.erase(first);
}
}
} else {
threshold = 8;
}
std::vector<Range> mergedRanges = {ranges.front()};
size_t i;
for (i = 1; i < ranges.size() - 1; ++i) {
auto left = mergedRanges.end() - 1;
auto curr = ranges.begin() + i;
auto right = ranges.begin() + i + 1;
if (curr->isZero && curr->end - curr->start <= threshold) {
left->end = right->end;
++i;
} else {
mergedRanges.push_back(*curr);
}
}
if (i < ranges.size()) {
mergedRanges.push_back(ranges.back());
}
std::swap(ranges, mergedRanges);
}
void MemoryPacking::optimizeSegmentOps(Module* module) {
struct Optimizer : WalkerPass<PostWalker<Optimizer>> {
bool isFunctionParallel() override { return true; }
bool requiresNonNullableLocalFixups() override { return false; }
std::unique_ptr<Pass> create() override {
return std::make_unique<Optimizer>();
}
bool needsRefinalizing;
void visitMemoryInit(MemoryInit* curr) {
Builder builder(*getModule());
auto* segment = getModule()->getDataSegment(curr->segment);
size_t maxRuntimeSize = segment->isPassive ? segment->data.size() : 0;
bool mustNop = false;
bool mustTrap = false;
auto* offset = curr->offset->dynCast<Const>();
auto* size = curr->size->dynCast<Const>();
if (offset && uint32_t(offset->value.geti32()) > maxRuntimeSize) {
mustTrap = true;
}
if (size && uint32_t(size->value.geti32()) > maxRuntimeSize) {
mustTrap = true;
}
if (offset && size) {
uint64_t offsetVal(offset->value.geti32());
uint64_t sizeVal(size->value.geti32());
if (offsetVal + sizeVal > maxRuntimeSize) {
mustTrap = true;
} else if (offsetVal == 0 && sizeVal == 0) {
mustNop = true;
}
}
assert(!mustNop || !mustTrap);
if (mustNop) {
replaceCurrent(
builder.makeIf(makeGtShiftedMemorySize(builder, *getModule(), curr),
builder.makeUnreachable()));
} else if (mustTrap) {
replaceCurrent(builder.blockify(builder.makeDrop(curr->dest),
builder.makeDrop(curr->offset),
builder.makeDrop(curr->size),
builder.makeUnreachable()));
needsRefinalizing = true;
} else if (!segment->isPassive) {
replaceCurrent(builder.makeIf(
builder.makeBinary(
OrInt32,
makeGtShiftedMemorySize(builder, *getModule(), curr),
builder.makeBinary(OrInt32, curr->offset, curr->size)),
builder.makeUnreachable()));
}
}
void visitDataDrop(DataDrop* curr) {
if (!getModule()->getDataSegment(curr->segment)->isPassive) {
ExpressionManipulator::nop(curr);
}
}
void doWalkFunction(Function* func) {
needsRefinalizing = false;
super::doWalkFunction(func);
if (needsRefinalizing) {
ReFinalize().walkFunctionInModule(func, getModule());
}
}
} optimizer;
optimizer.run(getPassRunner(), module);
}
void MemoryPacking::getSegmentReferrers(Module* module,
ReferrersMap& referrers) {
auto collectReferrers = [&](Function* func, ReferrersMap& referrers) {
if (func->imported()) {
return;
}
struct Collector
: WalkerPass<PostWalker<Collector, UnifiedExpressionVisitor<Collector>>> {
ReferrersMap& referrers;
Collector(ReferrersMap& referrers) : referrers(referrers) {}
void visitExpression(Expression* curr) {
#define DELEGATE_ID curr->_id
#define DELEGATE_START(id) [[maybe_unused]] auto* cast = curr->cast<id>();
#define DELEGATE_GET_FIELD(id, field) cast->field
#define DELEGATE_FIELD_TYPE(id, field)
#define DELEGATE_FIELD_HEAPTYPE(id, field)
#define DELEGATE_FIELD_CHILD(id, field)
#define DELEGATE_FIELD_OPTIONAL_CHILD(id, field)
#define DELEGATE_FIELD_INT(id, field)
#define DELEGATE_FIELD_INT_ARRAY(id, field)
#define DELEGATE_FIELD_LITERAL(id, field)
#define DELEGATE_FIELD_NAME(id, field)
#define DELEGATE_FIELD_NAME_VECTOR(id, field)
#define DELEGATE_FIELD_SCOPE_NAME_DEF(id, field)
#define DELEGATE_FIELD_SCOPE_NAME_USE(id, field)
#define DELEGATE_FIELD_SCOPE_NAME_USE_VECTOR(id, field)
#define DELEGATE_FIELD_ADDRESS(id, field)
#define DELEGATE_FIELD_NAME_KIND(id, field, kind) \
if (kind == ModuleItemKind::DataSegment) { \
referrers[cast->field].push_back(curr); \
}
#include "wasm-delegations-fields.def"
}
} collector(referrers);
collector.walkFunctionInModule(func, module);
};
ModuleUtils::ParallelFunctionAnalysis<ReferrersMap> analysis(
*module, collectReferrers);
for (auto& [_, funcReferrersMap] : analysis.map) {
for (auto& [i, segReferrers] : funcReferrersMap) {
referrers[i].insert(
referrers[i].end(), segReferrers.begin(), segReferrers.end());
}
}
}
void MemoryPacking::dropUnusedSegments(
Module* module,
std::vector<std::unique_ptr<DataSegment>>& segments,
ReferrersMap& referrers) {
std::vector<std::unique_ptr<DataSegment>> usedSegments;
for (size_t i = 0; i < segments.size(); ++i) {
bool used = false;
auto referrersIt = referrers.find(segments[i]->name);
bool hasReferrers = referrersIt != referrers.end();
if (segments[i]->isPassive) {
if (hasReferrers) {
for (auto* referrer : referrersIt->second) {
if (!referrer->is<DataDrop>()) {
used = true;
break;
}
}
}
} else {
used = true;
}
if (used) {
usedSegments.push_back(std::move(segments[i]));
} else if (hasReferrers) {
for (auto* referrer : referrersIt->second) {
ExpressionManipulator::nop(referrer);
}
}
}
std::swap(segments, usedSegments);
module->updateDataSegmentsMap();
}
void MemoryPacking::createSplitSegments(
Builder& builder,
const DataSegment* segment,
std::vector<Range>& ranges,
std::vector<std::unique_ptr<DataSegment>>& packed,
size_t segmentsRemaining) {
size_t segmentCount = 0;
bool hasExplicitName = false;
for (size_t i = 0; i < ranges.size(); ++i) {
Range& range = ranges[i];
if (range.isZero) {
continue;
}
Expression* offset = nullptr;
if (!segment->isPassive) {
if (auto* c = segment->offset->dynCast<Const>()) {
if (c->value.type == Type::i32) {
offset = builder.makeConst(int32_t(c->value.geti32() + range.start));
} else {
assert(c->value.type == Type::i64);
offset = builder.makeConst(int64_t(c->value.geti64() + range.start));
}
} else {
assert(ranges.size() == 1);
offset = segment->offset;
}
}
if (WebLimitations::MaxDataSegments <= packed.size() + segmentsRemaining) {
auto lastNonzero = ranges.end() - 1;
if (lastNonzero->isZero) {
--lastNonzero;
}
range.end = lastNonzero->end;
ranges.erase(ranges.begin() + i + 1, lastNonzero + 1);
}
Name name;
if (segment->name.is()) {
if (!segmentCount) {
name = segment->name;
hasExplicitName = segment->hasExplicitName;
} else {
name = segment->name.toString() + "." + std::to_string(segmentCount);
}
segmentCount++;
}
auto curr = Builder::makeDataSegment(name,
segment->memory,
segment->isPassive,
offset,
segment->data.data() + range.start,
range.end - range.start);
curr->hasExplicitName = hasExplicitName;
packed.push_back(std::move(curr));
}
}
void MemoryPacking::createReplacements(Module* module,
const std::vector<Range>& ranges,
const std::vector<Name>& segments,
const Referrers& referrers,
Replacements& replacements) {
if (ranges.size() == 1 && !ranges.front().isZero) {
return;
}
Builder builder(*module);
Name dropStateGlobal;
auto getDropStateGlobal = [&]() {
if (dropStateGlobal != Name()) {
return dropStateGlobal;
}
dropStateGlobal =
Names::getValidGlobalName(*module, "__mem_segment_drop_state");
module->addGlobal(builder.makeGlobal(dropStateGlobal,
Type::i32,
builder.makeConst(int32_t(0)),
Builder::Mutable));
return dropStateGlobal;
};
for (auto referrer : referrers) {
auto* init = referrer->dynCast<MemoryInit>();
if (init == nullptr) {
continue;
}
size_t start = init->offset->cast<Const>()->value.geti32();
size_t end = start + init->size->cast<Const>()->value.geti32();
size_t initIndex = 0;
size_t firstRangeIdx = 0;
while (firstRangeIdx < ranges.size() &&
ranges[firstRangeIdx].end <= start) {
if (!ranges[firstRangeIdx].isZero) {
++initIndex;
}
++firstRangeIdx;
}
if (start == end) {
Expression* result = builder.makeIf(
builder.makeBinary(
OrInt32,
makeGtShiftedMemorySize(builder, *module, init),
builder.makeGlobalGet(getDropStateGlobal(), Type::i32)),
builder.makeUnreachable());
replacements[init] = [result](Function*) { return result; };
continue;
}
assert(firstRangeIdx < ranges.size());
Expression* result = nullptr;
auto appendResult = [&](Expression* expr) {
result = result ? builder.blockify(result, expr) : expr;
};
Index* setVar = nullptr;
std::vector<Index*> getVars;
if (!init->dest->is<Const>()) {
auto set = builder.makeLocalSet(-1, init->dest);
setVar = &set->index;
appendResult(set);
}
if (ranges[firstRangeIdx].isZero) {
appendResult(
builder.makeIf(builder.makeGlobalGet(getDropStateGlobal(), Type::i32),
builder.makeUnreachable()));
}
size_t bytesWritten = 0;
for (size_t i = firstRangeIdx; i < ranges.size() && ranges[i].start < end;
++i) {
auto& range = ranges[i];
Expression* dest;
Type ptrType = module->getMemory(init->memory)->indexType;
if (auto* c = init->dest->dynCast<Const>()) {
dest =
builder.makeConstPtr(c->value.getInteger() + bytesWritten, ptrType);
} else {
auto* get = builder.makeLocalGet(-1, Type::i32);
getVars.push_back(&get->index);
dest = get;
if (bytesWritten > 0) {
Const* addend = builder.makeConst(int32_t(bytesWritten));
dest = builder.makeBinary(AddInt32, dest, addend);
}
}
size_t bytes = std::min(range.end, end) - std::max(range.start, start);
bytesWritten += bytes;
if (range.isZero) {
Expression* value = builder.makeConst(Literal::makeZero(Type::i32));
Expression* size = builder.makeConstPtr(bytes, ptrType);
appendResult(builder.makeMemoryFill(dest, value, size, init->memory));
} else {
size_t offsetBytes = std::max(start, range.start) - range.start;
Expression* offset = builder.makeConst(int32_t(offsetBytes));
Expression* size = builder.makeConst(int32_t(bytes));
appendResult(builder.makeMemoryInit(
segments[initIndex], dest, offset, size, init->memory));
initIndex++;
}
}
assert(result);
replacements[init] = [module, setVar, getVars, result](Function* function) {
if (setVar != nullptr) {
Index destVar = Builder(*module).addVar(function, Type::i32);
*setVar = destVar;
for (auto* getVar : getVars) {
*getVar = destVar;
}
}
return result;
};
}
for (auto drop : referrers) {
if (!drop->is<DataDrop>()) {
continue;
}
Expression* result = nullptr;
auto appendResult = [&](Expression* expr) {
result = result ? builder.blockify(result, expr) : expr;
};
if (dropStateGlobal != Name()) {
appendResult(
builder.makeGlobalSet(dropStateGlobal, builder.makeConst(int32_t(1))));
}
size_t dropIndex = 0;
for (auto range : ranges) {
if (!range.isZero) {
appendResult(builder.makeDataDrop(segments[dropIndex++]));
}
}
replacements[drop] = [result, module](Function*) {
return result ? result : Builder(*module).makeNop();
};
}
}
void MemoryPacking::replaceSegmentOps(Module* module,
Replacements& replacements) {
struct Replacer : WalkerPass<PostWalker<Replacer>> {
bool isFunctionParallel() override { return true; }
bool requiresNonNullableLocalFixups() override { return false; }
Replacements& replacements;
Replacer(Replacements& replacements) : replacements(replacements){};
std::unique_ptr<Pass> create() override {
return std::make_unique<Replacer>(replacements);
}
void visitMemoryInit(MemoryInit* curr) {
if (auto replacement = replacements.find(curr);
replacement != replacements.end()) {
replaceCurrent(replacement->second(getFunction()));
}
}
void visitDataDrop(DataDrop* curr) {
if (auto replacement = replacements.find(curr);
replacement != replacements.end()) {
replaceCurrent(replacement->second(getFunction()));
}
}
void visitArrayNewData(ArrayNewData* curr) {
if (auto replacement = replacements.find(curr);
replacement != replacements.end()) {
replaceCurrent(replacement->second(getFunction()));
}
}
} replacer(replacements);
replacer.run(getPassRunner(), module);
}
Pass* createMemoryPackingPass() { return new MemoryPacking(); }
}