[llvm] [TailRecElim] Handle discarded-call return conflicts with a shift accumulator (PR #221694)
Federico Bruzzone via llvm-commits
llvm-commits at lists.llvm.org
Mon Sep 7 03:08:36 PDT 2026
https://github.com/FedericoBruzzone created https://github.com/llvm/llvm-project/pull/221694
Fixes a follow-up miscompile from #181331, reported by @nikic in https://github.com/llvm/llvm-project/pull/181331#issuecomment-5509668449, similar but unrelated to #214503.
`findBaseCaseRetConstant` only scanned live `ret` instructions to check that a shift accumulator's base case is consistent everywhere. A sibling call site that discards its own recursive call and returns a fixed constant is eliminated earlier, so its `ret` is already gone by the time the scan runs. The scan missed it and `cleanupAndFinalize` would silently replace that
constant with the accumulator value.
This PR also check the false-value of already-inserted selects for consistency, bailing out the accumulator transform on conflict.
>From 1c4793e0e74726f0e382de0e225a3bc55168101e Mon Sep 17 00:00:00 2001
From: Federico Bruzzone <federico.bruzzone.i at gmail.com>
Date: Mon, 7 Sep 2026 12:02:23 +0200
Subject: [PATCH] [TailRecElim] Handle discarded-call return conflicts with a
shift accumulator
Signed-off-by: Federico Bruzzone <federico.bruzzone.i at gmail.com>
---
.../Scalar/TailRecursionElimination.cpp | 38 ++++++++++++++-----
.../TailCallElim/shl-accumulator-opt.ll | 36 ++++++++++++++++++
2 files changed, 64 insertions(+), 10 deletions(-)
diff --git a/llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp b/llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp
index b3df27433a9f5..99aafa3390eeb 100644
--- a/llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp
+++ b/llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp
@@ -50,6 +50,7 @@
//===----------------------------------------------------------------------===//
#include "llvm/Transforms/Scalar/TailRecursionElimination.h"
+#include "llvm/ADT/ArrayRef.h"
#include "llvm/ADT/STLExtras.h"
#include "llvm/ADT/SmallPtrSet.h"
#include "llvm/ADT/Statistic.h"
@@ -452,6 +453,11 @@ static bool isUnaryAccumulatorRecurrence(Instruction *I) {
// will be rewritten to return the accumulator, so all of them have to yield the
// same base-case constant. Return that constant, or nullptr on failure.
//
+// RetSelects are the selects already inserted for call sites eliminated via
+// the "found return value" mechanism instead of the accumulator one. Their
+// original `ret` is gone, so they'd otherwise be invisible to the scan above,
+// but they still have to agree on the same base-case constant.
+//
// FIXME: There is a room for improvement here in the future, e.g., consider
// non-constant values and multiple base cases -- e.g., we want to be able to
// handle code like:
@@ -462,10 +468,18 @@ static bool isUnaryAccumulatorRecurrence(Instruction *I) {
// return f(x-1) << 1;
// }
// ```
-static Constant *findBaseCaseRetConstant(Function &F,
- Instruction *AccRecInstr) {
+static Constant *findBaseCaseRetConstant(Function &F, Instruction *AccRecInstr,
+ ArrayRef<SelectInst *> RetSelects) {
Constant *BaseCaseVal = nullptr;
+ // Records C as the base-case constant the first time it's seen, and
+ // otherwise checks that it agrees with the one already on record.
+ auto AgreesWithBaseCase = [&](Constant *C) {
+ if (!BaseCaseVal)
+ BaseCaseVal = C;
+ return BaseCaseVal == C;
+ };
+
for (BasicBlock &BB : F) {
auto *RI = dyn_cast<ReturnInst>(BB.getTerminator());
if (!RI || !RI->getReturnValue())
@@ -483,12 +497,13 @@ static Constant *findBaseCaseRetConstant(Function &F,
// not eliminated) must be rejected: returning the accumulator in its place
// would drop that computation.
auto *C = dyn_cast<Constant>(RV);
- if (!C)
+ if (!C || !AgreesWithBaseCase(C))
return nullptr;
+ }
- if (!BaseCaseVal)
- BaseCaseVal = C;
- else if (BaseCaseVal != C)
+ for (SelectInst *SI : RetSelects) {
+ auto *C = dyn_cast<Constant>(SI->getFalseValue());
+ if (!C || !AgreesWithBaseCase(C))
return nullptr;
}
@@ -498,8 +513,9 @@ static Constant *findBaseCaseRetConstant(Function &F,
// This function checks whether the instruction I can be used
// to perform accumulator recursion elimination for the
// call instruction CI.
-static Constant *canTransformAccumulatorRecursion(Instruction *I,
- CallInst *CI) {
+static Constant *
+canTransformAccumulatorRecursion(Instruction *I, CallInst *CI,
+ ArrayRef<SelectInst *> RetSelects) {
bool IsUnaryAccumulatorRecurrence = isUnaryAccumulatorRecurrence(I);
if ((!I->isAssociative() || !I->isCommutative()) &&
!IsUnaryAccumulatorRecurrence)
@@ -517,7 +533,8 @@ static Constant *canTransformAccumulatorRecursion(Instruction *I,
// findTRECandidate guarantees CI is a recursive call to its own
// function, so scan the enclosing function for the base-case return.
- AccInitVal = findBaseCaseRetConstant(*CI->getFunction(), /*AccRecInstr=*/I);
+ AccInitVal = findBaseCaseRetConstant(*CI->getFunction(), /*AccRecInstr=*/I,
+ RetSelects);
if (!AccInitVal)
return nullptr;
} else {
@@ -815,7 +832,8 @@ bool TailRecursionEliminator::eliminateCall(CallInst *CI) {
// arithmetic operation that could be transformed using accumulator
// recursion elimination. Check to see if this is the case, and if so,
// remember which instruction accumulates for later.
- Constant *AccInitVal = canTransformAccumulatorRecursion(&*BBI, CI);
+ Constant *AccInitVal =
+ canTransformAccumulatorRecursion(&*BBI, CI, RetSelects);
if (AccPN || !AccInitVal)
return false; // We cannot eliminate the tail recursion!
diff --git a/llvm/test/Transforms/TailCallElim/shl-accumulator-opt.ll b/llvm/test/Transforms/TailCallElim/shl-accumulator-opt.ll
index d9d0308dc163f..9104d3d93e02c 100644
--- a/llvm/test/Transforms/TailCallElim/shl-accumulator-opt.ll
+++ b/llvm/test/Transforms/TailCallElim/shl-accumulator-opt.ll
@@ -239,3 +239,39 @@ other:
%res = sub i32 %call2, 3
ret i32 %res
}
+
+; Negative test: a sibling call site that discards its own recursive call and
+; unconditionally returns a fixed constant must still agree with the shift
+; accumulator's base case.
+; int f(int x) {
+; if (x == 0) return 1;
+; if (x == 3) { f(x - 1); return 2; }
+; return f(x - 1) << 1;
+; }
+define i32 @test_neg_discarded_call_conflicting_base(i32 %x) {
+; CHECK-LABEL: define i32 @test_neg_discarded_call_conflicting_base(
+; CHECK-NOT: accumulator.tr
+; CHECK: select {{.*}}, i32 2
+;
+entry:
+ %isbase = icmp eq i32 %x, 0
+ br i1 %isbase, label %base, label %rec
+
+rec:
+ %issp = icmp eq i32 %x, 3
+ br i1 %issp, label %special, label %normal
+
+special:
+ %d = sub i32 %x, 1
+ %c1 = tail call i32 @test_neg_discarded_call_conflicting_base(i32 %d)
+ ret i32 2
+
+normal:
+ %dec = sub i32 %x, 1
+ %c = tail call i32 @test_neg_discarded_call_conflicting_base(i32 %dec)
+ %shl = shl i32 %c, 1
+ ret i32 %shl
+
+base:
+ ret i32 1
+}
More information about the llvm-commits
mailing list