[llvm] Remove the optional bitcast between a musttail call and its ret (PR #201280)
Arseniy Obolenskiy via llvm-commits
llvm-commits at lists.llvm.org
Wed Jun 3 00:22:46 PDT 2026
https://github.com/aobolensk created https://github.com/llvm/llvm-project/pull/201280
Under opaque pointers the only bitcast the verifier could accept in this position is a no-op ptr->ptr cast
Drop it and reduce isTypeCongruent to a plain type equality check
>From dda639b5cada4a5e16a8dfa4e48e17fec17a5f43 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Wed, 3 Jun 2026 09:20:13 +0200
Subject: [PATCH] Remove the optional bitcast between a musttail call and its
ret
Under opaque pointers the only bitcast the verifier could accept in this position is a no-op ptr->ptr cast
Drop it and reduce isTypeCongruent to a plain type equality check
---
llvm/lib/IR/BasicBlock.cpp | 8 ----
llvm/lib/IR/Verifier.cpp | 42 ++++---------------
llvm/lib/Transforms/Utils/InlineFunction.cpp | 23 +---------
.../AddressSanitizer/musttail.ll | 14 -------
.../ThreadSanitizer/tsan_musttail.ll | 13 ------
.../Transforms/CallSiteSplitting/musttail.ll | 24 -----------
.../test/Transforms/SafeStack/X86/musttail.ll | 19 ---------
7 files changed, 10 insertions(+), 133 deletions(-)
diff --git a/llvm/lib/IR/BasicBlock.cpp b/llvm/lib/IR/BasicBlock.cpp
index d611e9c2d0a96..03e38b3b89061 100644
--- a/llvm/lib/IR/BasicBlock.cpp
+++ b/llvm/lib/IR/BasicBlock.cpp
@@ -239,14 +239,6 @@ const CallInst *BasicBlock::getTerminatingMustTailCall() const {
if (Value *RV = RI->getReturnValue()) {
if (RV != Prev)
return nullptr;
-
- // Look through the optional bitcast.
- if (auto *BI = dyn_cast<BitCastInst>(Prev)) {
- RV = BI->getOperand(0);
- Prev = BI->getPrevNode();
- if (!Prev || RV != Prev)
- return nullptr;
- }
}
if (auto *CI = dyn_cast<CallInst>(Prev)) {
diff --git a/llvm/lib/IR/Verifier.cpp b/llvm/lib/IR/Verifier.cpp
index c9639d1420bfc..bc6c5882b1076 100644
--- a/llvm/lib/IR/Verifier.cpp
+++ b/llvm/lib/IR/Verifier.cpp
@@ -4246,18 +4246,6 @@ void Verifier::verifyTailCCMustTailAttrs(const AttrBuilder &Attrs,
Twine("byref attribute not allowed in ") + Context);
}
-/// Two types are "congruent" if they are identical, or if they are both pointer
-/// types with different pointee types and the same address space.
-static bool isTypeCongruent(Type *L, Type *R) {
- if (L == R)
- return true;
- PointerType *PL = dyn_cast<PointerType>(L);
- PointerType *PR = dyn_cast<PointerType>(R);
- if (!PL || !PR)
- return false;
- return PL->getAddressSpace() == PR->getAddressSpace();
-}
-
static AttrBuilder getParameterABIAttributes(LLVMContext& C, unsigned I, AttributeList Attrs) {
static const Attribute::AttrKind ABIAttrs[] = {
Attribute::StructRet, Attribute::ByVal, Attribute::InAlloca,
@@ -4287,32 +4275,21 @@ void Verifier::verifyMustTailCall(CallInst &CI) {
FunctionType *CalleeTy = CI.getFunctionType();
Check(CallerTy->isVarArg() == CalleeTy->isVarArg(),
"cannot guarantee tail call due to mismatched varargs", &CI);
- Check(isTypeCongruent(CallerTy->getReturnType(), CalleeTy->getReturnType()),
+ Check(CallerTy->getReturnType() == CalleeTy->getReturnType(),
"cannot guarantee tail call due to mismatched return types", &CI);
// - The calling conventions of the caller and callee must match.
Check(F->getCallingConv() == CI.getCallingConv(),
"cannot guarantee tail call due to mismatched calling conv", &CI);
- // - The call must immediately precede a :ref:`ret <i_ret>` instruction,
- // or a pointer bitcast followed by a ret instruction.
- // - The ret instruction must return the (possibly bitcasted) value
- // produced by the call or void.
- Value *RetVal = &CI;
+ // - The call must immediately precede a :ref:`ret <i_ret>` instruction.
+ // - The ret instruction must return the value produced by the call or void.
Instruction *Next = CI.getNextNode();
- // Handle the optional bitcast.
- if (BitCastInst *BI = dyn_cast_or_null<BitCastInst>(Next)) {
- Check(BI->getOperand(0) == RetVal,
- "bitcast following musttail call must use the call", BI);
- RetVal = BI;
- Next = BI->getNextNode();
- }
-
// Check the return.
ReturnInst *Ret = dyn_cast_or_null<ReturnInst>(Next);
- Check(Ret, "musttail call must precede a ret with an optional bitcast", &CI);
- Check(!Ret->getReturnValue() || Ret->getReturnValue() == RetVal ||
+ Check(Ret, "musttail call must precede a ret", &CI);
+ Check(!Ret->getReturnValue() || Ret->getReturnValue() == &CI ||
isa<UndefValue>(Ret->getReturnValue()),
"musttail call result must be returned", Ret);
@@ -4341,16 +4318,13 @@ void Verifier::verifyMustTailCall(CallInst &CI) {
return;
}
- // - The caller and callee prototypes must match. Pointer types of
- // parameters or return types may differ in pointee type, but not
- // address space.
+ // - The caller and callee prototypes must match.
if (!CI.getIntrinsicID()) {
Check(CallerTy->getNumParams() == CalleeTy->getNumParams(),
"cannot guarantee tail call due to mismatched parameter counts", &CI);
for (unsigned I = 0, E = CallerTy->getNumParams(); I != E; ++I) {
- Check(
- isTypeCongruent(CallerTy->getParamType(I), CalleeTy->getParamType(I)),
- "cannot guarantee tail call due to mismatched parameter types", &CI);
+ Check(CallerTy->getParamType(I) == CalleeTy->getParamType(I),
+ "cannot guarantee tail call due to mismatched parameter types", &CI);
}
}
diff --git a/llvm/lib/Transforms/Utils/InlineFunction.cpp b/llvm/lib/Transforms/Utils/InlineFunction.cpp
index 1d3f66509b1c5..be186ffbf7e42 100644
--- a/llvm/lib/Transforms/Utils/InlineFunction.cpp
+++ b/llvm/lib/Transforms/Utils/InlineFunction.cpp
@@ -3300,32 +3300,13 @@ void llvm::InlineFunctionImpl(CallBase &CB, InlineFunctionInfo &IFI,
// musttail. Therefore it's safe to return without merging control into the
// phi below.
if (InlinedMustTailCalls) {
- // Check if we need to bitcast the result of any musttail calls.
- Type *NewRetTy = Caller->getReturnType();
- bool NeedBitCast = !CB.use_empty() && CB.getType() != NewRetTy;
-
// Handle the returns preceded by musttail calls separately.
SmallVector<ReturnInst *, 8> NormalReturns;
for (ReturnInst *RI : Returns) {
CallInst *ReturnedMustTail =
RI->getParent()->getTerminatingMustTailCall();
- if (!ReturnedMustTail) {
+ if (!ReturnedMustTail)
NormalReturns.push_back(RI);
- continue;
- }
- if (!NeedBitCast)
- continue;
-
- // Delete the old return and any preceding bitcast.
- BasicBlock *CurBB = RI->getParent();
- auto *OldCast = dyn_cast_or_null<BitCastInst>(RI->getReturnValue());
- RI->eraseFromParent();
- if (OldCast)
- OldCast->eraseFromParent();
-
- // Insert a new bitcast and return with the right type.
- IRBuilder<> Builder(CurBB);
- Builder.CreateRet(Builder.CreateBitCast(ReturnedMustTail, NewRetTy));
}
// Leave behind the normal returns so we can merge control flow.
@@ -3547,7 +3528,7 @@ void llvm::InlineFunctionImpl(CallBase &CB, InlineFunctionInfo &IFI,
CB.eraseFromParent();
// If we inlined any musttail calls and the original return is now
- // unreachable, delete it. It can only contain a bitcast and ret.
+ // unreachable, delete it. It can only contain a ret.
if (InlinedMustTailCalls && pred_empty(AfterCallBB))
AfterCallBB->eraseFromParent();
diff --git a/llvm/test/Instrumentation/AddressSanitizer/musttail.ll b/llvm/test/Instrumentation/AddressSanitizer/musttail.ll
index fed4521c195a4..8ec87693d1859 100644
--- a/llvm/test/Instrumentation/AddressSanitizer/musttail.ll
+++ b/llvm/test/Instrumentation/AddressSanitizer/musttail.ll
@@ -18,17 +18,3 @@ define i32 @call_foo(ptr %a) sanitize_address {
; CHECK-LABEL: define i32 @call_foo(ptr %a)
; CHECK: %r = musttail call i32 @foo(ptr %a)
; CHECK-NEXT: ret i32 %r
-
-
-define i32 @call_foo_cast(ptr %a) sanitize_address {
- %x = alloca [10 x i8], align 1
- call void @alloca_test_use(ptr %x)
- %r = musttail call i32 @foo(ptr %a)
- %t = bitcast i32 %r to i32
- ret i32 %t
-}
-
-; CHECK-LABEL: define i32 @call_foo_cast(ptr %a)
-; CHECK: %r = musttail call i32 @foo(ptr %a)
-; CHECK-NEXT: %t = bitcast i32 %r to i32
-; CHECK-NEXT: ret i32 %t
diff --git a/llvm/test/Instrumentation/ThreadSanitizer/tsan_musttail.ll b/llvm/test/Instrumentation/ThreadSanitizer/tsan_musttail.ll
index 5e56aa2d11068..2d16a82f666d1 100644
--- a/llvm/test/Instrumentation/ThreadSanitizer/tsan_musttail.ll
+++ b/llvm/test/Instrumentation/ThreadSanitizer/tsan_musttail.ll
@@ -15,16 +15,3 @@ define i32 @call_preallocated_musttail(ptr preallocated(i32) %a) sanitize_thread
; CHECK: call void @__tsan_func_exit()
; CHECK-NEXT: %r = musttail call i32 @preallocated_musttail(ptr preallocated(i32) %a)
; CHECK-NEXT: ret i32 %r
-
-
-define i32 @call_preallocated_musttail_cast(ptr preallocated(i32) %a) sanitize_thread {
- %r = musttail call i32 @preallocated_musttail(ptr preallocated(i32) %a)
- %t = bitcast i32 %r to i32
- ret i32 %t
-}
-
-; CHECK-LABEL: define i32 @call_preallocated_musttail_cast(ptr preallocated(i32) %a)
-; CHECK: call void @__tsan_func_exit()
-; CHECK-NEXT: %r = musttail call i32 @preallocated_musttail(ptr preallocated(i32) %a)
-; CHECK-NEXT: %t = bitcast i32 %r to i32
-; CHECK-NEXT: ret i32 %t
diff --git a/llvm/test/Transforms/CallSiteSplitting/musttail.ll b/llvm/test/Transforms/CallSiteSplitting/musttail.ll
index 0f989a2ae4ad1..1993e1e97e403 100644
--- a/llvm/test/Transforms/CallSiteSplitting/musttail.ll
+++ b/llvm/test/Transforms/CallSiteSplitting/musttail.ll
@@ -1,29 +1,5 @@
; RUN: opt < %s -passes=callsite-splitting -verify-dom-info -S | FileCheck %s
-;CHECK-LABEL: @caller
-;CHECK-LABEL: Top.split:
-;CHECK: %ca1 = musttail call ptr @callee(ptr null, ptr %b)
-;CHECK: %cb2 = bitcast ptr %ca1 to ptr
-;CHECK: ret ptr %cb2
-;CHECK-LABEL: TBB.split
-;CHECK: %ca3 = musttail call ptr @callee(ptr nonnull %a, ptr null)
-;CHECK: %cb4 = bitcast ptr %ca3 to ptr
-;CHECK: ret ptr %cb4
-define ptr @caller(ptr %a, ptr %b) {
-Top:
- %c = icmp eq ptr %a, null
- br i1 %c, label %Tail, label %TBB
-TBB:
- %c2 = icmp eq ptr %b, null
- br i1 %c2, label %Tail, label %End
-Tail:
- %ca = musttail call ptr @callee(ptr %a, ptr %b)
- %cb = bitcast ptr %ca to ptr
- ret ptr %cb
-End:
- ret ptr null
-}
-
define ptr @callee(ptr %a, ptr %b) noinline {
ret ptr %a
}
diff --git a/llvm/test/Transforms/SafeStack/X86/musttail.ll b/llvm/test/Transforms/SafeStack/X86/musttail.ll
index 0289729f53b50..317bc3920fb1b 100644
--- a/llvm/test/Transforms/SafeStack/X86/musttail.ll
+++ b/llvm/test/Transforms/SafeStack/X86/musttail.ll
@@ -24,22 +24,3 @@ define i32 @call_foo(ptr %a) safestack {
%r = musttail call i32 @foo(ptr %a)
ret i32 %r
}
-
-define i32 @call_foo_cast(ptr %a) safestack {
-; CHECK-LABEL: @call_foo_cast(
-; CHECK-NEXT: [[UNSAFE_STACK_PTR:%.*]] = load ptr, ptr @__safestack_unsafe_stack_ptr, align 8
-; CHECK-NEXT: [[UNSAFE_STACK_STATIC_TOP:%.*]] = getelementptr i8, ptr [[UNSAFE_STACK_PTR]], i32 -16
-; CHECK-NEXT: store ptr [[UNSAFE_STACK_STATIC_TOP]], ptr @__safestack_unsafe_stack_ptr, align 8
-; CHECK-NEXT: [[TMP1:%.*]] = getelementptr i8, ptr [[UNSAFE_STACK_PTR]], i32 -10
-; CHECK-NEXT: call void @alloca_test_use(ptr [[TMP1]])
-; CHECK-NEXT: store ptr [[UNSAFE_STACK_PTR]], ptr @__safestack_unsafe_stack_ptr, align 8
-; CHECK-NEXT: [[R:%.*]] = musttail call i32 @foo(ptr [[A:%.*]])
-; CHECK-NEXT: [[T:%.*]] = bitcast i32 [[R]] to i32
-; CHECK-NEXT: ret i32 [[T]]
-;
- %x = alloca [10 x i8], align 1
- call void @alloca_test_use(ptr %x)
- %r = musttail call i32 @foo(ptr %a)
- %t = bitcast i32 %r to i32
- ret i32 %t
-}
More information about the llvm-commits
mailing list