1//===- DXILMemIntrinsics.cpp - Eliminate Memory Intrinsics ----------------===//
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 "DXILMemIntrinsics.h"
10#include "DirectX.h"
11#include "llvm/Analysis/DXILResource.h"
12#include "llvm/IR/IRBuilder.h"
13#include "llvm/IR/IntrinsicInst.h"
14#include "llvm/IR/IntrinsicsDirectX.h"
15#include "llvm/IR/Module.h"
16
17#define DEBUG_TYPE "dxil-mem-intrinsics"
18
19using namespace llvm;
20
21void expandMemSet(MemSetInst *MemSet) {
22 IRBuilder<> Builder(MemSet);
23 Value *Dst = MemSet->getDest();
24 Value *Val = MemSet->getValue();
25 ConstantInt *LengthCI = dyn_cast<ConstantInt>(Val: MemSet->getLength());
26 assert(LengthCI && "Expected length to be a ConstantInt");
27
28 [[maybe_unused]] const DataLayout &DL = Builder.getDataLayout();
29 [[maybe_unused]] uint64_t OrigLength = LengthCI->getZExtValue();
30
31 AllocaInst *Alloca = dyn_cast<AllocaInst>(Val: Dst);
32
33 assert(Alloca && "Expected memset on an Alloca");
34 assert(OrigLength == Alloca->getAllocationSize(DL)->getFixedValue() &&
35 "Expected for memset size to match DataLayout size");
36
37 Type *AllocatedTy = Alloca->getAllocatedType();
38 ArrayType *ArrTy = dyn_cast<ArrayType>(Val: AllocatedTy);
39 assert(ArrTy && "Expected Alloca for an Array Type");
40
41 Type *ElemTy = ArrTy->getElementType();
42 uint64_t Size = ArrTy->getArrayNumElements();
43
44 [[maybe_unused]] uint64_t ElemSize = DL.getTypeStoreSize(Ty: ElemTy);
45
46 assert(ElemSize > 0 && "Size must be set");
47 assert(OrigLength == ElemSize * Size && "Size in bytes must match");
48
49 Value *TypedVal = Val;
50
51 if (Val->getType() != ElemTy)
52 TypedVal = Builder.CreateIntCast(V: Val, DestTy: ElemTy, isSigned: false);
53
54 for (uint64_t I = 0; I < Size; ++I) {
55 Value *Zero = Builder.getInt32(C: 0);
56 Value *Offset = Builder.getInt32(C: I);
57 Value *Ptr = GetElementPtrInst::Create(PointeeType: ArrTy, Ptr: Dst, IdxList: {Zero, Offset}, NameStr: "gep",
58 InsertBefore: MemSet->getIterator());
59 Builder.CreateStore(Val: TypedVal, Ptr);
60 }
61
62 MemSet->eraseFromParent();
63}
64
65static Type *getPointeeType(Value *Ptr, const DataLayout &DL) {
66 if (auto *GV = dyn_cast<GlobalVariable>(Val: Ptr))
67 return GV->getValueType();
68 if (auto *AI = dyn_cast<AllocaInst>(Val: Ptr))
69 return AI->getAllocatedType();
70
71 if (auto *II = dyn_cast<IntrinsicInst>(Val: Ptr)) {
72 if (II->getIntrinsicID() == Intrinsic::dx_resource_getpointer) {
73 Type *Ty = cast<dxil::AnyResourceExtType>(Val: II->getArgOperand(i: 0)->getType())
74 ->getResourceType();
75 assert(Ty && "getpointer used on untyped resource");
76 return Ty;
77 }
78 }
79
80 if (auto *GEP = dyn_cast<GEPOperator>(Val: Ptr)) {
81 Type *Ty = GEP->getResultElementType();
82 if (!Ty->isIntegerTy(BitWidth: 8))
83 return Ty;
84
85 // We have ptradd, so we have to hope there's enough information to work out
86 // what we're indexing.
87 Type *IndexedType = getPointeeType(Ptr: GEP->getPointerOperand(), DL);
88 if (auto *AT = dyn_cast<ArrayType>(Val: IndexedType))
89 return AT->getElementType();
90
91 if (auto *ST = dyn_cast<StructType>(Val: IndexedType)) {
92 // Indexing a struct should always be constant
93 APInt ConstantOffset(DL.getIndexTypeSizeInBits(Ty: GEP->getType()), 0);
94 [[maybe_unused]] bool IsConst =
95 GEP->accumulateConstantOffset(DL, Offset&: ConstantOffset);
96 assert(IsConst && "Non-constant GEP into struct?");
97
98 // Now, work out what we'll find at that offset.
99 const StructLayout *Layout = DL.getStructLayout(Ty: ST);
100 unsigned Idx =
101 Layout->getElementContainingOffset(FixedOffset: ConstantOffset.getZExtValue());
102
103 return ST->getTypeAtIndex(N: Idx);
104 }
105
106 llvm_unreachable("Could not infer type from GEP");
107 }
108
109 llvm_unreachable("Could not calculate pointee type");
110}
111
112static size_t flattenTypes(Type *ContainerTy, const DataLayout &DL,
113 SmallVectorImpl<std::pair<Type *, size_t>> &FlatTys,
114 size_t NextOffset = 0) {
115 if (auto *AT = dyn_cast<ArrayType>(Val: ContainerTy)) {
116 for (uint64_t I = 0, E = AT->getNumElements(); I != E; ++I)
117 NextOffset = flattenTypes(ContainerTy: AT->getElementType(), DL, FlatTys, NextOffset);
118 return NextOffset;
119 }
120 if (auto *ST = dyn_cast<StructType>(Val: ContainerTy)) {
121 for (Type *Ty : ST->elements())
122 NextOffset = flattenTypes(ContainerTy: Ty, DL, FlatTys, NextOffset);
123 return NextOffset;
124 }
125
126 FlatTys.emplace_back(Args&: ContainerTy, Args&: NextOffset);
127 return NextOffset + DL.getTypeStoreSize(Ty: ContainerTy);
128}
129
130void expandMemCpy(MemCpyInst *MemCpy) {
131 IRBuilder<> Builder(MemCpy);
132 Value *Dst = MemCpy->getDest();
133 Value *Src = MemCpy->getSource();
134 ConstantInt *LengthCI = dyn_cast<ConstantInt>(Val: MemCpy->getLength());
135 assert(LengthCI && "Expected Length to be a ConstantInt");
136 assert(!MemCpy->isVolatile() && "Handling for volatile not implemented");
137
138 uint64_t ByteLength = LengthCI->getZExtValue();
139 // If length to copy is zero, no memcpy is needed.
140 if (ByteLength == 0)
141 return;
142
143 const DataLayout &DL = Builder.getDataLayout();
144
145 SmallVector<std::pair<Type *, size_t>> FlattenedTypes;
146 [[maybe_unused]] size_t MaxLength =
147 flattenTypes(ContainerTy: getPointeeType(Ptr: Dst, DL), DL, FlatTys&: FlattenedTypes);
148 assert(MaxLength >= ByteLength && "Dst not large enough for memcpy");
149
150 LLVM_DEBUG({
151 // Check if Src is layout compatible with Dst. This should always be true
152 // unless the frontend did something wrong.
153 SmallVector<std::pair<Type *, size_t>> SrcTypes;
154 size_t SrcLength = flattenTypes(getPointeeType(Src, DL), DL, SrcTypes);
155 assert(SrcLength >= ByteLength && "Src not large enough for memcpy");
156 for (const auto &[LHS, RHS] : zip(FlattenedTypes, SrcTypes)) {
157 auto &[DstTy, DstOffset] = LHS;
158 auto &[SrcTy, SrcOffset] = RHS;
159 assert(DstTy == SrcTy && "Mismatched types for memcpy");
160 assert(DstOffset == SrcOffset && "Incompatible layouts for memcpy");
161 if (DstOffset >= ByteLength)
162 break;
163 }
164 });
165
166 for (const auto &[Ty, Offset] : FlattenedTypes) {
167 if (Offset >= ByteLength)
168 break;
169 // TODO: Should we skip padding types here?
170 Value *ByteOffset = Builder.getInt32(C: Offset);
171 Value *SrcPtr = Builder.CreateInBoundsPtrAdd(Ptr: Src, Offset: ByteOffset);
172 Value *SrcVal = Builder.CreateLoad(Ty, Ptr: SrcPtr);
173 Value *DstPtr = Builder.CreateInBoundsPtrAdd(Ptr: Dst, Offset: ByteOffset);
174 Builder.CreateStore(Val: SrcVal, Ptr: DstPtr);
175 }
176
177 MemCpy->eraseFromParent();
178}
179
180void expandMemMove(MemMoveInst *MemMove) {
181 report_fatal_error(reason: "memmove expansion is not implemented yet.");
182}
183
184static bool eliminateMemIntrinsics(Module &M) {
185 bool HadMemIntrinsicUses = false;
186 for (auto &F : make_early_inc_range(Range: M.functions())) {
187 Intrinsic::ID IID = F.getIntrinsicID();
188 switch (IID) {
189 case Intrinsic::memcpy:
190 case Intrinsic::memcpy_inline:
191 case Intrinsic::memmove:
192 case Intrinsic::memset:
193 case Intrinsic::memset_inline:
194 break;
195 default:
196 continue;
197 }
198 for (User *U : make_early_inc_range(Range: F.users())) {
199 HadMemIntrinsicUses = true;
200 if (auto *MemSet = dyn_cast<MemSetInst>(Val: U))
201 expandMemSet(MemSet);
202 else if (auto *MemCpy = dyn_cast<MemCpyInst>(Val: U))
203 expandMemCpy(MemCpy);
204 else if (auto *MemMove = dyn_cast<MemMoveInst>(Val: U))
205 expandMemMove(MemMove);
206 else
207 llvm_unreachable("Unhandled memory intrinsic");
208 }
209 assert(F.user_empty() && "Mem intrinsic not eliminated?");
210 F.eraseFromParent();
211 }
212 return HadMemIntrinsicUses;
213}
214
215PreservedAnalyses DXILMemIntrinsics::run(Module &M, ModuleAnalysisManager &) {
216 if (eliminateMemIntrinsics(M))
217 return PreservedAnalyses::none();
218 return PreservedAnalyses::all();
219}
220
221class DXILMemIntrinsicsLegacy : public ModulePass {
222public:
223 bool runOnModule(Module &M) override { return eliminateMemIntrinsics(M); }
224 DXILMemIntrinsicsLegacy() : ModulePass(ID) {}
225
226 static char ID; // Pass identification.
227};
228char DXILMemIntrinsicsLegacy::ID = 0;
229
230INITIALIZE_PASS_BEGIN(DXILMemIntrinsicsLegacy, DEBUG_TYPE,
231 "DXIL Memory Intrinsic Elimination", false, false)
232INITIALIZE_PASS_END(DXILMemIntrinsicsLegacy, DEBUG_TYPE,
233 "DXIL Memory Intrinsic Elimination", false, false)
234
235ModulePass *llvm::createDXILMemIntrinsicsLegacyPass() {
236 return new DXILMemIntrinsicsLegacy();
237}
238