1//===- DXILWriterPass.cpp - Bitcode writing pass --------------------------===//
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// DXILWriterPass implementation.
10//
11//===----------------------------------------------------------------------===//
12
13#include "DXILWriterPass.h"
14#include "DXILBitcodeWriter.h"
15#include "MCTargetDesc/DirectXContainerObjectWriter.h"
16#include "llvm/ADT/STLExtras.h"
17#include "llvm/ADT/StringRef.h"
18#include "llvm/Analysis/ModuleSummaryAnalysis.h"
19#include "llvm/IR/Constants.h"
20#include "llvm/IR/DebugInfo.h"
21#include "llvm/IR/DerivedTypes.h"
22#include "llvm/IR/GlobalVariable.h"
23#include "llvm/IR/IntrinsicInst.h"
24#include "llvm/IR/Intrinsics.h"
25#include "llvm/IR/LLVMContext.h"
26#include "llvm/IR/Module.h"
27#include "llvm/IR/PassManager.h"
28#include "llvm/InitializePasses.h"
29#include "llvm/Pass.h"
30#include "llvm/Support/Alignment.h"
31#include "llvm/Support/CommandLine.h"
32#include "llvm/Transforms/Utils/Cloning.h"
33#include "llvm/Transforms/Utils/ModuleUtils.h"
34
35using namespace llvm;
36using namespace llvm::dxil;
37
38cl::opt<bool> dxil::SourceInDebugModule(
39 "dx-source-in-debug-module",
40 cl::desc("Embed source code into debug module on DirectX target"),
41 cl::init(Val: false));
42
43namespace {
44class WriteDXILPass : public llvm::ModulePass {
45 raw_ostream &OS; // raw_ostream to print on
46
47public:
48 static char ID; // Pass identification, replacement for typeid
49 WriteDXILPass() : ModulePass(ID), OS(dbgs()) {
50 initializeWriteDXILPassPass(*PassRegistry::getPassRegistry());
51 }
52
53 explicit WriteDXILPass(raw_ostream &o) : ModulePass(ID), OS(o) {
54 initializeWriteDXILPassPass(*PassRegistry::getPassRegistry());
55 }
56
57 StringRef getPassName() const override { return "Bitcode Writer"; }
58
59 bool runOnModule(Module &M) override {
60 WriteDXILToFile(M, Out&: OS);
61 return false;
62 }
63 void getAnalysisUsage(AnalysisUsage &AU) const override {
64 AU.setPreservesAll();
65 }
66};
67
68static void legalizeLifetimeIntrinsics(Module &M) {
69 LLVMContext &Ctx = M.getContext();
70 Type *I64Ty = IntegerType::get(C&: Ctx, NumBits: 64);
71 Type *PtrTy = PointerType::get(C&: Ctx, AddressSpace: 0);
72 Intrinsic::ID LifetimeIIDs[2] = {Intrinsic::lifetime_start,
73 Intrinsic::lifetime_end};
74 for (Intrinsic::ID &IID : LifetimeIIDs) {
75 Function *F = M.getFunction(Name: Intrinsic::getName(Id: IID, OverloadTys: {PtrTy}, M: &M));
76 if (!F)
77 continue;
78
79 // Get or insert an LLVM 3.7-compliant lifetime intrinsic function of the
80 // form `void @llvm.lifetime.[start/end](i64, ptr)` with the NoUnwind
81 // attribute
82 AttributeList Attr;
83 Attr = Attr.addFnAttribute(C&: Ctx, Kind: Attribute::NoUnwind);
84 FunctionCallee LifetimeCallee = M.getOrInsertFunction(
85 Name: Intrinsic::getBaseName(id: IID), AttributeList: Attr, RetTy: Type::getVoidTy(C&: Ctx), Args: I64Ty, Args: PtrTy);
86
87 // Replace all calls to lifetime intrinsics with calls to the
88 // LLVM 3.7-compliant version of the lifetime intrinsic
89 for (User *U : make_early_inc_range(Range: F->users())) {
90 CallInst *CI = dyn_cast<CallInst>(Val: U);
91 assert(CI &&
92 "Expected user of a lifetime intrinsic function to be a CallInst");
93
94 // LLVM 3.7 lifetime intrinics require an i8* operand, so we insert
95 // a bitcast to ensure that is the case
96 Value *PtrOperand = CI->getArgOperand(i: 0);
97 PointerType *PtrOpPtrTy = cast<PointerType>(Val: PtrOperand->getType());
98 Value *NoOpBitCast = CastInst::Create(Instruction::BitCast, S: PtrOperand,
99 Ty: PtrOpPtrTy, Name: "", InsertBefore: CI->getIterator());
100
101 // LLVM 3.7 lifetime intrinsics have an explicit size operand, whose value
102 // we can obtain from the pointer operand which must be an AllocaInst (as
103 // of https://github.com/llvm/llvm-project/pull/149310)
104 AllocaInst *AI = dyn_cast<AllocaInst>(Val: PtrOperand);
105 assert(AI &&
106 "The pointer operand of a lifetime intrinsic call must be an "
107 "AllocaInst");
108 std::optional<TypeSize> AllocSize =
109 AI->getAllocationSize(DL: CI->getDataLayout());
110 assert(AllocSize.has_value() &&
111 "Expected the allocation size of AllocaInst to be known");
112 CallInst *NewCI = CallInst::Create(
113 Func: LifetimeCallee,
114 Args: {ConstantInt::get(Ty: I64Ty, V: AllocSize.value().getFixedValue()),
115 NoOpBitCast},
116 NameStr: "", InsertBefore: CI->getIterator());
117 for (Attribute ParamAttr : CI->getParamAttributes(ArgNo: 0))
118 NewCI->addParamAttr(ArgNo: 1, Attr: ParamAttr);
119
120 CI->eraseFromParent();
121 }
122
123 F->eraseFromParent();
124 }
125}
126
127static void removeLifetimeIntrinsics(Module &M) {
128 Intrinsic::ID LifetimeIIDs[2] = {Intrinsic::lifetime_start,
129 Intrinsic::lifetime_end};
130 for (Intrinsic::ID &IID : LifetimeIIDs) {
131 Function *F = M.getFunction(Name: Intrinsic::getBaseName(id: IID));
132 if (!F)
133 continue;
134
135 for (User *U : make_early_inc_range(Range: F->users())) {
136 CallInst *CI = dyn_cast<CallInst>(Val: U);
137 assert(CI && "Expected user of lifetime function to be a CallInst");
138 BitCastInst *BCI = dyn_cast<BitCastInst>(Val: CI->getArgOperand(i: 1));
139 assert(BCI && "Expected pointer operand of CallInst to be a BitCastInst");
140 CI->eraseFromParent();
141 BCI->eraseFromParent();
142 }
143 F->eraseFromParent();
144 }
145}
146
147static void replaceNamedMetadataArray(Module &M, StringRef Name,
148 ArrayRef<Metadata *> NewOps) {
149 NamedMDNode *NMD = M.getNamedMetadata(Name);
150 if (!NMD)
151 return;
152 NMD->eraseFromParent();
153 M.getOrInsertNamedMetadata(Name)->addOperand(
154 M: MDTuple::get(Context&: M.getContext(), MDs: NewOps));
155}
156
157class EmbedDXILPass : public llvm::ModulePass {
158 std::string writeModule(Module &M, bool HasDebugInfo, bool WriteDebug) {
159 std::string Data;
160 llvm::raw_string_ostream OS(Data);
161
162 if (HasDebugInfo) {
163 if (WriteDebug) {
164 if (!SourceInDebugModule) {
165 // Replace dx.source metadata nodes with stubs.
166 LLVMContext &Ctx = M.getContext();
167 MDString *EmptyString = MDString::get(Context&: Ctx, Str: "");
168 replaceNamedMetadataArray(M, Name: "dx.source.contents",
169 NewOps: {EmptyString, EmptyString});
170 replaceNamedMetadataArray(M, Name: "dx.source.defines", NewOps: {});
171 replaceNamedMetadataArray(M, Name: "dx.source.mainFileName", NewOps: {EmptyString});
172 replaceNamedMetadataArray(M, Name: "dx.source.args", NewOps: {});
173 }
174 } else {
175 // If we have an ILDB part, strip DXIL from all debug info.
176 StripDebugInfo(M);
177
178 // Also, manually remove debug version flags and dx.source nodes.
179 if (NamedMDNode *Flags = M.getModuleFlagsMetadata()) {
180 SmallVector<llvm::Module::ModuleFlagEntry, 4> FlagEntries;
181 M.getModuleFlagsMetadata(Flags&: FlagEntries);
182 Flags->eraseFromParent();
183 for (llvm::Module::ModuleFlagEntry &Entry : FlagEntries) {
184 if (Entry.Key->getString() == "Dwarf Version" ||
185 Entry.Key->getString() == "Debug Info Version") {
186 continue;
187 }
188 M.addModuleFlag(Behavior: Entry.Behavior, Key: Entry.Key->getString(), Val: Entry.Val);
189 }
190 }
191 for (NamedMDNode &NMD : llvm::make_early_inc_range(Range: M.named_metadata()))
192 if (NMD.getName().starts_with(Prefix: "dx.source"))
193 NMD.eraseFromParent();
194 }
195 } else {
196#ifdef EXPENSIVE_CHECKS
197 assert(
198 StripDebugInfo(M) == false &&
199 "The module must not contain any debug info here."
200 "Shader modules with debug info must have !DICompileUnit metadata.");
201#endif
202 }
203 WriteDXILToFile(M, Out&: OS);
204 return Data;
205 }
206
207 GlobalVariable *createSectionGlobal(Module &M, StringRef Data,
208 StringRef GlobalName,
209 StringRef SectionName) {
210 Constant *ModuleConstant =
211 ConstantDataArray::get(Context&: M.getContext(), Elts: arrayRefFromStringRef(Input: Data));
212 auto *GV = new llvm::GlobalVariable(M, ModuleConstant->getType(), true,
213 GlobalValue::PrivateLinkage,
214 ModuleConstant, GlobalName);
215 GV->setSection(SectionName);
216 GV->setAlignment(Align(4));
217 return GV;
218 }
219
220public:
221 static char ID; // Pass identification, replacement for typeid
222 EmbedDXILPass() : ModulePass(ID) {
223 initializeEmbedDXILPassPass(*PassRegistry::getPassRegistry());
224 }
225
226 StringRef getPassName() const override { return "DXIL Embedder"; }
227
228 bool runOnModule(Module &M) override {
229 // Perform late legalization of lifetime intrinsics that would otherwise
230 // fail the Module Verifier if performed in an earlier pass
231 legalizeLifetimeIntrinsics(M);
232
233 bool HasDebugInfo = !M.debug_compile_units().empty();
234
235 if (SlimDebug && EmbedDebug)
236 reportFatalUsageError(reason: "/Qembed_debug is not compatible with /Zs");
237 if (!HasDebugInfo && EmbedDebug)
238 reportFatalUsageError(
239 reason: "Missing debug info for embedding into the container");
240 if (!HasDebugInfo && !PdbDebugPath.empty())
241 reportFatalUsageError(reason: "Missing debug info for writing to the PDB file");
242
243 std::string ILDBData;
244 if (HasDebugInfo) {
245 // Write DXIL with debug info to ILDB part.
246 // Clone the module to avoid alternating it with DebugInfoPass
247 // before stripping the debug info later.
248 ILDBData =
249 writeModule(M&: *llvm::CloneModule(M), HasDebugInfo, /*WriteDebug=*/true);
250 }
251
252 // Clone the module to save dx.source metadata nodes from stripping, as they
253 // are needed for DXILMetadataAnalysisWrapperPass.
254 std::string DXILData =
255 writeModule(M&: *llvm::CloneModule(M), HasDebugInfo, /*WriteDebug=*/false);
256
257 // We no longer need lifetime intrinsics after bitcode serialization, so we
258 // simply remove them to keep the Module Verifier happy after our
259 // not-so-legal legalizations
260 removeLifetimeIntrinsics(M);
261
262 SmallVector<GlobalValue *, 2> Globals;
263 if (HasDebugInfo) {
264 // Create a GV after both parts are written, otherwise it gets
265 // added to DXIL when `writeModule` is called the second time.
266 Globals.emplace_back(Args: createSectionGlobal(M, Data: ILDBData, GlobalName: "dx.ildb", SectionName: "ILDB"));
267 }
268 Globals.emplace_back(Args: createSectionGlobal(M, Data: DXILData, GlobalName: "dx.dxil", SectionName: "DXIL"));
269 appendToCompilerUsed(M, Values: Globals);
270 return true;
271 }
272
273 void getAnalysisUsage(AnalysisUsage &AU) const override {
274 AU.setPreservesAll();
275 }
276};
277} // namespace
278
279char WriteDXILPass::ID = 0;
280INITIALIZE_PASS_BEGIN(WriteDXILPass, "dxil-write-bitcode", "Write Bitcode",
281 false, true)
282INITIALIZE_PASS_DEPENDENCY(ModuleSummaryIndexWrapperPass)
283INITIALIZE_PASS_END(WriteDXILPass, "dxil-write-bitcode", "Write Bitcode", false,
284 true)
285
286ModulePass *llvm::createDXILWriterPass(raw_ostream &Str) {
287 return new WriteDXILPass(Str);
288}
289
290char EmbedDXILPass::ID = 0;
291INITIALIZE_PASS(EmbedDXILPass, "dxil-embed", "Embed DXIL", false, true)
292
293ModulePass *llvm::createDXILEmbedderPass() { return new EmbedDXILPass(); }
294