java-topology/defects/llvm/patch/llvm-0002-hotcoldsplitting-getOutliningPenalty-Region-O-R2.patch

47 lines
2 KiB
Diff
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# UNDF: UNDF-2026-000000153
# UNDF: (leave blank)
# CWE-407: HotColdSplitting getOutliningPenalty Region membership O(R²)
#
# In getOutliningPenalty, Region is an ArrayRef<BasicBlock*>. Two nested
# loops call is_contained(Region, ...) for successor/predecessor membership:
#
# 1) Line 322: for BB in Region → for SuccBB in successors(BB) →
# is_contained(Region, SuccBB) → O(R² × S)
#
# 2) Line 340: for ExitBB → for PHI → for incoming →
# is_contained(Region, getIncomingBlock) → O(E × P × V × R)
#
# Fix: build a SmallPtrSet<BasicBlock*> from Region once at function entry,
# then use O(1) set lookups instead of O(R) linear scans.
#
# Severity: MEDIUM — cold regions in large functions can have hundreds of
# blocks; HotColdSplitting runs as part of the default optimization pipeline.
#
--- a/llvm/lib/Transforms/IPO/HotColdSplitting.cpp
+++ b/llvm/lib/Transforms/IPO/HotColdSplitting.cpp
@@ -299,6 +299,9 @@
static int getOutliningPenalty(ArrayRef<BasicBlock *> Region,
unsigned NumInputs, unsigned NumOutputs) {
int Penalty = SplittingThreshold;
+ // Build a set for O(1) region membership tests (was O(R) linear scan).
+ SmallPtrSet<BasicBlock *, 16> RegionSet(Region.begin(), Region.end());
+
LLVM_DEBUG(dbgs() << "Applying penalty for splitting: " << Penalty << "\n");
// If the splitting threshold is set at or below zero, skip the usual
@@ -321,7 +324,7 @@
for (BasicBlock *SuccBB : successors(BB)) {
- if (!is_contained(Region, SuccBB)) {
+ if (!RegionSet.count(SuccBB)) {
NoBlocksReturn = false;
SuccsOutsideRegion.insert(SuccBB);
}
@@ -340,7 +343,7 @@
int NumIncomingVals = 0;
for (unsigned i = 0; i < PN.getNumIncomingValues(); ++i)
- if (llvm::is_contained(Region, PN.getIncomingBlock(i))) {
+ if (RegionSet.count(PN.getIncomingBlock(i))) {
++NumIncomingVals;
if (NumIncomingVals > 1) {
++NumSplitExitPhis;