1//===--- ExpandMemCmp.cpp - Expand memcmp() to load/stores ----------------===//
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 tries to expand memcmp() calls into optimally-sized loads and
10// compares for the target.
11//
12//===----------------------------------------------------------------------===//
13
14#include "llvm/Transforms/Scalar/ExpandMemCmp.h"
15#include "ScalarOptions.h"
16#include "llvm/ADT/Statistic.h"
17#include "llvm/Analysis/ConstantFolding.h"
18#include "llvm/Analysis/DomTreeUpdater.h"
19#include "llvm/Analysis/LazyBlockFrequencyInfo.h"
20#include "llvm/Analysis/ProfileSummaryInfo.h"
21#include "llvm/Analysis/TargetLibraryInfo.h"
22#include "llvm/Analysis/TargetTransformInfo.h"
23#include "llvm/Analysis/ValueTracking.h"
24#include "llvm/IR/Dominators.h"
25#include "llvm/IR/IRBuilder.h"
26#include "llvm/IR/InstIterator.h"
27#include "llvm/IR/PatternMatch.h"
28#include "llvm/IR/ProfDataUtils.h"
29#include "llvm/Transforms/Utils/BasicBlockUtils.h"
30#include "llvm/Transforms/Utils/Local.h"
31#include "llvm/Transforms/Utils/SizeOpts.h"
32#include <optional>
33
34using namespace llvm;
35using namespace llvm::PatternMatch;
36
37#define DEBUG_TYPE "expand-memcmp"
38
39STATISTIC(NumMemCmpCalls, "Number of memcmp calls");
40STATISTIC(NumMemCmpNotConstant, "Number of memcmp calls without constant size");
41STATISTIC(NumMemCmpGreaterThanMax,
42 "Number of memcmp calls with size greater than max size");
43STATISTIC(NumMemCmpInlined, "Number of inlined memcmp calls");
44
45namespace {
46
47// Return the known alignment of the pointer argument \p ArgNo of \p CI,
48// combining the alignment of the underlying pointer value with any align
49// attribute on the call site itself.
50static Align getMemCmpArgAlignment(const CallInst *CI, unsigned ArgNo,
51 const DataLayout &DL) {
52 Align A = CI->getArgOperand(i: ArgNo)->getPointerAlignment(DL);
53 if (MaybeAlign ParamAlign = CI->getParamAlign(ArgNo))
54 A = std::max(a: A, b: *ParamAlign);
55 return A;
56}
57
58// This class provides helper functions to expand a memcmp library call into an
59// inline expansion.
60class MemCmpExpansion {
61 struct LoadPair {
62 Value *Lhs = nullptr;
63 Value *Rhs = nullptr;
64 };
65
66 struct ResultBlock {
67 BasicBlock *BB = nullptr;
68 PHINode *PhiSrc1 = nullptr;
69 PHINode *PhiSrc2 = nullptr;
70
71 ResultBlock() = default;
72 };
73
74 CallInst *const CI = nullptr;
75 ResultBlock ResBlock;
76 const uint64_t Size;
77 unsigned MaxLoadSize = 0;
78 const uint64_t NumLoadsPerBlock;
79 const unsigned MaxBytesPerBlock;
80 unsigned MaxBlockSize = 0;
81 std::vector<BasicBlock *> LoadCmpBlocks;
82 BasicBlock *EndBlock = nullptr;
83 PHINode *PhiRes = nullptr;
84 const bool IsUsedForZeroCmp;
85 const DataLayout &DL;
86 const TargetTransformInfo &TTI;
87 // The known common alignment of the two source pointers.
88 const Align CommonAlign;
89 DomTreeUpdater *DTU = nullptr;
90 IRBuilder<> Builder;
91 // Represents the decomposition in blocks of the expansion. For example,
92 // comparing 33 bytes on X86+sse can be done with 2x16-byte loads and
93 // 1x1-byte load, which would be represented as [{16, 0}, {16, 16}, {1, 32}.
94 struct LoadEntry {
95 LoadEntry(unsigned LoadSize, uint64_t Offset)
96 : LoadSize(LoadSize), Offset(Offset) {
97 }
98
99 // The size of the load for this block, in bytes.
100 unsigned LoadSize;
101 // The offset of this load from the base pointer, in bytes.
102 uint64_t Offset;
103 };
104 using LoadEntryVector = SmallVector<LoadEntry, 8>;
105 LoadEntryVector LoadSequence;
106
107 void createLoadCmpBlocks();
108 void createResultBlock();
109 void setupResultBlockPHINodes();
110 void setupEndBlockPHINodes();
111 Value *getCompareLoadPairs(unsigned BlockIndex, unsigned &LoadIndex);
112 LoadPair getPackedLoadPair(unsigned BlockIndex, unsigned &LoadIndex);
113 void emitLoadCompareBlock(unsigned BlockIndex, unsigned &LoadIndex);
114 void emitLoadCompareBlockMultipleLoads(unsigned BlockIndex,
115 unsigned &LoadIndex);
116 void emitLoadCompareByteBlock(unsigned BlockIndex, unsigned OffsetBytes);
117 void emitMemCmpResultBlock();
118 Value *getMemCmpExpansionZeroCase();
119 Value *getMemCmpEqZeroOneBlock();
120 Value *getMemCmpOneBlock();
121 Value *getMemCmpOneBlockMultipleLoads();
122 Value *getMemCmpResult(const LoadPair &Loads);
123 LoadPair getLoadPair(Type *LoadSizeType, Type *BSwapSizeType,
124 Type *CmpSizeType, unsigned OffsetBytes);
125
126 // Return true if a load of `LoadSize` bytes at `Offset` from the base
127 // pointers is accessible on the target: either it is naturally aligned given
128 // the known common base alignment, or the target allows a misaligned access
129 // of that width.
130 bool isAccessAllowed(unsigned LoadSize, uint64_t Offset) const;
131
132 static LoadEntryVector
133 computeGreedyLoadSequence(uint64_t Size, llvm::ArrayRef<unsigned> LoadSizes,
134 unsigned MaxNumLoads);
135 LoadEntryVector computeOverlappingLoadSequence(uint64_t Size,
136 unsigned MaxLoadSize,
137 unsigned MaxNumLoads) const;
138
139 void optimiseLoadSequence(
140 LoadEntryVector &LoadSequence,
141 const TargetTransformInfo::MemCmpExpansionOptions &Options,
142 bool IsUsedForZeroCmp) const;
143
144public:
145 MemCmpExpansion(CallInst *CI, uint64_t Size,
146 const TargetTransformInfo::MemCmpExpansionOptions &Options,
147 const bool IsUsedForZeroCmp, const DataLayout &TheDataLayout,
148 DomTreeUpdater *DTU, const TargetTransformInfo &TTI,
149 Align CommonAlign, unsigned MaxBytesPerBlock);
150
151 unsigned getNumBlocks();
152 unsigned getNumLoadsInBlock(unsigned LoadIndex) const;
153 unsigned getNumBytesInBlock(unsigned LoadIndex, unsigned NumLoads) const;
154 uint64_t getNumLoads() const { return LoadSequence.size(); }
155
156 Value *getMemCmpExpansion();
157};
158
159// Return true if a load of `LoadSize` bytes at `Offset` from the base pointers
160// is accessible on the target: either it is naturally aligned given the known
161// common base alignment, or the target allows a misaligned access of that
162// width. We query whether the access is *allowed*, not whether it is *fast*,
163// matching the historical behavior of forming unaligned loads whenever the
164// target permits them.
165static bool isAccessAllowed(const CallInst *CI, const TargetTransformInfo &TTI,
166 Align CommonAlign, unsigned LoadSize,
167 uint64_t Offset) {
168 // The access is naturally aligned when the known alignment is at least the
169 // load width. LoadSize is not necessarily a power of two here: some targets
170 // like RISC-V add non-power-of-two load sizes for vector memcmp, so compare
171 // against the raw width rather than constructing an Align, which would
172 // require a power of two.
173 Align AccessAlign = commonAlignment(A: CommonAlign, Offset);
174 if (AccessAlign.value() >= LoadSize)
175 return true;
176 unsigned AS = CI->getArgOperand(i: 0)->getType()->getPointerAddressSpace();
177 return TTI.allowsMisalignedMemoryAccesses(Context&: CI->getContext(), BitWidth: LoadSize * 8, AddressSpace: AS,
178 Alignment: AccessAlign);
179}
180
181// Return true if a load of `LoadSize` bytes at `Offset` from the base pointers
182// is accessible on the target given the known common base alignment. This gates
183// the (power-of-two) overlapping loads; tail expansions are always legalized by
184// the backend and skip this check.
185bool MemCmpExpansion::isAccessAllowed(unsigned LoadSize,
186 uint64_t Offset) const {
187 return ::isAccessAllowed(CI, TTI, CommonAlign, LoadSize, Offset);
188}
189
190MemCmpExpansion::LoadEntryVector
191MemCmpExpansion::computeGreedyLoadSequence(uint64_t Size,
192 llvm::ArrayRef<unsigned> LoadSizes,
193 const unsigned MaxNumLoads) {
194 LoadEntryVector LoadSequence;
195 uint64_t Offset = 0;
196 while (Size && !LoadSizes.empty()) {
197 const unsigned LoadSize = LoadSizes.front();
198 const uint64_t NumLoadsForThisSize = Size / LoadSize;
199 if (LoadSequence.size() + NumLoadsForThisSize > MaxNumLoads) {
200 // Do not expand if the total number of loads is larger than what the
201 // target allows. Note that it's important that we exit before completing
202 // the expansion to avoid using a ton of memory to store the expansion for
203 // large sizes.
204 return {};
205 }
206 if (NumLoadsForThisSize > 0) {
207 for (uint64_t I = 0; I < NumLoadsForThisSize; ++I) {
208 LoadSequence.push_back(Elt: {LoadSize, Offset});
209 Offset += LoadSize;
210 }
211 Size = Size % LoadSize;
212 }
213 LoadSizes = LoadSizes.drop_front();
214 }
215 return LoadSequence;
216}
217
218MemCmpExpansion::LoadEntryVector
219MemCmpExpansion::computeOverlappingLoadSequence(
220 uint64_t Size, const unsigned MaxLoadSize,
221 const unsigned MaxNumLoads) const {
222 // These are already handled by the greedy approach.
223 if (Size < 2 || MaxLoadSize < 2)
224 return {};
225
226 // We try to do as many non-overlapping loads as possible starting from the
227 // beginning.
228 const uint64_t NumNonOverlappingLoads = Size / MaxLoadSize;
229 assert(NumNonOverlappingLoads && "there must be at least one load");
230 // There remain 0 to (MaxLoadSize - 1) bytes to load, this will be done with
231 // an overlapping load.
232 Size = Size - NumNonOverlappingLoads * MaxLoadSize;
233 // Bail if we do not need an overloapping store, this is already handled by
234 // the greedy approach.
235 if (Size == 0)
236 return {};
237 // Bail if the number of loads (non-overlapping + potential overlapping one)
238 // is larger than the max allowed.
239 if ((NumNonOverlappingLoads + 1) > MaxNumLoads)
240 return {};
241
242 // Add non-overlapping loads.
243 LoadEntryVector LoadSequence;
244 uint64_t Offset = 0;
245 for (uint64_t I = 0; I < NumNonOverlappingLoads; ++I) {
246 LoadSequence.push_back(Elt: {MaxLoadSize, Offset});
247 Offset += MaxLoadSize;
248 }
249
250 // Add the last overlapping load. Its offset is not a multiple of the load
251 // size, so it may be misaligned; bail if the target cannot access it.
252 assert(Size > 0 && Size < MaxLoadSize && "broken invariant");
253 uint64_t OverlapOffset = Offset - (MaxLoadSize - Size);
254 if (!isAccessAllowed(LoadSize: MaxLoadSize, Offset: OverlapOffset))
255 return {};
256
257 LoadSequence.push_back(Elt: {MaxLoadSize, OverlapOffset});
258 return LoadSequence;
259}
260
261void MemCmpExpansion::optimiseLoadSequence(
262 LoadEntryVector &LoadSequence,
263 const TargetTransformInfo::MemCmpExpansionOptions &Options,
264 bool IsUsedForZeroCmp) const {
265 // This part of code attempts to optimize the LoadSequence by merging allowed
266 // subsequences into single loads of allowed sizes from
267 // `MemCmpExpansionOptions::AllowedTailExpansions`. If it is for zero
268 // comparison or if no allowed tail expansions are specified, we exit early.
269 if (IsUsedForZeroCmp || Options.AllowedTailExpansions.empty())
270 return;
271
272 while (LoadSequence.size() >= 2) {
273 auto Last = LoadSequence[LoadSequence.size() - 1];
274 auto PreLast = LoadSequence[LoadSequence.size() - 2];
275
276 // Exit the loop if the two sequences are not contiguous
277 if (PreLast.Offset + PreLast.LoadSize != Last.Offset)
278 break;
279
280 auto LoadSize = Last.LoadSize + PreLast.LoadSize;
281 if (find(Range: Options.AllowedTailExpansions, Val: LoadSize) ==
282 Options.AllowedTailExpansions.end())
283 break;
284
285 // A merged load wider than MaxLoadSize can only be emitted when it is the
286 // sole load (getMemCmpOneBlock); in a multi-block expansion
287 // emitLoadCompareBlock requires every load to fit in MaxLoadSize (the
288 // result-block phis are sized to it). The per-call-site alignment filter
289 // can shrink MaxLoadSize, so stop merging when the result would still be
290 // multi-block and the merged load exceeds it.
291 if (LoadSize > MaxLoadSize && LoadSequence.size() > 2)
292 break;
293
294 // Remove the last two sequences and replace with the combined sequence
295 LoadSequence.pop_back();
296 LoadSequence.pop_back();
297 LoadSequence.emplace_back(Args&: LoadSize, Args&: PreLast.Offset);
298 }
299}
300
301// Initialize the basic block structure required for expansion of memcmp call
302// with given maximum load size and memcmp size parameter.
303// This structure includes:
304// 1. A list of load compare blocks - LoadCmpBlocks.
305// 2. An EndBlock, split from original instruction point, which is the block to
306// return from.
307// 3. ResultBlock, block to branch to for early exit when a
308// LoadCmpBlock finds a difference.
309MemCmpExpansion::MemCmpExpansion(
310 CallInst *const CI, uint64_t Size,
311 const TargetTransformInfo::MemCmpExpansionOptions &Options,
312 const bool IsUsedForZeroCmp, const DataLayout &TheDataLayout,
313 DomTreeUpdater *DTU, const TargetTransformInfo &TTI, Align CommonAlign,
314 unsigned MaxBytesPerBlock)
315 : CI(CI), Size(Size), NumLoadsPerBlock(Options.NumLoadsPerBlock),
316 MaxBytesPerBlock(MaxBytesPerBlock), IsUsedForZeroCmp(IsUsedForZeroCmp),
317 DL(TheDataLayout), TTI(TTI), CommonAlign(CommonAlign), DTU(DTU),
318 Builder(CI) {
319 assert(Size > 0 && "zero blocks");
320 assert(NumLoadsPerBlock > 0 && "zero loads per block");
321 // Scale the max size down if the target can load more bytes than we need.
322 llvm::ArrayRef<unsigned> LoadSizes(Options.LoadSizes);
323 while (!LoadSizes.empty() && LoadSizes.front() > Size) {
324 LoadSizes = LoadSizes.drop_front();
325 }
326 assert(!LoadSizes.empty() && "cannot load Size bytes");
327 MaxLoadSize = LoadSizes.front();
328 // Compute the decomposition.
329 LoadSequence =
330 computeGreedyLoadSequence(Size, LoadSizes, MaxNumLoads: Options.MaxNumLoads);
331 assert(LoadSequence.size() <= Options.MaxNumLoads && "broken invariant");
332 // If we allow overlapping loads and the load sequence is not already optimal,
333 // use overlapping loads.
334 if (Options.AllowOverlappingLoads &&
335 (LoadSequence.empty() || LoadSequence.size() > 2)) {
336 auto OverlappingLoads =
337 computeOverlappingLoadSequence(Size, MaxLoadSize, MaxNumLoads: Options.MaxNumLoads);
338 if (!OverlappingLoads.empty() &&
339 (LoadSequence.empty() ||
340 OverlappingLoads.size() < LoadSequence.size())) {
341 LoadSequence = OverlappingLoads;
342 }
343 }
344 assert(LoadSequence.size() <= Options.MaxNumLoads && "broken invariant");
345 optimiseLoadSequence(LoadSequence, Options, IsUsedForZeroCmp);
346
347 unsigned LoadIndex = 0;
348 while (LoadIndex < getNumLoads()) {
349 unsigned NumLoads = getNumLoadsInBlock(LoadIndex);
350 MaxBlockSize =
351 std::max(a: MaxBlockSize, b: getNumBytesInBlock(LoadIndex, NumLoads));
352 LoadIndex += NumLoads;
353 }
354 if (MaxBlockSize)
355 MaxBlockSize = PowerOf2Ceil(A: MaxBlockSize);
356}
357
358unsigned MemCmpExpansion::getNumLoadsInBlock(unsigned LoadIndex) const {
359 if (IsUsedForZeroCmp)
360 return std::min<uint64_t>(a: getNumLoads() - LoadIndex, b: NumLoadsPerBlock);
361
362 unsigned NumLoads = 0;
363 unsigned NumBytes = 0;
364 while (LoadIndex + NumLoads < getNumLoads() && NumLoads < NumLoadsPerBlock &&
365 NumBytes + LoadSequence[LoadIndex + NumLoads].LoadSize <=
366 MaxBytesPerBlock) {
367 NumBytes += LoadSequence[LoadIndex + NumLoads].LoadSize;
368 ++NumLoads;
369 }
370 assert(NumLoads && "at least one load must fit in a block");
371 return NumLoads;
372}
373
374unsigned MemCmpExpansion::getNumBytesInBlock(unsigned LoadIndex,
375 unsigned NumLoads) const {
376 unsigned NumBytes = 0;
377 for (unsigned I = 0; I != NumLoads; ++I)
378 NumBytes += LoadSequence[LoadIndex + I].LoadSize;
379 return NumBytes;
380}
381
382unsigned MemCmpExpansion::getNumBlocks() {
383 unsigned NumBlocks = 0;
384 for (unsigned LoadIndex = 0; LoadIndex < getNumLoads(); ++NumBlocks)
385 LoadIndex += getNumLoadsInBlock(LoadIndex);
386 return NumBlocks;
387}
388
389void MemCmpExpansion::createLoadCmpBlocks() {
390 for (unsigned i = 0; i < getNumBlocks(); i++) {
391 BasicBlock *BB = BasicBlock::Create(Context&: CI->getContext(), Name: "loadbb",
392 Parent: EndBlock->getParent(), InsertBefore: EndBlock);
393 LoadCmpBlocks.push_back(x: BB);
394 }
395}
396
397void MemCmpExpansion::createResultBlock() {
398 ResBlock.BB = BasicBlock::Create(Context&: CI->getContext(), Name: "res_block",
399 Parent: EndBlock->getParent(), InsertBefore: EndBlock);
400}
401
402MemCmpExpansion::LoadPair MemCmpExpansion::getLoadPair(Type *LoadSizeType,
403 Type *BSwapSizeType,
404 Type *CmpSizeType,
405 unsigned OffsetBytes) {
406 // Get the memory source at offset `OffsetBytes`.
407 Value *LhsSource = CI->getArgOperand(i: 0);
408 Value *RhsSource = CI->getArgOperand(i: 1);
409 Align LhsAlign = getMemCmpArgAlignment(CI, ArgNo: 0, DL);
410 Align RhsAlign = getMemCmpArgAlignment(CI, ArgNo: 1, DL);
411 if (OffsetBytes > 0) {
412 auto *ByteType = Type::getInt8Ty(C&: CI->getContext());
413 LhsSource = Builder.CreateConstGEP1_64(Ty: ByteType, Ptr: LhsSource, Idx0: OffsetBytes);
414 RhsSource = Builder.CreateConstGEP1_64(Ty: ByteType, Ptr: RhsSource, Idx0: OffsetBytes);
415 LhsAlign = commonAlignment(A: LhsAlign, Offset: OffsetBytes);
416 RhsAlign = commonAlignment(A: RhsAlign, Offset: OffsetBytes);
417 }
418
419 // Create a constant or a load from the source.
420 Value *Lhs = nullptr;
421 if (auto *C = dyn_cast<Constant>(Val: LhsSource))
422 Lhs = ConstantFoldLoadFromConstPtr(C, Ty: LoadSizeType, DL);
423 if (!Lhs)
424 Lhs = Builder.CreateAlignedLoad(Ty: LoadSizeType, Ptr: LhsSource, Align: LhsAlign);
425
426 Value *Rhs = nullptr;
427 if (auto *C = dyn_cast<Constant>(Val: RhsSource))
428 Rhs = ConstantFoldLoadFromConstPtr(C, Ty: LoadSizeType, DL);
429 if (!Rhs)
430 Rhs = Builder.CreateAlignedLoad(Ty: LoadSizeType, Ptr: RhsSource, Align: RhsAlign);
431
432 // Zero extend if Byte Swap intrinsic has different type
433 if (BSwapSizeType && LoadSizeType != BSwapSizeType) {
434 Lhs = Builder.CreateZExt(V: Lhs, DestTy: BSwapSizeType);
435 Rhs = Builder.CreateZExt(V: Rhs, DestTy: BSwapSizeType);
436 }
437
438 // Swap bytes if required.
439 if (BSwapSizeType) {
440 Function *Bswap = Intrinsic::getOrInsertDeclaration(
441 M: CI->getModule(), id: Intrinsic::bswap, OverloadTys: BSwapSizeType);
442 Lhs = Builder.CreateCall(Callee: Bswap, Args: Lhs);
443 Rhs = Builder.CreateCall(Callee: Bswap, Args: Rhs);
444 }
445
446 // Zero extend if required.
447 if (CmpSizeType != nullptr && CmpSizeType != Lhs->getType()) {
448 Lhs = Builder.CreateZExt(V: Lhs, DestTy: CmpSizeType);
449 Rhs = Builder.CreateZExt(V: Rhs, DestTy: CmpSizeType);
450 }
451 return {.Lhs: Lhs, .Rhs: Rhs};
452}
453
454// This function creates the IR instructions for loading and comparing 1 byte.
455// It loads 1 byte from each source of the memcmp parameters with the given
456// GEPIndex. It then subtracts the two loaded values and adds this result to the
457// final phi node for selecting the memcmp result.
458void MemCmpExpansion::emitLoadCompareByteBlock(unsigned BlockIndex,
459 unsigned OffsetBytes) {
460 BasicBlock *BB = LoadCmpBlocks[BlockIndex];
461 Builder.SetInsertPoint(BB);
462 const LoadPair Loads =
463 getLoadPair(LoadSizeType: Type::getInt8Ty(C&: CI->getContext()), BSwapSizeType: nullptr,
464 CmpSizeType: Type::getInt32Ty(C&: CI->getContext()), OffsetBytes);
465 Value *Diff = Builder.CreateSub(LHS: Loads.Lhs, RHS: Loads.Rhs);
466
467 PhiRes->addIncoming(V: Diff, BB);
468
469 if (BlockIndex < (LoadCmpBlocks.size() - 1)) {
470 // Early exit branch if difference found to EndBlock. Otherwise, continue to
471 // next LoadCmpBlock,
472 Value *Cmp = Builder.CreateICmp(P: ICmpInst::ICMP_NE, LHS: Diff,
473 RHS: ConstantInt::get(Ty: Diff->getType(), V: 0));
474 Builder.CreateCondBr(Cond: Cmp, True: EndBlock, False: LoadCmpBlocks[BlockIndex + 1]);
475 if (DTU)
476 DTU->applyUpdates(
477 Updates: {{DominatorTree::Insert, BB, EndBlock},
478 {DominatorTree::Insert, BB, LoadCmpBlocks[BlockIndex + 1]}});
479 } else {
480 // The last block has an unconditional branch to EndBlock.
481 Builder.CreateBr(Dest: EndBlock);
482 if (DTU)
483 DTU->applyUpdates(Updates: {{DominatorTree::Insert, BB, EndBlock}});
484 }
485}
486
487/// Generate an equality comparison for one or more pairs of loaded values.
488/// This is used in the case where the memcmp() call is compared equal or not
489/// equal to zero.
490Value *MemCmpExpansion::getCompareLoadPairs(unsigned BlockIndex,
491 unsigned &LoadIndex) {
492 assert(LoadIndex < getNumLoads() &&
493 "getCompareLoadPairs() called with no remaining loads");
494 std::vector<Value *> XorList, OrList;
495 Value *Diff = nullptr;
496
497 const unsigned NumLoads = getNumLoadsInBlock(LoadIndex);
498
499 // For a single-block expansion, start inserting before the memcmp call.
500 if (LoadCmpBlocks.empty())
501 Builder.SetInsertPoint(CI);
502 else
503 Builder.SetInsertPoint(LoadCmpBlocks[BlockIndex]);
504
505 Value *Cmp = nullptr;
506 // If we have multiple loads per block, we need to generate a composite
507 // comparison using xor+or. The type for the combinations is the largest load
508 // type.
509 IntegerType *const MaxLoadType =
510 NumLoads == 1 ? nullptr
511 : IntegerType::get(C&: CI->getContext(), NumBits: MaxLoadSize * 8);
512
513 for (unsigned i = 0; i < NumLoads; ++i, ++LoadIndex) {
514 const LoadEntry &CurLoadEntry = LoadSequence[LoadIndex];
515 const LoadPair Loads = getLoadPair(
516 LoadSizeType: IntegerType::get(C&: CI->getContext(), NumBits: CurLoadEntry.LoadSize * 8), BSwapSizeType: nullptr,
517 CmpSizeType: MaxLoadType, OffsetBytes: CurLoadEntry.Offset);
518
519 if (NumLoads != 1) {
520 // If we have multiple loads per block, we need to generate a composite
521 // comparison using xor+or.
522 Diff = Builder.CreateXor(LHS: Loads.Lhs, RHS: Loads.Rhs);
523 Diff = Builder.CreateZExt(V: Diff, DestTy: MaxLoadType);
524 XorList.push_back(x: Diff);
525 } else {
526 // If there's only one load per block, we just compare the loaded values.
527 Cmp = Builder.CreateICmpNE(LHS: Loads.Lhs, RHS: Loads.Rhs);
528 }
529 }
530
531 auto pairWiseOr = [&](std::vector<Value *> &InList) -> std::vector<Value *> {
532 std::vector<Value *> OutList;
533 for (unsigned i = 0; i < InList.size() - 1; i = i + 2) {
534 Value *Or = Builder.CreateOr(LHS: InList[i], RHS: InList[i + 1]);
535 OutList.push_back(x: Or);
536 }
537 if (InList.size() % 2 != 0)
538 OutList.push_back(x: InList.back());
539 return OutList;
540 };
541
542 if (!Cmp) {
543 // Pairwise OR the XOR results.
544 OrList = pairWiseOr(XorList);
545
546 // Pairwise OR the OR results until one result left.
547 while (OrList.size() != 1) {
548 OrList = pairWiseOr(OrList);
549 }
550
551 assert(Diff && "Failed to find comparison diff");
552 Cmp = Builder.CreateICmpNE(LHS: OrList[0], RHS: ConstantInt::get(Ty: Diff->getType(), V: 0));
553 }
554
555 return Cmp;
556}
557
558MemCmpExpansion::LoadPair
559MemCmpExpansion::getPackedLoadPair(unsigned BlockIndex, unsigned &LoadIndex) {
560 assert(LoadIndex < getNumLoads() &&
561 "getPackedLoadPair() called with no remaining loads");
562 if (LoadCmpBlocks.empty())
563 Builder.SetInsertPoint(CI);
564 else
565 Builder.SetInsertPoint(LoadCmpBlocks[BlockIndex]);
566
567 const unsigned NumLoads = getNumLoadsInBlock(LoadIndex);
568 const unsigned NumBytes = getNumBytesInBlock(LoadIndex, NumLoads);
569 // Pack the loads so that the byte at the lowest address occupies the most
570 // significant bits. An unsigned comparison of the packed values therefore
571 // has the same lexicographic ordering as memcmp.
572 auto *BlockType = IntegerType::get(C&: CI->getContext(), NumBits: MaxBlockSize * 8);
573 Value *PackedLhs = ConstantInt::get(Ty: BlockType, V: 0);
574 Value *PackedRhs = ConstantInt::get(Ty: BlockType, V: 0);
575 unsigned RemainingBytes = NumBytes;
576
577 for (unsigned I = 0; I != NumLoads; ++I, ++LoadIndex) {
578 const LoadEntry &Entry = LoadSequence[LoadIndex];
579 auto *LoadType = IntegerType::get(C&: CI->getContext(), NumBits: Entry.LoadSize * 8);
580 auto *BSwapType = DL.isLittleEndian() && Entry.LoadSize != 1
581 ? IntegerType::get(C&: CI->getContext(),
582 NumBits: PowerOf2Ceil(A: Entry.LoadSize * 8))
583 : nullptr;
584 LoadPair Loads = getLoadPair(LoadSizeType: LoadType, BSwapSizeType: BSwapType, CmpSizeType: BlockType, OffsetBytes: Entry.Offset);
585
586 if (BSwapType && BSwapType->getIntegerBitWidth() != Entry.LoadSize * 8) {
587 unsigned Padding = BSwapType->getIntegerBitWidth() - Entry.LoadSize * 8;
588 Loads.Lhs = Builder.CreateLShr(LHS: Loads.Lhs, RHS: Padding);
589 Loads.Rhs = Builder.CreateLShr(LHS: Loads.Rhs, RHS: Padding);
590 }
591
592 RemainingBytes -= Entry.LoadSize;
593 unsigned Shift = RemainingBytes * 8;
594 PackedLhs =
595 Builder.CreateOr(LHS: PackedLhs, RHS: Builder.CreateShl(LHS: Loads.Lhs, RHS: Shift));
596 PackedRhs =
597 Builder.CreateOr(LHS: PackedRhs, RHS: Builder.CreateShl(LHS: Loads.Rhs, RHS: Shift));
598 }
599
600 return {.Lhs: PackedLhs, .Rhs: PackedRhs};
601}
602
603void MemCmpExpansion::emitLoadCompareBlockMultipleLoads(unsigned BlockIndex,
604 unsigned &LoadIndex) {
605 Value *Cmp = getCompareLoadPairs(BlockIndex, LoadIndex);
606
607 BasicBlock *NextBB = (BlockIndex == (LoadCmpBlocks.size() - 1))
608 ? EndBlock
609 : LoadCmpBlocks[BlockIndex + 1];
610 // Early exit branch if difference found to ResultBlock. Otherwise,
611 // continue to next LoadCmpBlock or EndBlock.
612 BasicBlock *BB = Builder.GetInsertBlock();
613 CondBrInst *CmpBr = Builder.CreateCondBr(Cond: Cmp, True: ResBlock.BB, False: NextBB);
614 setExplicitlyUnknownBranchWeightsIfProfiled(I&: *CmpBr, DEBUG_TYPE,
615 F: CI->getFunction());
616 if (DTU)
617 DTU->applyUpdates(Updates: {{DominatorTree::Insert, BB, ResBlock.BB},
618 {DominatorTree::Insert, BB, NextBB}});
619
620 // Add a phi edge for the last LoadCmpBlock to Endblock with a value of 0
621 // since early exit to ResultBlock was not taken (no difference was found in
622 // any of the bytes).
623 if (BlockIndex == LoadCmpBlocks.size() - 1) {
624 Value *Zero = ConstantInt::get(Ty: Type::getInt32Ty(C&: CI->getContext()), V: 0);
625 PhiRes->addIncoming(V: Zero, BB: LoadCmpBlocks[BlockIndex]);
626 }
627}
628
629// This function creates the IR intructions for loading and comparing using the
630// given LoadSize. It loads the number of bytes specified by LoadSize from each
631// source of the memcmp parameters. It then does a subtract to see if there was
632// a difference in the loaded values. If a difference is found, it branches
633// with an early exit to the ResultBlock for calculating which source was
634// larger. Otherwise, it falls through to the either the next LoadCmpBlock or
635// the EndBlock if this is the last LoadCmpBlock. Loading 1 byte is handled with
636// a special case through emitLoadCompareByteBlock. The special handling can
637// simply subtract the loaded values and add it to the result phi node.
638void MemCmpExpansion::emitLoadCompareBlock(unsigned BlockIndex,
639 unsigned &LoadIndex) {
640 const unsigned NumLoads = getNumLoadsInBlock(LoadIndex);
641 if (NumLoads == 1 && LoadSequence[LoadIndex].LoadSize == 1) {
642 MemCmpExpansion::emitLoadCompareByteBlock(BlockIndex,
643 OffsetBytes: LoadSequence[LoadIndex].Offset);
644 ++LoadIndex;
645 return;
646 }
647
648 LoadPair Loads;
649 if (NumLoads == 1) {
650 const LoadEntry &Entry = LoadSequence[LoadIndex++];
651 auto *LoadType = IntegerType::get(C&: CI->getContext(), NumBits: Entry.LoadSize * 8);
652 auto *BSwapType = DL.isLittleEndian()
653 ? IntegerType::get(C&: CI->getContext(),
654 NumBits: PowerOf2Ceil(A: Entry.LoadSize * 8))
655 : nullptr;
656 auto *CmpType = IntegerType::get(C&: CI->getContext(), NumBits: MaxBlockSize * 8);
657 Builder.SetInsertPoint(LoadCmpBlocks[BlockIndex]);
658 Loads = getLoadPair(LoadSizeType: LoadType, BSwapSizeType: BSwapType, CmpSizeType: CmpType, OffsetBytes: Entry.Offset);
659 } else {
660 Loads = getPackedLoadPair(BlockIndex, LoadIndex);
661 }
662
663 // Add the loaded values to the phi nodes for calculating memcmp result only
664 // if result is not used in a zero equality.
665 if (!IsUsedForZeroCmp) {
666 ResBlock.PhiSrc1->addIncoming(V: Loads.Lhs, BB: LoadCmpBlocks[BlockIndex]);
667 ResBlock.PhiSrc2->addIncoming(V: Loads.Rhs, BB: LoadCmpBlocks[BlockIndex]);
668 }
669
670 Value *Cmp = Builder.CreateICmp(P: ICmpInst::ICMP_EQ, LHS: Loads.Lhs, RHS: Loads.Rhs);
671 BasicBlock *NextBB = (BlockIndex == (LoadCmpBlocks.size() - 1))
672 ? EndBlock
673 : LoadCmpBlocks[BlockIndex + 1];
674 // Early exit branch if difference found to ResultBlock. Otherwise, continue
675 // to next LoadCmpBlock or EndBlock.
676 BasicBlock *BB = Builder.GetInsertBlock();
677 CondBrInst *CmpBr = Builder.CreateCondBr(Cond: Cmp, True: NextBB, False: ResBlock.BB);
678 setExplicitlyUnknownBranchWeightsIfProfiled(I&: *CmpBr, DEBUG_TYPE,
679 F: CI->getFunction());
680 if (DTU)
681 DTU->applyUpdates(Updates: {{DominatorTree::Insert, BB, NextBB},
682 {DominatorTree::Insert, BB, ResBlock.BB}});
683
684 // Add a phi edge for the last LoadCmpBlock to Endblock with a value of 0
685 // since early exit to ResultBlock was not taken (no difference was found in
686 // any of the bytes).
687 if (BlockIndex == LoadCmpBlocks.size() - 1) {
688 Value *Zero = ConstantInt::get(Ty: Type::getInt32Ty(C&: CI->getContext()), V: 0);
689 PhiRes->addIncoming(V: Zero, BB: LoadCmpBlocks[BlockIndex]);
690 }
691}
692
693// This function populates the ResultBlock with a sequence to calculate the
694// memcmp result. It compares the two loaded source values and returns -1 if
695// src1 < src2 and 1 if src1 > src2.
696void MemCmpExpansion::emitMemCmpResultBlock() {
697 // Special case: if memcmp result is used in a zero equality, result does not
698 // need to be calculated and can simply return 1.
699 if (IsUsedForZeroCmp) {
700 BasicBlock::iterator InsertPt = ResBlock.BB->getFirstInsertionPt();
701 Builder.SetInsertPoint(InsertPt);
702 Value *Res = ConstantInt::get(Ty: Type::getInt32Ty(C&: CI->getContext()), V: 1);
703 PhiRes->addIncoming(V: Res, BB: ResBlock.BB);
704 Builder.CreateBr(Dest: EndBlock);
705 if (DTU)
706 DTU->applyUpdates(Updates: {{DominatorTree::Insert, ResBlock.BB, EndBlock}});
707 return;
708 }
709 BasicBlock::iterator InsertPt = ResBlock.BB->getFirstInsertionPt();
710 Builder.SetInsertPoint(InsertPt);
711
712 Value *Cmp = Builder.CreateICmp(P: ICmpInst::ICMP_ULT, LHS: ResBlock.PhiSrc1,
713 RHS: ResBlock.PhiSrc2);
714
715 Value *Res =
716 Builder.CreateSelect(C: Cmp, True: Constant::getAllOnesValue(Ty: Builder.getInt32Ty()),
717 False: ConstantInt::get(Ty: Builder.getInt32Ty(), V: 1));
718 setExplicitlyUnknownBranchWeightsIfProfiled(I&: *cast<Instruction>(Val: Res),
719 DEBUG_TYPE, F: CI->getFunction());
720
721 PhiRes->addIncoming(V: Res, BB: ResBlock.BB);
722 Builder.CreateBr(Dest: EndBlock);
723 if (DTU)
724 DTU->applyUpdates(Updates: {{DominatorTree::Insert, ResBlock.BB, EndBlock}});
725}
726
727void MemCmpExpansion::setupResultBlockPHINodes() {
728 Type *MaxLoadType = IntegerType::get(C&: CI->getContext(), NumBits: MaxBlockSize * 8);
729 Builder.SetInsertPoint(ResBlock.BB);
730 ResBlock.PhiSrc1 = Builder.CreatePHI(Ty: MaxLoadType, NumReservedValues: getNumBlocks(), Name: "phi.src1");
731 ResBlock.PhiSrc2 = Builder.CreatePHI(Ty: MaxLoadType, NumReservedValues: getNumBlocks(), Name: "phi.src2");
732}
733
734void MemCmpExpansion::setupEndBlockPHINodes() {
735 Builder.SetInsertPoint(EndBlock->begin());
736 PhiRes = Builder.CreatePHI(Ty: Type::getInt32Ty(C&: CI->getContext()), NumReservedValues: 2, Name: "phi.res");
737}
738
739Value *MemCmpExpansion::getMemCmpExpansionZeroCase() {
740 unsigned LoadIndex = 0;
741 // This loop populates each of the LoadCmpBlocks with the IR sequence to
742 // handle multiple loads per block.
743 for (unsigned I = 0; I < getNumBlocks(); ++I) {
744 emitLoadCompareBlockMultipleLoads(BlockIndex: I, LoadIndex);
745 }
746
747 emitMemCmpResultBlock();
748 return PhiRes;
749}
750
751/// A memcmp expansion that compares equality with 0 and only has one block of
752/// load and compare can bypass the compare, branch, and phi IR that is required
753/// in the general case.
754Value *MemCmpExpansion::getMemCmpEqZeroOneBlock() {
755 unsigned LoadIndex = 0;
756 Value *Cmp = getCompareLoadPairs(BlockIndex: 0, LoadIndex);
757 assert(LoadIndex == getNumLoads() && "some entries were not consumed");
758 return Builder.CreateZExt(V: Cmp, DestTy: Type::getInt32Ty(C&: CI->getContext()));
759}
760
761/// A memcmp expansion that only has one block of load and compare can bypass
762/// the compare, branch, and phi IR that is required in the general case.
763/// This function also analyses users of memcmp, and if there is only one user
764/// from which we can conclude that only 2 out of 3 memcmp outcomes really
765/// matter, then it generates more efficient code with only one comparison.
766Value *MemCmpExpansion::getMemCmpOneBlock() {
767 bool NeedsBSwap = DL.isLittleEndian() && Size != 1;
768 Type *LoadSizeType = IntegerType::get(C&: CI->getContext(), NumBits: Size * 8);
769 Type *BSwapSizeType =
770 NeedsBSwap ? IntegerType::get(C&: CI->getContext(), NumBits: PowerOf2Ceil(A: Size * 8))
771 : nullptr;
772 Type *MaxLoadType =
773 IntegerType::get(C&: CI->getContext(),
774 NumBits: std::max(a: MaxLoadSize, b: (unsigned)PowerOf2Ceil(A: Size)) * 8);
775
776 // The i8 and i16 cases don't need compares. We zext the loaded values and
777 // subtract them to get the suitable negative, zero, or positive i32 result.
778 if (Size == 1 || Size == 2) {
779 const LoadPair Loads = getLoadPair(LoadSizeType, BSwapSizeType,
780 CmpSizeType: Builder.getInt32Ty(), /*Offset*/ OffsetBytes: 0);
781 return Builder.CreateSub(LHS: Loads.Lhs, RHS: Loads.Rhs);
782 }
783
784 const LoadPair Loads = getLoadPair(LoadSizeType, BSwapSizeType, CmpSizeType: MaxLoadType,
785 /*Offset*/ OffsetBytes: 0);
786
787 return getMemCmpResult(Loads);
788}
789
790Value *MemCmpExpansion::getMemCmpOneBlockMultipleLoads() {
791 unsigned LoadIndex = 0;
792 LoadPair Loads = getPackedLoadPair(/*BlockIndex=*/0, LoadIndex);
793 assert(LoadIndex == getNumLoads() && "some entries were not consumed");
794 return getMemCmpResult(Loads);
795}
796
797Value *MemCmpExpansion::getMemCmpResult(const LoadPair &Loads) {
798 // If a user of memcmp cares only about two outcomes, for example:
799 // bool result = memcmp(a, b, NBYTES) > 0;
800 // We can generate more optimal code with a smaller number of operations
801 if (CI->hasOneUser()) {
802 auto *UI = cast<Instruction>(Val: *CI->user_begin());
803 CmpPredicate Pred = ICmpInst::Predicate::BAD_ICMP_PREDICATE;
804 bool NeedsZExt = false;
805 // This is a special case because instead of checking if the result is less
806 // than zero:
807 // bool result = memcmp(a, b, NBYTES) < 0;
808 // Compiler is clever enough to generate the following code:
809 // bool result = memcmp(a, b, NBYTES) >> 31;
810 if (match(V: UI,
811 P: m_LShr(L: m_Value(),
812 R: m_SpecificInt(V: CI->getType()->getIntegerBitWidth() - 1)))) {
813 Pred = ICmpInst::ICMP_SLT;
814 NeedsZExt = true;
815 } else if (match(V: UI, P: m_SpecificICmp(MatchPred: ICmpInst::ICMP_SGT, L: m_Specific(V: CI),
816 R: m_AllOnes()))) {
817 // Adjust predicate as if it compared with 0.
818 Pred = ICmpInst::ICMP_SGE;
819 } else if (match(V: UI, P: m_SpecificICmp(MatchPred: ICmpInst::ICMP_SLT, L: m_Specific(V: CI),
820 R: m_One()))) {
821 // Adjust predicate as if it compared with 0.
822 Pred = ICmpInst::ICMP_SLE;
823 } else {
824 // In case of a successful match this call will set `Pred` variable
825 match(V: UI, P: m_ICmp(Pred, L: m_Specific(V: CI), R: m_Zero()));
826 }
827 // Generate new code and remove the original memcmp call and the user
828 if (ICmpInst::isSigned(Pred)) {
829 Value *Cmp = Builder.CreateICmp(P: ICmpInst::getUnsignedPredicate(Pred),
830 LHS: Loads.Lhs, RHS: Loads.Rhs);
831 auto *Result = NeedsZExt ? Builder.CreateZExt(V: Cmp, DestTy: UI->getType()) : Cmp;
832 UI->replaceAllUsesWith(V: Result);
833 UI->eraseFromParent();
834 CI->eraseFromParent();
835 return nullptr;
836 }
837 }
838
839 // The result of memcmp is negative, zero, or positive.
840 return Builder.CreateIntrinsic(RetTy: Builder.getInt32Ty(), ID: Intrinsic::ucmp,
841 Args: {Loads.Lhs, Loads.Rhs});
842}
843
844// This function expands the memcmp call into an inline expansion and returns
845// the memcmp result. Returns nullptr if the memcmp is already replaced.
846Value *MemCmpExpansion::getMemCmpExpansion() {
847 // Create the basic block framework for a multi-block expansion.
848 if (getNumBlocks() != 1) {
849 BasicBlock *StartBlock = CI->getParent();
850 EndBlock = SplitBlock(Old: StartBlock, SplitPt: CI, DTU, /*LI=*/nullptr,
851 /*MSSAU=*/nullptr, BBName: "endblock");
852 setupEndBlockPHINodes();
853 createResultBlock();
854
855 // If return value of memcmp is not used in a zero equality, we need to
856 // calculate which source was larger. The calculation requires the
857 // two loaded source values of each load compare block.
858 // These will be saved in the phi nodes created by setupResultBlockPHINodes.
859 if (!IsUsedForZeroCmp) setupResultBlockPHINodes();
860
861 // Create the number of required load compare basic blocks.
862 createLoadCmpBlocks();
863
864 // Update the terminator added by SplitBlock to branch to the first
865 // LoadCmpBlock.
866 StartBlock->getTerminator()->setSuccessor(Idx: 0, BB: LoadCmpBlocks[0]);
867 if (DTU)
868 DTU->applyUpdates(Updates: {{DominatorTree::Insert, StartBlock, LoadCmpBlocks[0]},
869 {DominatorTree::Delete, StartBlock, EndBlock}});
870 }
871
872 Builder.SetCurrentDebugLocation(CI->getDebugLoc());
873
874 if (IsUsedForZeroCmp)
875 return getNumBlocks() == 1 ? getMemCmpEqZeroOneBlock()
876 : getMemCmpExpansionZeroCase();
877
878 if (getNumBlocks() == 1)
879 return getNumLoads() == 1 ? getMemCmpOneBlock()
880 : getMemCmpOneBlockMultipleLoads();
881
882 unsigned LoadIndex = 0;
883 for (unsigned I = 0; I < getNumBlocks(); ++I) {
884 emitLoadCompareBlock(BlockIndex: I, LoadIndex);
885 }
886
887 emitMemCmpResultBlock();
888 return PhiRes;
889}
890
891// This function checks to see if an expansion of memcmp can be generated.
892// It checks for constant compare size that is less than the max inline size.
893// If an expansion cannot occur, returns false to leave as a library call.
894// Otherwise, the library call is replaced with a new IR instruction sequence.
895/// We want to transform:
896/// %call = call signext i32 @memcmp(i8* %0, i8* %1, i64 15)
897/// To:
898/// loadbb:
899/// %0 = bitcast i32* %buffer2 to i8*
900/// %1 = bitcast i32* %buffer1 to i8*
901/// %2 = bitcast i8* %1 to i64*
902/// %3 = bitcast i8* %0 to i64*
903/// %4 = load i64, i64* %2
904/// %5 = load i64, i64* %3
905/// %6 = call i64 @llvm.bswap.i64(i64 %4)
906/// %7 = call i64 @llvm.bswap.i64(i64 %5)
907/// %8 = sub i64 %6, %7
908/// %9 = icmp ne i64 %8, 0
909/// br i1 %9, label %res_block, label %loadbb1
910/// res_block: ; preds = %loadbb2,
911/// %loadbb1, %loadbb
912/// %phi.src1 = phi i64 [ %6, %loadbb ], [ %22, %loadbb1 ], [ %36, %loadbb2 ]
913/// %phi.src2 = phi i64 [ %7, %loadbb ], [ %23, %loadbb1 ], [ %37, %loadbb2 ]
914/// %10 = icmp ult i64 %phi.src1, %phi.src2
915/// %11 = select i1 %10, i32 -1, i32 1
916/// br label %endblock
917/// loadbb1: ; preds = %loadbb
918/// %12 = bitcast i32* %buffer2 to i8*
919/// %13 = bitcast i32* %buffer1 to i8*
920/// %14 = bitcast i8* %13 to i32*
921/// %15 = bitcast i8* %12 to i32*
922/// %16 = getelementptr i32, i32* %14, i32 2
923/// %17 = getelementptr i32, i32* %15, i32 2
924/// %18 = load i32, i32* %16
925/// %19 = load i32, i32* %17
926/// %20 = call i32 @llvm.bswap.i32(i32 %18)
927/// %21 = call i32 @llvm.bswap.i32(i32 %19)
928/// %22 = zext i32 %20 to i64
929/// %23 = zext i32 %21 to i64
930/// %24 = sub i64 %22, %23
931/// %25 = icmp ne i64 %24, 0
932/// br i1 %25, label %res_block, label %loadbb2
933/// loadbb2: ; preds = %loadbb1
934/// %26 = bitcast i32* %buffer2 to i8*
935/// %27 = bitcast i32* %buffer1 to i8*
936/// %28 = bitcast i8* %27 to i16*
937/// %29 = bitcast i8* %26 to i16*
938/// %30 = getelementptr i16, i16* %28, i16 6
939/// %31 = getelementptr i16, i16* %29, i16 6
940/// %32 = load i16, i16* %30
941/// %33 = load i16, i16* %31
942/// %34 = call i16 @llvm.bswap.i16(i16 %32)
943/// %35 = call i16 @llvm.bswap.i16(i16 %33)
944/// %36 = zext i16 %34 to i64
945/// %37 = zext i16 %35 to i64
946/// %38 = sub i64 %36, %37
947/// %39 = icmp ne i64 %38, 0
948/// br i1 %39, label %res_block, label %loadbb3
949/// loadbb3: ; preds = %loadbb2
950/// %40 = bitcast i32* %buffer2 to i8*
951/// %41 = bitcast i32* %buffer1 to i8*
952/// %42 = getelementptr i8, i8* %41, i8 14
953/// %43 = getelementptr i8, i8* %40, i8 14
954/// %44 = load i8, i8* %42
955/// %45 = load i8, i8* %43
956/// %46 = zext i8 %44 to i32
957/// %47 = zext i8 %45 to i32
958/// %48 = sub i32 %46, %47
959/// br label %endblock
960/// endblock: ; preds = %res_block,
961/// %loadbb3
962/// %phi.res = phi i32 [ %48, %loadbb3 ], [ %11, %res_block ]
963/// ret i32 %phi.res
964static bool expandMemCmp(CallInst *CI, const TargetTransformInfo *TTI,
965 const DataLayout *DL, ProfileSummaryInfo *PSI,
966 BlockFrequencyInfo *BFI, DomTreeUpdater *DTU,
967 const bool IsBCmp) {
968 const ScalarOptions &Opts = ScalarOptions::Global;
969 NumMemCmpCalls++;
970
971 // Early exit from expansion if -Oz.
972 if (CI->getFunction()->hasMinSize())
973 return false;
974
975 // Early exit from expansion if size is not a constant.
976 ConstantInt *SizeCast = dyn_cast<ConstantInt>(Val: CI->getArgOperand(i: 2));
977 if (!SizeCast) {
978 NumMemCmpNotConstant++;
979 return false;
980 }
981 const uint64_t SizeVal = SizeCast->getZExtValue();
982
983 if (SizeVal == 0) {
984 return false;
985 }
986 // TTI call to check if target would like to expand memcmp. Also, get the
987 // available load sizes.
988 const bool IsUsedForZeroCmp =
989 IsBCmp || isOnlyUsedInZeroEqualityComparison(CtxI: CI);
990 bool OptForSize = llvm::shouldOptimizeForSize(BB: CI->getParent(), PSI, BFI);
991 auto Options = TTI->enableMemCmpExpansion(OptSize: OptForSize,
992 IsZeroCmp: IsUsedForZeroCmp);
993 if (!Options) return false;
994
995 if (Opts.memcmp_num_loads_per_block)
996 Options.NumLoadsPerBlock = *Opts.memcmp_num_loads_per_block;
997
998 if (OptForSize && Opts.max_loads_per_memcmp_opt_size)
999 Options.MaxNumLoads = *Opts.max_loads_per_memcmp_opt_size;
1000
1001 if (!OptForSize && Opts.max_loads_per_memcmp)
1002 Options.MaxNumLoads = *Opts.max_loads_per_memcmp;
1003
1004 // Keep only the load sizes the target can access at the base alignment:
1005 // either the access is naturally aligned, or the target allows a misaligned
1006 // access of that width. This lets strict-alignment targets expand compares
1007 // whose pointers happen to be sufficiently aligned, while still falling back
1008 // to the libcall when no load size fits. Because the greedy load sequence
1009 // only places a load of size S at an offset that is a multiple of S, a size
1010 // kept here is always accessible in that sequence; overlapping loads and
1011 // merged tail expansions are checked separately against their actual offsets
1012 // in MemCmpExpansion.
1013 const Align CommonAlign = std::min(a: getMemCmpArgAlignment(CI, ArgNo: 0, DL: *DL),
1014 b: getMemCmpArgAlignment(CI, ArgNo: 1, DL: *DL));
1015 // Remember the target's preferred load width before filtering inaccessible
1016 // sizes. It is also the maximum width of the value formed by packing the
1017 // surviving loads in one ordering-compare block.
1018 const unsigned MaxBytesPerBlock = Options.LoadSizes.front();
1019 llvm::erase_if(C&: Options.LoadSizes, P: [&](unsigned LoadSize) {
1020 return !isAccessAllowed(CI, TTI: *TTI, CommonAlign, LoadSize, /*Offset=*/0);
1021 });
1022 // If the filter removed every load size, bail out to the libcall: the
1023 // MemCmpExpansion constructor asserts that at least one load size remains.
1024 // In practice all in-tree targets include a byte load size, which is
1025 // accessible at any alignment and therefore always survives the filter.
1026 if (Options.LoadSizes.empty())
1027 return false;
1028
1029 MemCmpExpansion Expansion(CI, SizeVal, Options, IsUsedForZeroCmp, *DL, DTU,
1030 *TTI, CommonAlign, MaxBytesPerBlock);
1031
1032 // Don't expand if this will require more loads than desired by the target.
1033 if (Expansion.getNumLoads() == 0) {
1034 NumMemCmpGreaterThanMax++;
1035 return false;
1036 }
1037
1038 NumMemCmpInlined++;
1039
1040 if (Value *Res = Expansion.getMemCmpExpansion()) {
1041 // Replace call with result of expansion and erase call.
1042 CI->replaceAllUsesWith(V: Res);
1043 CI->eraseFromParent();
1044 }
1045
1046 return true;
1047}
1048
1049static PreservedAnalyses runImpl(Function &F, const TargetLibraryInfo *TLI,
1050 const TargetTransformInfo *TTI,
1051 ProfileSummaryInfo *PSI,
1052 BlockFrequencyInfo *BFI, DominatorTree *DT) {
1053 std::optional<DomTreeUpdater> DTU;
1054 if (DT)
1055 DTU.emplace(args&: DT, args: DomTreeUpdater::UpdateStrategy::Lazy);
1056
1057 const DataLayout& DL = F.getDataLayout();
1058 SmallVector<std::pair<CallInst *, LibFunc>, 8> MemCmpCalls;
1059 for (Instruction &I : instructions(F)) {
1060 if (auto *CI = dyn_cast<CallInst>(Val: &I)) {
1061 LibFunc Func = TLI->getLibFunc(CB: *CI);
1062 if (Func == LibFunc_memcmp || Func == LibFunc_bcmp)
1063 MemCmpCalls.push_back(Elt: {CI, Func});
1064 }
1065 }
1066
1067 bool MadeChanges = false;
1068 for (const auto &[CI, Func] : MemCmpCalls) {
1069 if (expandMemCmp(CI, TTI, DL: &DL, PSI, BFI, DTU: DTU ? &*DTU : nullptr,
1070 IsBCmp: Func == LibFunc_bcmp))
1071 MadeChanges = true;
1072 }
1073
1074 if (MadeChanges)
1075 for (BasicBlock &BB : F)
1076 SimplifyInstructionsInBlock(BB: &BB);
1077 if (!MadeChanges)
1078 return PreservedAnalyses::all();
1079 PreservedAnalyses PA;
1080 PA.preserve<DominatorTreeAnalysis>();
1081 return PA;
1082}
1083
1084} // namespace
1085
1086PreservedAnalyses ExpandMemCmpPass::run(Function &F,
1087 FunctionAnalysisManager &FAM) {
1088 // Don't expand memcmp in sanitized functions — sanitizers intercept memcmp
1089 // calls to check for memory errors, and expanding would bypass that.
1090 if (F.hasFnAttribute(Kind: Attribute::SanitizeAddress) ||
1091 F.hasFnAttribute(Kind: Attribute::SanitizeMemory) ||
1092 F.hasFnAttribute(Kind: Attribute::SanitizeThread) ||
1093 F.hasFnAttribute(Kind: Attribute::SanitizeHWAddress))
1094 return PreservedAnalyses::all();
1095
1096 const auto &TLI = FAM.getResult<TargetLibraryAnalysis>(IR&: F);
1097 const auto &TTI = FAM.getResult<TargetIRAnalysis>(IR&: F);
1098 auto *PSI = FAM.getResult<ModuleAnalysisManagerFunctionProxy>(IR&: F)
1099 .getCachedResult<ProfileSummaryAnalysis>(IR&: *F.getParent());
1100 BlockFrequencyInfo *BFI = (PSI && PSI->hasProfileSummary())
1101 ? &FAM.getResult<BlockFrequencyAnalysis>(IR&: F)
1102 : nullptr;
1103 auto *DT = FAM.getCachedResult<DominatorTreeAnalysis>(IR&: F);
1104
1105 return runImpl(F, TLI: &TLI, TTI: &TTI, PSI, BFI, DT);
1106}
1107