[llvm] [TailRecElim] Handle discarded-call return conflicts with a shift accumulator (PR #221694)

via llvm-commits llvm-commits at lists.llvm.org
Mon Sep 7 03:09:17 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-llvm-transforms

Author: Federico Bruzzone (FedericoBruzzone)

<details>
<summary>Changes</summary>

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.



---
Full diff: https://github.com/llvm/llvm-project/pull/221694.diff


2 Files Affected:

- (modified) llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp (+28-10) 
- (modified) llvm/test/Transforms/TailCallElim/shl-accumulator-opt.ll (+36) 


``````````diff
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
+}

``````````

</details>


https://github.com/llvm/llvm-project/pull/221694


More information about the llvm-commits mailing list