[llvm] [AllocToken] Merge !alloc_token metadata across optimization passes (PR #211516)

Marco Elver via llvm-commits llvm-commits at lists.llvm.org
Fri Jul 31 06:08:39 PDT 2026


https://github.com/melver updated https://github.com/llvm/llvm-project/pull/211516

>From f3e949985d9ed2d7d1d4e93694ed810333c3f3ca Mon Sep 17 00:00:00 2001
From: Marco Elver <elver at google.com>
Date: Thu, 23 Jul 2026 12:37:10 +0200
Subject: [PATCH] [AllocToken] Merge !alloc_token metadata across optimization
 passes

Previously, optimization transforms like DeadStoreElimination (folding
malloc to calloc), SimplifyLibCalls (folding realloc(null, n) to
malloc), and combineMetadata (merging equivalent allocations in
GVN/SimplifyCFG) dropped !alloc_token metadata or stripped it when type
names differed. This caused allocation sites of pointer-containing types
to lose their metadata.

Fix it by implementing MDNode::getMergedAllocTokenMetadata() to merge
!alloc_token metadata by ORing the pointer-containing boolean flag and
concatenating type names with a pipe separator, updating combineMetadata
to use it, and preserving !alloc_token metadata when folding allocation
libcalls in DSE and SimplifyLibCalls.

Assisted-by: Antigravity:gemini
---
 llvm/include/llvm/IR/Metadata.h               |  2 +
 llvm/lib/IR/Metadata.cpp                      | 41 +++++++++++++++++++
 .../Scalar/DeadStoreElimination.cpp           |  3 ++
 llvm/lib/Transforms/Utils/Local.cpp           |  5 +--
 .../lib/Transforms/Utils/SimplifyLibCalls.cpp |  9 +++-
 .../DeadStoreElimination/alloc-token-dse.ll   | 17 ++++++++
 .../InstCombine/alloc-token-realloc.ll        | 15 +++++++
 .../SimplifyCFG/merge-calls-alloc-token.ll    | 29 ++++++++++++-
 8 files changed, 115 insertions(+), 6 deletions(-)
 create mode 100644 llvm/test/Transforms/DeadStoreElimination/alloc-token-dse.ll
 create mode 100644 llvm/test/Transforms/InstCombine/alloc-token-realloc.ll

diff --git a/llvm/include/llvm/IR/Metadata.h b/llvm/include/llvm/IR/Metadata.h
index c0d0703135f68..5b458fa14f0b1 100644
--- a/llvm/include/llvm/IR/Metadata.h
+++ b/llvm/include/llvm/IR/Metadata.h
@@ -1466,6 +1466,8 @@ class MDNode : public Metadata {
   LLVM_ABI static MDNode *getMergedCallsiteMetadata(MDNode *A, MDNode *B);
   LLVM_ABI static MDNode *getMergedCalleeTypeMetadata(const MDNode *A,
                                                       const MDNode *B);
+  LLVM_ABI static MDNode *getMergedAllocTokenMetadata(const MDNode *A,
+                                                      const MDNode *B);
 
   /// Convert !captures metadata to CaptureComponents. MD may be nullptr.
   LLVM_ABI static CaptureComponents toCaptureComponents(const MDNode *MD);
diff --git a/llvm/lib/IR/Metadata.cpp b/llvm/lib/IR/Metadata.cpp
index 491c788fc4445..0a4141ee2362e 100644
--- a/llvm/lib/IR/Metadata.cpp
+++ b/llvm/lib/IR/Metadata.cpp
@@ -21,6 +21,7 @@
 #include "llvm/ADT/SetVector.h"
 #include "llvm/ADT/SmallPtrSet.h"
 #include "llvm/ADT/SmallSet.h"
+#include "llvm/ADT/SmallString.h"
 #include "llvm/ADT/SmallVector.h"
 #include "llvm/ADT/StringMap.h"
 #include "llvm/ADT/StringRef.h"
@@ -1332,6 +1333,46 @@ MDNode *MDNode::getMergedCalleeTypeMetadata(const MDNode *A, const MDNode *B) {
   return MDNode::get(A->getContext(), AB);
 }
 
+MDNode *MDNode::getMergedAllocTokenMetadata(const MDNode *A, const MDNode *B) {
+  // Drop !alloc_token metadata if either instruction lacks it to avoid mis-
+  // classifying unclassified allocations, where the fallback token must be
+  // used instead.
+  if (!A || !B)
+    return nullptr;
+  if (A == B)
+    return const_cast<MDNode *>(A);
+  if (A->getNumOperands() != 2 || B->getNumOperands() != 2)
+    return nullptr;
+  auto *CIA = mdconst::dyn_extract_or_null<ConstantInt>(A->getOperand(1));
+  auto *CIB = mdconst::dyn_extract_or_null<ConstantInt>(B->getOperand(1));
+  if (!CIA || !CIB)
+    return nullptr;
+
+  MDString *NameA = dyn_cast<MDString>(A->getOperand(0));
+  MDString *NameB = dyn_cast<MDString>(B->getOperand(0));
+  if (!NameA || !NameB)
+    return nullptr;
+
+  if (NameA == NameB)
+    return CIA->isOne() ? const_cast<MDNode *>(A) : const_cast<MDNode *>(B);
+
+  LLVMContext &Ctx = A->getContext();
+  StringRef StrA = NameA->getString();
+  StringRef StrB = NameB->getString();
+
+  SmallString<64> Buffer;
+  Buffer.reserve(StrA.size() + 1 + StrB.size());
+  Buffer.append(StrA);
+  Buffer.push_back('|');
+  Buffer.append(StrB);
+
+  bool MergedContainsPointer = CIA->isOne() || CIB->isOne();
+  Metadata *Ops[] = {MDString::get(Ctx, Buffer),
+                     ConstantAsMetadata::get(ConstantInt::get(
+                         Type::getInt1Ty(Ctx), MergedContainsPointer))};
+  return MDNode::get(Ctx, Ops);
+}
+
 MDNode *MDNode::getMostGenericRange(MDNode *A, MDNode *B) {
   // Given two ranges, we want to compute the union of the ranges. This
   // is slightly complicated by having to combine the intervals and merge
diff --git a/llvm/lib/Transforms/Scalar/DeadStoreElimination.cpp b/llvm/lib/Transforms/Scalar/DeadStoreElimination.cpp
index 2e49252c3cd7e..110a3c20654ac 100644
--- a/llvm/lib/Transforms/Scalar/DeadStoreElimination.cpp
+++ b/llvm/lib/Transforms/Scalar/DeadStoreElimination.cpp
@@ -2358,6 +2358,9 @@ bool DSEState::tryFoldIntoCalloc(MemoryDef *Def, const Value *DefUO) {
   if (!Calloc)
     return false;
 
+  if (MDNode *MD = Malloc->getMetadata(LLVMContext::MD_alloc_token))
+    cast<Instruction>(Calloc)->setMetadata(LLVMContext::MD_alloc_token, MD);
+
   MemorySSAUpdater Updater(&MSSA);
   auto *NewAccess = Updater.createMemoryAccessAfter(cast<Instruction>(Calloc),
                                                     nullptr, MallocDef);
diff --git a/llvm/lib/Transforms/Utils/Local.cpp b/llvm/lib/Transforms/Utils/Local.cpp
index b17740c0bc192..003ce1d2d0f2e 100644
--- a/llvm/lib/Transforms/Utils/Local.cpp
+++ b/llvm/lib/Transforms/Utils/Local.cpp
@@ -3067,9 +3067,8 @@ static void combineMetadata(Instruction *K, const Instruction *J,
                                            MDNode::toCaptureComponents(KMD)));
         break;
       case LLVMContext::MD_alloc_token:
-        // Preserve !alloc_token if both K and J have it, and they are equal.
-        if (KMD != JMD)
-          K->setMetadata(Kind, nullptr);
+        if (!AAOnly && KMD != JMD)
+          K->setMetadata(Kind, MDNode::getMergedAllocTokenMetadata(KMD, JMD));
         break;
       }
   }
diff --git a/llvm/lib/Transforms/Utils/SimplifyLibCalls.cpp b/llvm/lib/Transforms/Utils/SimplifyLibCalls.cpp
index e4d71a172c716..b870c7ed49509 100644
--- a/llvm/lib/Transforms/Utils/SimplifyLibCalls.cpp
+++ b/llvm/lib/Transforms/Utils/SimplifyLibCalls.cpp
@@ -1725,8 +1725,13 @@ Value *LibCallSimplifier::optimizeMemSet(CallInst *CI, IRBuilderBase &B) {
 }
 
 Value *LibCallSimplifier::optimizeRealloc(CallInst *CI, IRBuilderBase &B) {
-  if (isa<ConstantPointerNull>(CI->getArgOperand(0)))
-    return copyFlags(*CI, emitMalloc(CI->getArgOperand(1), B, DL, TLI));
+  if (isa<ConstantPointerNull>(CI->getArgOperand(0))) {
+    Value *Malloc = emitMalloc(CI->getArgOperand(1), B, DL, TLI);
+    if (auto *MallocCI = dyn_cast_or_null<CallInst>(Malloc))
+      if (MDNode *MD = CI->getMetadata(LLVMContext::MD_alloc_token))
+        MallocCI->setMetadata(LLVMContext::MD_alloc_token, MD);
+    return copyFlags(*CI, Malloc);
+  }
 
   return nullptr;
 }
diff --git a/llvm/test/Transforms/DeadStoreElimination/alloc-token-dse.ll b/llvm/test/Transforms/DeadStoreElimination/alloc-token-dse.ll
new file mode 100644
index 0000000000000..cd5eb0e138df8
--- /dev/null
+++ b/llvm/test/Transforms/DeadStoreElimination/alloc-token-dse.ll
@@ -0,0 +1,17 @@
+; RUN: opt < %s -passes=dse -S | FileCheck %s
+
+declare noalias ptr @malloc(i64) allockind("alloc,uninitialized")
+declare noalias ptr @calloc(i64, i64) allockind("alloc,zeroed")
+declare void @llvm.memset.p0.i64(ptr, i8, i64, i1)
+
+define ptr @test_dse_malloc_memset_to_calloc() {
+; CHECK-LABEL: define ptr @test_dse_malloc_memset_to_calloc()
+; CHECK-NEXT: %calloc = call ptr @calloc(i64 1, i64 16), !alloc_token [[META:![0-9]+]]
+; CHECK-NEXT: ret ptr %calloc
+  %1 = tail call ptr @malloc(i64 16), !alloc_token !0
+  call void @llvm.memset.p0.i64(ptr %1, i8 0, i64 16, i1 false)
+  ret ptr %1
+}
+
+!0 = !{!"StructA", i1 true}
+; CHECK: [[META]] = !{!"StructA", i1 true}
diff --git a/llvm/test/Transforms/InstCombine/alloc-token-realloc.ll b/llvm/test/Transforms/InstCombine/alloc-token-realloc.ll
new file mode 100644
index 0000000000000..6b5f665facbe5
--- /dev/null
+++ b/llvm/test/Transforms/InstCombine/alloc-token-realloc.ll
@@ -0,0 +1,15 @@
+; RUN: opt < %s -passes=instcombine -S | FileCheck %s
+
+declare ptr @realloc(ptr allocptr, i64) allockind("realloc") allocsize(1)
+declare noalias ptr @malloc(i64) allockind("alloc,uninitialized")
+
+define ptr @test_realloc_null_alloc_token() {
+; CHECK-LABEL: define ptr @test_realloc_null_alloc_token()
+; CHECK-NEXT: %malloc = call dereferenceable_or_null(16) ptr @malloc(i64 16), !alloc_token [[META:![0-9]+]]
+; CHECK-NEXT: ret ptr %malloc
+  %call = call ptr @realloc(ptr null, i64 16), !alloc_token !0
+  ret ptr %call
+}
+
+!0 = !{!"StructA", i1 true}
+; CHECK: [[META]] = !{!"StructA", i1 true}
diff --git a/llvm/test/Transforms/SimplifyCFG/merge-calls-alloc-token.ll b/llvm/test/Transforms/SimplifyCFG/merge-calls-alloc-token.ll
index 42d3dcc92712d..47aec75edbe74 100644
--- a/llvm/test/Transforms/SimplifyCFG/merge-calls-alloc-token.ll
+++ b/llvm/test/Transforms/SimplifyCFG/merge-calls-alloc-token.ll
@@ -32,7 +32,7 @@ define ptr @test_merge_alloc_token_different(i1 %b) {
 ; CHECK-LABEL: define ptr @test_merge_alloc_token_different(
 ; CHECK-SAME: i1 [[B:%.*]]) {
 ; CHECK-NEXT:  [[ENTRY:.*:]]
-; CHECK-NEXT:    [[CALL:%.*]] = call ptr @_Znwm(i64 4)
+; CHECK-NEXT:    [[CALL:%.*]] = call ptr @_Znwm(i64 4), !alloc_token [[META1:![0-9]+]]
 ; CHECK-NEXT:    ret ptr [[CALL]]
 ;
 entry:
@@ -97,8 +97,35 @@ if.end:
   ret ptr %x.0
 }
 
+define ptr @test_merge_alloc_token_both_ptrs(i1 %b) {
+; CHECK-LABEL: define ptr @test_merge_alloc_token_both_ptrs(
+; CHECK-SAME: i1 [[B:%.*]]) {
+; CHECK-NEXT:  [[ENTRY:.*:]]
+; CHECK-NEXT:    [[CALL:%.*]] = call ptr @_Znwm(i64 16), !alloc_token [[META2:![0-9]+]]
+; CHECK-NEXT:    ret ptr [[CALL]]
+;
+entry:
+  br i1 %b, label %if.then, label %if.else
+
+if.then:
+  %call = call ptr @_Znwm(i64 16), !alloc_token !2
+  br label %if.end
+
+if.else:
+  %call1 = call ptr @_Znwm(i64 16), !alloc_token !3
+  br label %if.end
+
+if.end:
+  %x.0 = phi ptr [ %call, %if.then ], [ %call1, %if.else ]
+  ret ptr %x.0
+}
+
 !0 = !{!"int", i1 0}
 !1 = !{!"char[4]", i1 0}
+!2 = !{!"StructA", i1 1}
+!3 = !{!"StructB", i1 1}
 ;.
 ; CHECK: [[META0]] = !{!"int", i1 false}
+; CHECK: [[META1]] = !{!"int|char[4]", i1 false}
+; CHECK: [[META2]] = !{!"StructA|StructB", i1 true}
 ;.



More information about the llvm-commits mailing list