1//===- JumpTableToSwitch.cpp ----------------------------------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8
9#include "llvm/Transforms/Scalar/JumpTableToSwitch.h"
10#include "ScalarOptions.h"
11#include "llvm/ADT/STLExtras.h"
12#include "llvm/ADT/SmallVector.h"
13#include "llvm/ADT/Statistic.h"
14#include "llvm/Analysis/ConstantFolding.h"
15#include "llvm/Analysis/DomTreeUpdater.h"
16#include "llvm/Analysis/OptimizationRemarkEmitter.h"
17#include "llvm/Analysis/PostDominators.h"
18#include "llvm/IR/IRBuilder.h"
19#include "llvm/IR/LLVMContext.h"
20#include "llvm/IR/ProfDataUtils.h"
21#include "llvm/ProfileData/InstrProf.h"
22#include "llvm/Transforms/Utils/BasicBlockUtils.h"
23#include <limits>
24
25using namespace llvm;
26
27#define DEBUG_TYPE "jump-table-to-switch"
28
29STATISTIC(NumEligibleJumpTables, "The number of jump tables seen by the pass "
30 "that can be converted if deemed profitable.");
31STATISTIC(NumJumpTablesConverted,
32 "The number of jump tables converted into switches.");
33
34namespace {
35struct JumpTableTy {
36 Value *Index;
37 SmallVector<Function *, 10> Funcs;
38};
39} // anonymous namespace
40
41static std::optional<JumpTableTy> parseJumpTable(GetElementPtrInst *GEP,
42 PointerType *PtrTy,
43 FunctionType *CallFTy) {
44 const ScalarOptions &Opts = ScalarOptions::Global;
45 Constant *Ptr = dyn_cast<Constant>(Val: GEP->getPointerOperand());
46 if (!Ptr)
47 return std::nullopt;
48
49 GlobalVariable *GV = dyn_cast<GlobalVariable>(Val: Ptr);
50 if (!GV || !GV->isConstant() || !GV->hasDefinitiveInitializer())
51 return std::nullopt;
52
53 Function &F = *GEP->getParent()->getParent();
54 const DataLayout &DL = F.getDataLayout();
55 const unsigned BitWidth =
56 DL.getIndexSizeInBits(AS: GEP->getPointerAddressSpace());
57 SmallMapVector<Value *, APInt, 4> VariableOffsets;
58 APInt ConstantOffset(BitWidth, 0);
59 if (!GEP->collectOffset(DL, BitWidth, VariableOffsets, ConstantOffset))
60 return std::nullopt;
61 if (VariableOffsets.size() != 1)
62 return std::nullopt;
63 // TODO: consider supporting more general patterns
64 if (!ConstantOffset.isZero())
65 return std::nullopt;
66 APInt StrideBytes = VariableOffsets.front().second;
67 const uint64_t JumpTableSizeBytes = GV->getGlobalSize(DL);
68 if (JumpTableSizeBytes % StrideBytes.getZExtValue() != 0)
69 return std::nullopt;
70 ++NumEligibleJumpTables;
71 const uint64_t N = JumpTableSizeBytes / StrideBytes.getZExtValue();
72 if (N > Opts.jump_table_to_switch_size_threshold)
73 return std::nullopt;
74
75 JumpTableTy JumpTable;
76 JumpTable.Index = VariableOffsets.front().first;
77 JumpTable.Funcs.reserve(N);
78 for (uint64_t Index = 0; Index < N; ++Index) {
79 // ConstantOffset is zero.
80 APInt Offset = Index * StrideBytes;
81 Constant *C =
82 ConstantFoldLoadFromConst(C: GV->getInitializer(), Ty: PtrTy, Offset, DL);
83 auto *Func = dyn_cast_or_null<Function>(Val: C);
84 if (!Func || Func->isDeclaration() || Func->getFunctionType() != CallFTy ||
85 Func->getInstructionCount() >
86 Opts.jump_table_to_switch_function_size_threshold)
87 return std::nullopt;
88 JumpTable.Funcs.push_back(Elt: Func);
89 }
90 return JumpTable;
91}
92
93static BasicBlock *
94expandToSwitch(CallBase *CB, const JumpTableTy &JT, DomTreeUpdater &DTU,
95 OptimizationRemarkEmitter &ORE,
96 llvm::function_ref<GlobalValue::GUID(const Function &)>
97 GetGuidForFunction) {
98 ++NumJumpTablesConverted;
99 const bool IsVoid = CB->getType() == Type::getVoidTy(C&: CB->getContext());
100
101 SmallVector<DominatorTree::UpdateType, 8> DTUpdates;
102 BasicBlock *BB = CB->getParent();
103 BasicBlock *Tail = SplitBlock(Old: BB, SplitPt: CB, DTU: &DTU, LI: nullptr, MSSAU: nullptr,
104 BBName: BB->getName() + Twine(".tail"));
105 DTUpdates.push_back(Elt: {DominatorTree::Delete, BB, Tail});
106 BB->getTerminator()->eraseFromParent();
107
108 Function &F = *BB->getParent();
109 BasicBlock *BBUnreachable = BasicBlock::Create(
110 Context&: F.getContext(), Name: "default.switch.case.unreachable", Parent: &F, InsertBefore: Tail);
111 IRBuilder<> BuilderUnreachable(BBUnreachable);
112 BuilderUnreachable.CreateUnreachable();
113
114 IRBuilder<> Builder(BB);
115 SwitchInst *Switch = Builder.CreateSwitch(V: JT.Index, Dest: BBUnreachable);
116 DTUpdates.push_back(Elt: {DominatorTree::Insert, BB, BBUnreachable});
117
118 IRBuilder<> BuilderTail(CB);
119 PHINode *PHI =
120 IsVoid ? nullptr : BuilderTail.CreatePHI(Ty: CB->getType(), NumReservedValues: JT.Funcs.size());
121 const auto *ProfMD = CB->getMetadata(KindID: LLVMContext::MD_prof);
122
123 SmallVector<uint64_t> BranchWeights;
124 DenseMap<GlobalValue::GUID, uint64_t> GuidToCounter;
125 const bool HadProfile = isValueProfileMD(ProfileData: ProfMD);
126 if (HadProfile) {
127 // The assumptions, coming in, are that the functions in JT.Funcs are
128 // defined in this module (from parseJumpTable).
129 assert(llvm::all_of(
130 JT.Funcs, [](const Function *F) { return F && !F->isDeclaration(); }));
131 BranchWeights.reserve(N: JT.Funcs.size() + 1);
132 // The first is the default target, which is the unreachable block created
133 // above.
134 BranchWeights.push_back(Elt: 0U);
135 uint64_t TotalCount = 0;
136 auto Targets = getValueProfDataFromInst(
137 Inst: *CB, ValueKind: InstrProfValueKind::IPVK_IndirectCallTarget,
138 MaxNumValueData: std::numeric_limits<uint32_t>::max(), TotalC&: TotalCount);
139
140 for (const auto &[G, C] : Targets) {
141 [[maybe_unused]] auto It = GuidToCounter.insert(KV: {G, C});
142 // We should always be inserting as it is verifier-enforced IR invariant
143 // that VP metadata does not have duplicate values.
144 assert(It.second);
145 }
146 }
147 for (auto [Index, Func] : llvm::enumerate(First: JT.Funcs)) {
148 BasicBlock *B = BasicBlock::Create(Context&: Func->getContext(),
149 Name: "call." + Twine(Index), Parent: &F, InsertBefore: Tail);
150 DTUpdates.push_back(Elt: {DominatorTree::Insert, BB, B});
151 DTUpdates.push_back(Elt: {DominatorTree::Insert, B, Tail});
152
153 CallBase *Call = cast<CallBase>(Val: CB->clone());
154 // The MD_prof metadata (VP kind), if it existed, can be dropped, it doesn't
155 // make sense on a direct call. Note that the values are used for the branch
156 // weights of the switch.
157 Call->setMetadata(KindID: LLVMContext::MD_prof, Node: nullptr);
158 Call->setCalledFunction(Func);
159 Call->insertInto(ParentBB: B, It: B->end());
160 Switch->addCase(
161 OnVal: cast<ConstantInt>(Val: ConstantInt::get(Ty: JT.Index->getType(), V: Index)), Dest: B);
162 GlobalValue::GUID FctID = GetGuidForFunction(*Func);
163 // It'd be OK to _not_ find target functions in GuidToCounter, e.g. suppose
164 // just some of the jump targets are taken (for the given profile).
165 BranchWeights.push_back(Elt: FctID == 0U ? 0U
166 : GuidToCounter.lookup_or(Val: FctID, Default: 0U));
167 UncondBrInst::Create(Target: Tail, InsertBefore: B);
168 if (PHI)
169 PHI->addIncoming(V: Call, BB: B);
170 }
171 DTU.applyUpdates(Updates: DTUpdates);
172 ORE.emit(RemarkBuilder: [&]() {
173 return OptimizationRemark(DEBUG_TYPE, "ReplacedJumpTableWithSwitch", CB)
174 << "expanded indirect call into switch";
175 });
176 // Only set branch weights on the switch if we have non-zero branch weights.
177 // We can have no non-zero branch weights while having VP metadata if for
178 // example, all of the functions are external and not instrumented.
179 if (HadProfile && llvm::any_of(Range&: BranchWeights, P: not_equal_to(Arg: 0))) {
180 setBranchWeights(I&: *Switch, Weights: downscaleWeights(Weights: BranchWeights),
181 /*IsExpected=*/false);
182 } else
183 setExplicitlyUnknownBranchWeights(I&: *Switch, DEBUG_TYPE);
184 if (PHI)
185 CB->replaceAllUsesWith(V: PHI);
186 CB->eraseFromParent();
187 return Tail;
188}
189
190PreservedAnalyses JumpTableToSwitchPass::run(Function &F,
191 FunctionAnalysisManager &AM) {
192 OptimizationRemarkEmitter &ORE =
193 AM.getResult<OptimizationRemarkEmitterAnalysis>(IR&: F);
194 DominatorTree *DT = AM.getCachedResult<DominatorTreeAnalysis>(IR&: F);
195 PostDominatorTree *PDT = AM.getCachedResult<PostDominatorTreeAnalysis>(IR&: F);
196 DomTreeUpdater DTU(DT, PDT, DomTreeUpdater::UpdateStrategy::Lazy);
197 bool Changed = false;
198 auto FuncToGuid = [&](const Function &Fct) {
199 if (const auto MaybeGUID = Fct.getGUIDIfAssigned(); MaybeGUID)
200 return *MaybeGUID;
201
202 return Function::getGUIDAssumingExternalLinkage(
203 GlobalName: getIRPGOObjectName(GO: Fct, InLTO));
204 };
205
206 for (BasicBlock &BB : make_early_inc_range(Range&: F)) {
207 BasicBlock *CurrentBB = &BB;
208 while (CurrentBB) {
209 BasicBlock *SplittedOutTail = nullptr;
210 for (Instruction &I : make_early_inc_range(Range&: *CurrentBB)) {
211 auto *Call = dyn_cast<CallInst>(Val: &I);
212 if (!Call || Call->getCalledFunction() || Call->isMustTailCall())
213 continue;
214 auto *L = dyn_cast<LoadInst>(Val: Call->getCalledOperand());
215 // Skip atomic or volatile loads.
216 if (!L || !L->isSimple())
217 continue;
218 auto *GEP = dyn_cast<GetElementPtrInst>(Val: L->getPointerOperand());
219 if (!GEP)
220 continue;
221 auto *PtrTy = dyn_cast<PointerType>(Val: L->getType());
222 assert(PtrTy && "call operand must be a pointer");
223 std::optional<JumpTableTy> JumpTable =
224 parseJumpTable(GEP, PtrTy, CallFTy: Call->getFunctionType());
225 if (!JumpTable)
226 continue;
227 SplittedOutTail =
228 expandToSwitch(CB: Call, JT: *JumpTable, DTU, ORE, GetGuidForFunction: FuncToGuid);
229 Changed = true;
230 break;
231 }
232 CurrentBB = SplittedOutTail ? SplittedOutTail : nullptr;
233 }
234 }
235
236 if (!Changed)
237 return PreservedAnalyses::all();
238
239 PreservedAnalyses PA;
240 if (DT)
241 PA.preserve<DominatorTreeAnalysis>();
242 if (PDT)
243 PA.preserve<PostDominatorTreeAnalysis>();
244 return PA;
245}
246