1//===- llvm-extract.cpp - LLVM function extraction utility ----------------===//
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 utility changes the input module to only contain a single function,
10// which is primarily used for debugging transformations.
11//
12//===----------------------------------------------------------------------===//
13
14#include "llvm/ADT/SetVector.h"
15#include "llvm/ADT/SmallPtrSet.h"
16#include "llvm/Bitcode/BitcodeWriterPass.h"
17#include "llvm/IR/DataLayout.h"
18#include "llvm/IR/IRPrintingPasses.h"
19#include "llvm/IR/Instructions.h"
20#include "llvm/IR/LLVMContext.h"
21#include "llvm/IR/Module.h"
22#include "llvm/IR/Verifier.h"
23#include "llvm/IRPrinter/IRPrintingPasses.h"
24#include "llvm/IRReader/IRReader.h"
25#include "llvm/Passes/PassBuilder.h"
26#include "llvm/Support/CommandLine.h"
27#include "llvm/Support/Error.h"
28#include "llvm/Support/FileSystem.h"
29#include "llvm/Support/InitLLVM.h"
30#include "llvm/Support/Regex.h"
31#include "llvm/Support/SourceMgr.h"
32#include "llvm/Support/SystemUtils.h"
33#include "llvm/Support/ToolOutputFile.h"
34#include "llvm/Transforms/IPO.h"
35#include "llvm/Transforms/IPO/BlockExtractor.h"
36#include "llvm/Transforms/IPO/ExtractGV.h"
37#include "llvm/Transforms/IPO/GlobalDCE.h"
38#include "llvm/Transforms/IPO/StripDeadPrototypes.h"
39#include "llvm/Transforms/IPO/StripSymbols.h"
40#include <memory>
41#include <utility>
42
43using namespace llvm;
44
45static cl::OptionCategory ExtractCat("llvm-extract Options");
46
47// InputFilename - The filename to read from.
48static cl::opt<std::string> InputFilename(cl::Positional,
49 cl::desc("<input bitcode file>"),
50 cl::init(Val: "-"),
51 cl::value_desc("filename"));
52
53static cl::opt<std::string> OutputFilename("o",
54 cl::desc("Specify output filename"),
55 cl::value_desc("filename"),
56 cl::init(Val: "-"), cl::cat(ExtractCat));
57
58static cl::opt<bool> Force("f", cl::desc("Enable binary output on terminals"),
59 cl::cat(ExtractCat));
60
61static cl::opt<bool> NoVerify("disable-verify",
62 cl::desc("Do not run the verifier"),
63 cl::cat(ExtractCat), cl::Hidden);
64
65static cl::opt<bool> DeleteFn("delete",
66 cl::desc("Delete specified Globals from Module"),
67 cl::cat(ExtractCat));
68
69static cl::opt<bool> KeepConstInit("keep-const-init",
70 cl::desc("Keep initializers of constants"),
71 cl::cat(ExtractCat));
72
73static cl::opt<bool>
74 Recursive("recursive", cl::desc("Recursively extract all called functions"),
75 cl::cat(ExtractCat));
76
77// ExtractFuncs - The functions to extract from the module.
78static cl::list<std::string>
79 ExtractFuncs("func", cl::desc("Specify function to extract"),
80 cl::value_desc("function"), cl::cat(ExtractCat));
81
82// ExtractRegExpFuncs - The functions, matched via regular expression, to
83// extract from the module.
84static cl::list<std::string>
85 ExtractRegExpFuncs("rfunc",
86 cl::desc("Specify function(s) to extract using a "
87 "regular expression"),
88 cl::value_desc("rfunction"), cl::cat(ExtractCat));
89
90// ExtractBlocks - The blocks to extract from the module.
91static cl::list<std::string> ExtractBlocks(
92 "bb",
93 cl::desc(
94 "Specify <function, basic block1[;basic block2...]> pairs to extract.\n"
95 "Each pair will create a function.\n"
96 "If multiple basic blocks are specified in one pair,\n"
97 "the first block in the sequence should dominate the rest.\n"
98 "If an unnamed basic block is to be extracted,\n"
99 "'%' should be added before the basic block variable names.\n"
100 "eg:\n"
101 " --bb=f:bb1;bb2 will extract one function with both bb1 and bb2;\n"
102 " --bb=f:bb1 --bb=f:bb2 will extract two functions, one with bb1, one "
103 "with bb2.\n"
104 " --bb=f:%1 will extract one function with basic block 1;"),
105 cl::value_desc("function:bb1[;bb2...]"), cl::cat(ExtractCat));
106
107// ExtractAlias - The alias to extract from the module.
108static cl::list<std::string>
109 ExtractAliases("alias", cl::desc("Specify alias to extract"),
110 cl::value_desc("alias"), cl::cat(ExtractCat));
111
112// ExtractRegExpAliases - The aliases, matched via regular expression, to
113// extract from the module.
114static cl::list<std::string>
115 ExtractRegExpAliases("ralias",
116 cl::desc("Specify alias(es) to extract using a "
117 "regular expression"),
118 cl::value_desc("ralias"), cl::cat(ExtractCat));
119
120// ExtractGlobals - The globals to extract from the module.
121static cl::list<std::string>
122 ExtractGlobals("glob", cl::desc("Specify global to extract"),
123 cl::value_desc("global"), cl::cat(ExtractCat));
124
125// ExtractRegExpGlobals - The globals, matched via regular expression, to
126// extract from the module...
127static cl::list<std::string>
128 ExtractRegExpGlobals("rglob",
129 cl::desc("Specify global(s) to extract using a "
130 "regular expression"),
131 cl::value_desc("rglobal"), cl::cat(ExtractCat));
132
133static cl::opt<bool> OutputAssembly("S",
134 cl::desc("Write output as LLVM assembly"),
135 cl::Hidden, cl::cat(ExtractCat));
136
137int main(int argc, char **argv) {
138 InitLLVM X(argc, argv);
139
140 LLVMContext Context;
141 cl::HideUnrelatedOptions(Category&: ExtractCat);
142 cl::ParseCommandLineOptions(argc, argv, Overview: "llvm extractor\n");
143
144 // Use lazy loading, since we only care about selected global values.
145 SMDiagnostic Err;
146 std::unique_ptr<Module> M = getLazyIRFileModule(Filename: InputFilename, Err, Context);
147
148 if (!M) {
149 Err.print(ProgName: argv[0], S&: errs());
150 return 1;
151 }
152
153 if (!NoVerify && verifyModule(M: *M, OS: &errs())) {
154 errs() << argv[0] << ": " << InputFilename
155 << ": error: input module is broken!\n";
156 return 1;
157 }
158
159 // Use SetVector to avoid duplicates.
160 SetVector<GlobalValue *> GVs;
161
162 // Figure out which aliases we should extract.
163 for (size_t i = 0, e = ExtractAliases.size(); i != e; ++i) {
164 GlobalAlias *GA = M->getNamedAlias(Name: ExtractAliases[i]);
165 if (!GA) {
166 errs() << argv[0] << ": program doesn't contain alias named '"
167 << ExtractAliases[i] << "'!\n";
168 return 1;
169 }
170 GVs.insert(X: GA);
171 }
172
173 // Extract aliases via regular expression matching.
174 for (size_t i = 0, e = ExtractRegExpAliases.size(); i != e; ++i) {
175 std::string Error;
176 Regex RegEx(ExtractRegExpAliases[i]);
177 if (!RegEx.isValid(Error)) {
178 errs() << argv[0] << ": '" << ExtractRegExpAliases[i] << "' "
179 "invalid regex: " << Error;
180 }
181 bool match = false;
182 for (Module::alias_iterator GA = M->alias_begin(), E = M->alias_end();
183 GA != E; GA++) {
184 if (RegEx.match(String: GA->getName())) {
185 GVs.insert(X: &*GA);
186 match = true;
187 }
188 }
189 if (!match) {
190 errs() << argv[0] << ": program doesn't contain global named '"
191 << ExtractRegExpAliases[i] << "'!\n";
192 return 1;
193 }
194 }
195
196 // Figure out which globals we should extract.
197 for (size_t i = 0, e = ExtractGlobals.size(); i != e; ++i) {
198 GlobalValue *GV = M->getNamedGlobal(Name: ExtractGlobals[i]);
199 if (!GV) {
200 errs() << argv[0] << ": program doesn't contain global named '"
201 << ExtractGlobals[i] << "'!\n";
202 return 1;
203 }
204 GVs.insert(X: GV);
205 }
206
207 // Extract globals via regular expression matching.
208 for (size_t i = 0, e = ExtractRegExpGlobals.size(); i != e; ++i) {
209 std::string Error;
210 Regex RegEx(ExtractRegExpGlobals[i]);
211 if (!RegEx.isValid(Error)) {
212 errs() << argv[0] << ": '" << ExtractRegExpGlobals[i] << "' "
213 "invalid regex: " << Error;
214 }
215 bool match = false;
216 for (auto &GV : M->globals()) {
217 if (RegEx.match(String: GV.getName())) {
218 GVs.insert(X: &GV);
219 match = true;
220 }
221 }
222 if (!match) {
223 errs() << argv[0] << ": program doesn't contain global named '"
224 << ExtractRegExpGlobals[i] << "'!\n";
225 return 1;
226 }
227 }
228
229 // Figure out which functions we should extract.
230 for (size_t i = 0, e = ExtractFuncs.size(); i != e; ++i) {
231 GlobalValue *GV = M->getFunction(Name: ExtractFuncs[i]);
232 if (!GV) {
233 errs() << argv[0] << ": program doesn't contain function named '"
234 << ExtractFuncs[i] << "'!\n";
235 return 1;
236 }
237 GVs.insert(X: GV);
238 }
239 // Extract functions via regular expression matching.
240 for (size_t i = 0, e = ExtractRegExpFuncs.size(); i != e; ++i) {
241 std::string Error;
242 StringRef RegExStr = ExtractRegExpFuncs[i];
243 Regex RegEx(RegExStr);
244 if (!RegEx.isValid(Error)) {
245 errs() << argv[0] << ": '" << ExtractRegExpFuncs[i] << "' "
246 "invalid regex: " << Error;
247 }
248 bool match = false;
249 for (Module::iterator F = M->begin(), E = M->end(); F != E;
250 F++) {
251 if (RegEx.match(String: F->getName())) {
252 GVs.insert(X: &*F);
253 match = true;
254 }
255 }
256 if (!match) {
257 errs() << argv[0] << ": program doesn't contain global named '"
258 << ExtractRegExpFuncs[i] << "'!\n";
259 return 1;
260 }
261 }
262
263 // Figure out which BasicBlocks we should extract.
264 SmallVector<std::pair<Function *, SmallVector<StringRef, 16>>, 2> BBMap;
265 for (StringRef StrPair : ExtractBlocks) {
266 SmallVector<StringRef, 16> BBNames;
267 auto BBInfo = StrPair.split(Separator: ':');
268 // Get the function.
269 Function *F = M->getFunction(Name: BBInfo.first);
270 if (!F) {
271 errs() << argv[0] << ": program doesn't contain a function named '"
272 << BBInfo.first << "'!\n";
273 return 1;
274 }
275 // Add the function to the materialize list, and store the basic block names
276 // to check after materialization.
277 GVs.insert(X: F);
278 BBInfo.second.split(A&: BBNames, Separator: ';', /*MaxSplit=*/-1, /*KeepEmpty=*/false);
279 BBMap.push_back(Elt: {F, std::move(BBNames)});
280 }
281
282 // Use *argv instead of argv[0] to work around a wrong GCC warning.
283 ExitOnError ExitOnErr(std::string(*argv) + ": error reading input: ");
284
285 if (Recursive) {
286 std::vector<llvm::Function *> Workqueue;
287 for (GlobalValue *GV : GVs) {
288 if (auto *F = dyn_cast<Function>(Val: GV)) {
289 Workqueue.push_back(x: F);
290 }
291 }
292 while (!Workqueue.empty()) {
293 Function *F = &*Workqueue.back();
294 Workqueue.pop_back();
295 ExitOnErr(F->materialize());
296 for (auto &BB : *F) {
297 for (auto &I : BB) {
298 CallBase *CB = dyn_cast<CallBase>(Val: &I);
299 if (!CB)
300 continue;
301 Function *CF = CB->getCalledFunction();
302 if (!CF)
303 continue;
304 if (CF->isDeclaration() || !GVs.insert(X: CF))
305 continue;
306 Workqueue.push_back(x: CF);
307 }
308 }
309 }
310 }
311
312 auto Materialize = [&](GlobalValue &GV) { ExitOnErr(GV.materialize()); };
313
314 // Materialize requisite global values.
315 if (!DeleteFn) {
316 for (size_t i = 0, e = GVs.size(); i != e; ++i)
317 Materialize(*GVs[i]);
318 } else {
319 // Deleting. Materialize every GV that's *not* in GVs.
320 SmallPtrSet<GlobalValue *, 8> GVSet(llvm::from_range, GVs);
321 for (auto &F : *M) {
322 if (!GVSet.count(Ptr: &F))
323 Materialize(F);
324 }
325 }
326
327 {
328 std::vector<GlobalValue *> Gvs(GVs.begin(), GVs.end());
329 LoopAnalysisManager LAM;
330 FunctionAnalysisManager FAM;
331 CGSCCAnalysisManager CGAM;
332 ModuleAnalysisManager MAM;
333
334 PassBuilder PB;
335
336 PB.registerModuleAnalyses(MAM);
337 PB.registerCGSCCAnalyses(CGAM);
338 PB.registerFunctionAnalyses(FAM);
339 PB.registerLoopAnalyses(LAM);
340 PB.crossRegisterProxies(LAM, FAM, CGAM, MAM);
341
342 ModulePassManager PM;
343 PM.addPass(Pass: ExtractGVPass(Gvs, DeleteFn, KeepConstInit));
344 PM.run(IR&: *M, AM&: MAM);
345
346 // Now that we have all the GVs we want, mark the module as fully
347 // materialized.
348 // FIXME: should the GVExtractionPass handle this?
349 ExitOnErr(M->materializeAll());
350 }
351
352 // Extract the specified basic blocks from the module and erase the existing
353 // functions.
354 if (!ExtractBlocks.empty()) {
355 // Figure out which BasicBlocks we should extract.
356 std::vector<std::vector<BasicBlock *>> GroupOfBBs;
357 for (auto &P : BBMap) {
358 std::vector<BasicBlock *> BBs;
359 for (StringRef BBName : P.second) {
360 // The function has been materialized, so add its matching basic blocks
361 // to the block extractor list, or fail if a name is not found.
362 auto Res = llvm::find_if(Range&: *P.first, P: [&](const BasicBlock &BB) {
363 return BB.getNameOrAsOperand() == BBName;
364 });
365 if (Res == P.first->end()) {
366 errs() << argv[0] << ": function " << P.first->getName()
367 << " doesn't contain a basic block named '" << BBName
368 << "'!\n";
369 return 1;
370 }
371 BBs.push_back(x: &*Res);
372 }
373 GroupOfBBs.push_back(x: BBs);
374 }
375
376 LoopAnalysisManager LAM;
377 FunctionAnalysisManager FAM;
378 CGSCCAnalysisManager CGAM;
379 ModuleAnalysisManager MAM;
380
381 PassBuilder PB;
382
383 PB.registerModuleAnalyses(MAM);
384 PB.registerCGSCCAnalyses(CGAM);
385 PB.registerFunctionAnalyses(FAM);
386 PB.registerLoopAnalyses(LAM);
387 PB.crossRegisterProxies(LAM, FAM, CGAM, MAM);
388
389 ModulePassManager PM;
390 PM.addPass(Pass: BlockExtractorPass(std::move(GroupOfBBs), true));
391 PM.run(IR&: *M, AM&: MAM);
392 }
393
394 // In addition to deleting all other functions, we also want to spiff it
395 // up a little bit. Do this now.
396
397 LoopAnalysisManager LAM;
398 FunctionAnalysisManager FAM;
399 CGSCCAnalysisManager CGAM;
400 ModuleAnalysisManager MAM;
401
402 PassBuilder PB;
403
404 PB.registerModuleAnalyses(MAM);
405 PB.registerCGSCCAnalyses(CGAM);
406 PB.registerFunctionAnalyses(FAM);
407 PB.registerLoopAnalyses(LAM);
408 PB.crossRegisterProxies(LAM, FAM, CGAM, MAM);
409
410 ModulePassManager PM;
411 if (!DeleteFn)
412 PM.addPass(Pass: GlobalDCEPass());
413 PM.addPass(Pass: StripDeadDebugInfoPass());
414 PM.addPass(Pass: StripDeadPrototypesPass());
415 PM.addPass(Pass: StripDeadCGProfilePass());
416
417 std::error_code EC;
418 ToolOutputFile Out(OutputFilename, EC, sys::fs::OF_None);
419 if (EC) {
420 errs() << EC.message() << '\n';
421 return 1;
422 }
423
424 if (OutputAssembly)
425 PM.addPass(Pass: PrintModulePass(Out.os(), "",
426 /*ShouldPreserveUseListOrder=*/false,
427 /*EmitSummaryIndex=*/false,
428 /*ShouldRenumberMetadata=*/true));
429 else if (Force || !CheckBitcodeOutputToConsole(stream_to_check&: Out.os()))
430 PM.addPass(
431 Pass: BitcodeWriterPass(Out.os(), /* ShouldPreserveUseListOrder */ true));
432
433 PM.run(IR&: *M, AM&: MAM);
434
435 // Declare success.
436 Out.keep();
437
438 return 0;
439}
440