1//===- BlockExtractor.cpp - Extracts blocks into their own functions ------===//
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// This pass extracts the specified basic blocks from the module into their
10// own functions.
11//
12//===----------------------------------------------------------------------===//
13
14#include "llvm/Transforms/IPO/BlockExtractor.h"
15#include "llvm/ADT/STLExtras.h"
16#include "llvm/ADT/Statistic.h"
17#include "llvm/IR/Instructions.h"
18#include "llvm/IR/Module.h"
19#include "llvm/IR/PassManager.h"
20#include "llvm/Support/CommandLine.h"
21#include "llvm/Support/Debug.h"
22#include "llvm/Support/MemoryBuffer.h"
23#include "llvm/Transforms/IPO.h"
24#include "llvm/Transforms/Utils/BasicBlockUtils.h"
25#include "llvm/Transforms/Utils/CodeExtractor.h"
26
27using namespace llvm;
28
29#define DEBUG_TYPE "block-extractor"
30
31STATISTIC(NumExtracted, "Number of basic blocks extracted");
32
33static cl::opt<std::string> BlockExtractorFile(
34 "extract-blocks-file", cl::value_desc("filename"),
35 cl::desc("A file containing list of basic blocks to extract"), cl::Hidden);
36
37static cl::opt<bool>
38 BlockExtractorEraseFuncs("extract-blocks-erase-funcs",
39 cl::desc("Erase the existing functions"),
40 cl::Hidden);
41namespace {
42class BlockExtractor {
43public:
44 BlockExtractor(bool EraseFunctions) : EraseFunctions(EraseFunctions) {}
45 bool runOnModule(Module &M);
46 void
47 init(const std::vector<std::vector<BasicBlock *>> &GroupsOfBlocksToExtract) {
48 GroupsOfBlocks = GroupsOfBlocksToExtract;
49 if (!BlockExtractorFile.empty())
50 loadFile();
51 }
52
53private:
54 std::vector<std::vector<BasicBlock *>> GroupsOfBlocks;
55 bool EraseFunctions;
56 /// Map a function name to groups of blocks.
57 SmallVector<std::pair<std::string, SmallVector<std::string, 4>>, 4>
58 BlocksByName;
59
60 void loadFile();
61 bool splitLandingPadPreds(Function &F);
62};
63
64} // end anonymous namespace
65
66/// Gets all of the blocks specified in the input file.
67void BlockExtractor::loadFile() {
68 auto ErrOrBuf = MemoryBuffer::getFile(Filename: BlockExtractorFile);
69 if (ErrOrBuf.getError())
70 report_fatal_error(reason: "BlockExtractor couldn't load the file.");
71 // Read the file.
72 auto &Buf = *ErrOrBuf;
73 SmallVector<StringRef, 16> Lines;
74 Buf->getBuffer().split(A&: Lines, Separator: '\n', /*MaxSplit=*/-1,
75 /*KeepEmpty=*/false);
76 for (const auto &Line : Lines) {
77 SmallVector<StringRef, 4> LineSplit;
78 Line.split(A&: LineSplit, Separator: ' ', /*MaxSplit=*/-1,
79 /*KeepEmpty=*/false);
80 if (LineSplit.empty())
81 continue;
82 if (LineSplit.size()!=2)
83 reportFatalUsageError(
84 reason: "Invalid line format, expecting lines like: 'funcname bb1[;bb2..]'");
85 SmallVector<StringRef, 4> BBNames;
86 LineSplit[1].split(A&: BBNames, Separator: ';', /*MaxSplit=*/-1,
87 /*KeepEmpty=*/false);
88 if (BBNames.empty())
89 report_fatal_error(reason: "Missing bbs name");
90 BlocksByName.push_back(
91 Elt: {std::string(LineSplit[0]), {BBNames.begin(), BBNames.end()}});
92 }
93}
94
95/// Extracts the landing pads to make sure all of them have only one
96/// predecessor.
97bool BlockExtractor::splitLandingPadPreds(Function &F) {
98 bool Changed = false;
99 for (BasicBlock &BB : F) {
100 for (Instruction &I : BB) {
101 if (!isa<InvokeInst>(Val: &I))
102 continue;
103 InvokeInst *II = cast<InvokeInst>(Val: &I);
104 BasicBlock *Parent = II->getParent();
105 BasicBlock *LPad = II->getUnwindDest();
106
107 // Look through the landing pad's predecessors. If one of them ends in an
108 // 'invoke', then we want to split the landing pad.
109 bool Split = false;
110 for (auto *PredBB : predecessors(BB: LPad)) {
111 if (PredBB->isLandingPad() && PredBB != Parent &&
112 isa<InvokeInst>(Val: Parent->getTerminator())) {
113 Split = true;
114 break;
115 }
116 }
117
118 if (!Split)
119 continue;
120
121 SmallVector<BasicBlock *, 2> NewBBs;
122 SplitLandingPadPredecessors(OrigBB: LPad, Preds: Parent, Suffix: ".1", Suffix2: ".2", NewBBs);
123 Changed = true;
124 }
125 }
126 return Changed;
127}
128
129bool BlockExtractor::runOnModule(Module &M) {
130 bool Changed = false;
131
132 // Get all the functions.
133 SmallVector<Function *, 4> Functions;
134 for (Function &F : M) {
135 Changed |= splitLandingPadPreds(F);
136 Functions.push_back(Elt: &F);
137 }
138
139 // Get all the blocks specified in the input file.
140 unsigned NextGroupIdx = GroupsOfBlocks.size();
141 GroupsOfBlocks.resize(new_size: NextGroupIdx + BlocksByName.size());
142 for (const auto &BInfo : BlocksByName) {
143 Function *F = M.getFunction(Name: BInfo.first);
144 if (!F)
145 reportFatalUsageError(
146 reason: "Invalid function name specified in the input file");
147 for (const auto &BBInfo : BInfo.second) {
148 auto Res = llvm::find_if(
149 Range&: *F, P: [&](const BasicBlock &BB) { return BB.getName() == BBInfo; });
150 if (Res == F->end())
151 reportFatalUsageError(reason: "Invalid block name specified in the input file");
152 GroupsOfBlocks[NextGroupIdx].push_back(x: &*Res);
153 }
154 ++NextGroupIdx;
155 }
156
157 // Extract each group of basic blocks.
158 for (auto &BBs : GroupsOfBlocks) {
159 SmallVector<BasicBlock *, 32> BlocksToExtractVec;
160 for (BasicBlock *BB : BBs) {
161 // Check if the module contains BB.
162 if (BB->getParent()->getParent() != &M)
163 reportFatalUsageError(reason: "Invalid basic block");
164 LLVM_DEBUG(dbgs() << "BlockExtractor: Extracting "
165 << BB->getParent()->getName() << ":" << BB->getName()
166 << "\n");
167 BlocksToExtractVec.push_back(Elt: BB);
168 if (const InvokeInst *II = dyn_cast<InvokeInst>(Val: BB->getTerminator()))
169 BlocksToExtractVec.push_back(Elt: II->getUnwindDest());
170 ++NumExtracted;
171 Changed = true;
172 }
173 CodeExtractorAnalysisCache CEAC(*BBs[0]->getParent());
174 Function *F = CodeExtractor(BlocksToExtractVec).extractCodeRegion(CEAC);
175 if (F)
176 LLVM_DEBUG(dbgs() << "Extracted group '" << (*BBs.begin())->getName()
177 << "' in: " << F->getName() << '\n');
178 else
179 LLVM_DEBUG(dbgs() << "Failed to extract for group '"
180 << (*BBs.begin())->getName() << "'\n");
181 }
182
183 // Erase the functions.
184 if (EraseFunctions || BlockExtractorEraseFuncs) {
185 for (Function *F : Functions) {
186 LLVM_DEBUG(dbgs() << "BlockExtractor: Trying to delete " << F->getName()
187 << "\n");
188 F->deleteBody();
189 }
190 // Set linkage as ExternalLinkage to avoid erasing unreachable functions.
191 for (Function &F : M)
192 F.setLinkage(GlobalValue::ExternalLinkage);
193 Changed = true;
194 }
195
196 return Changed;
197}
198
199BlockExtractorPass::BlockExtractorPass(
200 std::vector<std::vector<BasicBlock *>> &&GroupsOfBlocks,
201 bool EraseFunctions)
202 : GroupsOfBlocks(std::move(GroupsOfBlocks)),
203 EraseFunctions(EraseFunctions) {}
204
205PreservedAnalyses BlockExtractorPass::run(Module &M,
206 ModuleAnalysisManager &AM) {
207 BlockExtractor BE(EraseFunctions);
208 BE.init(GroupsOfBlocksToExtract: GroupsOfBlocks);
209 return BE.runOnModule(M) ? PreservedAnalyses::none()
210 : PreservedAnalyses::all();
211}
212