[llvm] [RISCV] Support tail calls for functions with sret parameters (PR #223999)

via llvm-commits llvm-commits at lists.llvm.org
Wed Sep 16 05:46:57 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-backend-risc-v

Author: renndong

<details>
<summary>Changes</summary>

This PR allows RISC-V tail calls from functions with an `sret`
parameter in a conservative set of cases.

RISC-V currently rejects a tail call whenever either the caller or the
callee has an `sret` parameter. This unnecessarily prevents tail-call
optimization.

For example:

```
%struct.Buffer = type { [7 x i32] }

; Function Attrs: mustprogress noinline nounwind uwtable
define void @<!-- -->foo(ptr sret(%struct.Buffer) align 4 %agg.result) {
entry:
  tail call void @<!-- -->bar(ptr %agg.result, i32 1, i32 2, i32 3, i32 4, i32 5, i32 6)
  ret void
}
```

The caller-provided return buffer is passed to `foo` in a0. The same
pointer is then passed as the bar's ordinary first argument, 
which also uses a0.  Before this change, RISC-V emits a normal
call with a prologue and epilogue:

```
      addi	sp, sp, -16
	.cfi_def_cfa_offset 16
	sd	ra, 8(sp)                       # 8-byte Folded Spill
	.cfi_offset ra, -8
	li	a1, 1
	li	a2, 2
	li	a3, 3
	li	a4, 4
	li	a5, 5
	li	a6, 6
	call	bar
	ld	ra, 8(sp)                       # 8-byte Folded Reload
	.cfi_restore ra
	addi	sp, sp, 16
	.cfi_def_cfa_offset 0
	ret
```

With this change, the call can be emitted directly as:

```
	li	a1, 1
	li	a2, 2
	li	a3, 3
	li	a4, 4
	li	a5, 5
	li	a6, 6
	tail	bar
```

This change is conservative and only considers cases where the caller’s sret pointer is passed to callees.


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


3 Files Affected:

- (modified) llvm/lib/Target/RISCV/RISCVISelLowering.cpp (+16-5) 
- (added) llvm/test/CodeGen/RISCV/tail-calls-sret.ll (+153) 
- (modified) llvm/test/CodeGen/RISCV/tail-calls.ll (+6-31) 


``````````diff
diff --git a/llvm/lib/Target/RISCV/RISCVISelLowering.cpp b/llvm/lib/Target/RISCV/RISCVISelLowering.cpp
index d05e1f85afb69..75c96b704b45e 100644
--- a/llvm/lib/Target/RISCV/RISCVISelLowering.cpp
+++ b/llvm/lib/Target/RISCV/RISCVISelLowering.cpp
@@ -27587,12 +27587,23 @@ bool RISCVTargetLowering::isEligibleForTailCallOptimization(
     if (VA.getLocInfo() == CCValAssign::Indirect)
       return false;
 
-  // Do not tail call opt if either caller or callee uses struct return
-  // semantics.
-  auto IsCallerStructRet = Caller.hasStructRetAttr();
+  // If the callee has an sret parameter, conservatively require it to receive
+  // the caller's sret pointer. If only the caller has an sret parameter, treat
+  // that pointer like an ordinary pointer when passing call arguments.
+  // TODO: Support other sret buffers that outlive the caller, such as globals.
+  auto IsCallerStructRet =
+      !Caller.arg_empty() && Caller.getArg(0)->hasStructRetAttr();
   auto IsCalleeStructRet = Outs.empty() ? false : Outs[0].Flags.isSRet();
-  if (IsCallerStructRet || IsCalleeStructRet)
-    return false;
+  if (IsCalleeStructRet) {
+    // Do not allow the tail call if the caller has no sret parameter.
+    if (!IsCallerStructRet)
+      return false;
+    // RISC-V passes the sret pointer as the first argument in a0. Require the
+    // callee's sret argument to be the caller's incoming sret pointer.
+    if (!CLI.CB || CLI.CB->arg_empty() ||
+        CLI.CB->getArgOperand(0) != Caller.getArg(0))
+      return false;
+  }
 
   // The callee has to preserve all registers the caller needs to preserve.
   const RISCVRegisterInfo *TRI = Subtarget.getRegisterInfo();
diff --git a/llvm/test/CodeGen/RISCV/tail-calls-sret.ll b/llvm/test/CodeGen/RISCV/tail-calls-sret.ll
new file mode 100644
index 0000000000000..8256916f3c293
--- /dev/null
+++ b/llvm/test/CodeGen/RISCV/tail-calls-sret.ll
@@ -0,0 +1,153 @@
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py UTC_ARGS: --version 5
+; RUN: llc -mtriple=riscv32 -verify-machineinstrs < %s | FileCheck %s --check-prefix=RV32
+; RUN: llc -mtriple=riscv64 -verify-machineinstrs < %s | FileCheck %s --check-prefix=RV64
+
+%struct.Buffer = type { [7 x i64] }
+
+declare void @forward_sret(ptr sret(%struct.Buffer), i64)
+declare void @use_pointer(ptr)
+declare void @use_as_second_arg(i32, ptr)
+
+; The caller's sret pointer can be passed as an ordinary first argument, such
+; as the `this` pointer of a C++ constructor.
+define void @caller_sret_as_first_arg(ptr noalias sret(%struct.Buffer) %result) {
+; RV32-LABEL: caller_sret_as_first_arg:
+; RV32:       # %bb.0: # %entry
+; RV32-NEXT:    tail use_pointer
+;
+; RV64-LABEL: caller_sret_as_first_arg:
+; RV64:       # %bb.0: # %entry
+; RV64-NEXT:    tail use_pointer
+entry:
+  tail call void @use_pointer(ptr %result)
+  ret void
+}
+
+; Matching caller and callee sret semantics can reuse the incoming
+; return buffer.
+define void @forward_result(ptr noalias sret(%struct.Buffer) %result, i64 %tag) {
+; RV32-LABEL: forward_result:
+; RV32:       # %bb.0: # %entry
+; RV32-NEXT:    tail forward_sret
+;
+; RV64-LABEL: forward_result:
+; RV64:       # %bb.0: # %entry
+; RV64-NEXT:    tail forward_sret
+entry:
+  tail call void @forward_sret(ptr sret(%struct.Buffer) %result, i64 %tag)
+  ret void
+}
+
+; A caller's sret pointer does not need to be forwarded to a callee without
+; sret semantics.
+define void @caller_sret_unused(ptr noalias sret(%struct.Buffer) %result, ptr %other) {
+; RV32-LABEL: caller_sret_unused:
+; RV32:       # %bb.0: # %entry
+; RV32-NEXT:    mv a0, a1
+; RV32-NEXT:    tail use_pointer
+;
+; RV64-LABEL: caller_sret_unused:
+; RV64:       # %bb.0: # %entry
+; RV64-NEXT:    mv a0, a1
+; RV64-NEXT:    tail use_pointer
+entry:
+  tail call void @use_pointer(ptr %other)
+  ret void
+}
+
+; The caller's sret pointer can also be passed as an ordinary non-first
+; argument.
+define void @caller_sret_as_second_arg(ptr noalias sret(%struct.Buffer) %result, i32 %tag) {
+; RV32-LABEL: caller_sret_as_second_arg:
+; RV32:       # %bb.0: # %entry
+; RV32-NEXT:    mv a2, a0
+; RV32-NEXT:    mv a0, a1
+; RV32-NEXT:    mv a1, a2
+; RV32-NEXT:    tail use_as_second_arg
+;
+; RV64-LABEL: caller_sret_as_second_arg:
+; RV64:       # %bb.0: # %entry
+; RV64-NEXT:    mv a2, a0
+; RV64-NEXT:    mv a0, a1
+; RV64-NEXT:    mv a1, a2
+; RV64-NEXT:    tail use_as_second_arg
+entry:
+  tail call void @use_as_second_arg(i32 %tag, ptr %result)
+  ret void
+}
+
+; Do not tail call when caller and callee use different sret buffers.
+define void @sret_not_forwarded(ptr noalias sret(%struct.Buffer) %result, ptr %other, i64 %tag) {
+; RV32-LABEL: sret_not_forwarded:
+; RV32:       # %bb.0: # %entry
+; RV32-NEXT:    addi sp, sp, -16
+; RV32-NEXT:    .cfi_def_cfa_offset 16
+; RV32-NEXT:    sw ra, 12(sp) # 4-byte Folded Spill
+; RV32-NEXT:    .cfi_offset ra, -4
+; RV32-NEXT:    mv a0, a1
+; RV32-NEXT:    mv a1, a2
+; RV32-NEXT:    mv a2, a3
+; RV32-NEXT:    call forward_sret
+; RV32-NEXT:    lw ra, 12(sp) # 4-byte Folded Reload
+; RV32-NEXT:    .cfi_restore ra
+; RV32-NEXT:    addi sp, sp, 16
+; RV32-NEXT:    .cfi_def_cfa_offset 0
+; RV32-NEXT:    ret
+;
+; RV64-LABEL: sret_not_forwarded:
+; RV64:       # %bb.0: # %entry
+; RV64-NEXT:    addi sp, sp, -16
+; RV64-NEXT:    .cfi_def_cfa_offset 16
+; RV64-NEXT:    sd ra, 8(sp) # 8-byte Folded Spill
+; RV64-NEXT:    .cfi_offset ra, -8
+; RV64-NEXT:    mv a0, a1
+; RV64-NEXT:    mv a1, a2
+; RV64-NEXT:    call forward_sret
+; RV64-NEXT:    ld ra, 8(sp) # 8-byte Folded Reload
+; RV64-NEXT:    .cfi_restore ra
+; RV64-NEXT:    addi sp, sp, 16
+; RV64-NEXT:    .cfi_def_cfa_offset 0
+; RV64-NEXT:    ret
+entry:
+  tail call void @forward_sret(ptr sret(%struct.Buffer) %other, i64 %tag)
+  ret void
+}
+
+; Do not tail call when the callee's sret pointer refers to the caller's local
+; stack frame.
+define void @local_sret_buffer(i64 %tag) {
+; RV32-LABEL: local_sret_buffer:
+; RV32:       # %bb.0: # %entry
+; RV32-NEXT:    addi sp, sp, -64
+; RV32-NEXT:    .cfi_def_cfa_offset 64
+; RV32-NEXT:    sw ra, 60(sp) # 4-byte Folded Spill
+; RV32-NEXT:    .cfi_offset ra, -4
+; RV32-NEXT:    mv a2, a1
+; RV32-NEXT:    mv a1, a0
+; RV32-NEXT:    mv a0, sp
+; RV32-NEXT:    call forward_sret
+; RV32-NEXT:    lw ra, 60(sp) # 4-byte Folded Reload
+; RV32-NEXT:    .cfi_restore ra
+; RV32-NEXT:    addi sp, sp, 64
+; RV32-NEXT:    .cfi_def_cfa_offset 0
+; RV32-NEXT:    ret
+;
+; RV64-LABEL: local_sret_buffer:
+; RV64:       # %bb.0: # %entry
+; RV64-NEXT:    addi sp, sp, -64
+; RV64-NEXT:    .cfi_def_cfa_offset 64
+; RV64-NEXT:    sd ra, 56(sp) # 8-byte Folded Spill
+; RV64-NEXT:    .cfi_offset ra, -8
+; RV64-NEXT:    mv a1, a0
+; RV64-NEXT:    mv a0, sp
+; RV64-NEXT:    call forward_sret
+; RV64-NEXT:    ld ra, 56(sp) # 8-byte Folded Reload
+; RV64-NEXT:    .cfi_restore ra
+; RV64-NEXT:    addi sp, sp, 64
+; RV64-NEXT:    .cfi_def_cfa_offset 0
+; RV64-NEXT:    ret
+entry:
+  %local = alloca %struct.Buffer, align 8
+  tail call void @forward_sret(ptr sret(%struct.Buffer) %local, i64 %tag)
+  ret void
+}
diff --git a/llvm/test/CodeGen/RISCV/tail-calls.ll b/llvm/test/CodeGen/RISCV/tail-calls.ll
index e43acfa56f3fd..bd12a9498d22e 100644
--- a/llvm/test/CodeGen/RISCV/tail-calls.ll
+++ b/llvm/test/CodeGen/RISCV/tail-calls.ll
@@ -1110,63 +1110,38 @@ entry:
   ret void
 }
 
-; Do not tail call optimize if caller uses structret semantics.
+; A caller using structret semantics does not prevent tail call optimization.
 declare void @callee_nostruct()
 define void @caller_struct(ptr sret(%struct.A) %a) nounwind {
 ; CHECK-LABEL: caller_struct:
 ; CHECK:       # %bb.0: # %entry
-; CHECK-NEXT:    addi sp, sp, -16
-; CHECK-NEXT:    sw ra, 12(sp) # 4-byte Folded Spill
-; CHECK-NEXT:    call callee_nostruct
-; CHECK-NEXT:    lw ra, 12(sp) # 4-byte Folded Reload
-; CHECK-NEXT:    addi sp, sp, 16
-; CHECK-NEXT:    ret
+; CHECK-NEXT:    tail callee_nostruct
 ;
 ; CHECK-CF-RV32-LABEL: caller_struct:
 ; CHECK-CF-RV32:       # %bb.0: # %entry
 ; CHECK-CF-RV32-NEXT:    lpad 0
-; CHECK-CF-RV32-NEXT:    addi sp, sp, -16
-; CHECK-CF-RV32-NEXT:    sw ra, 12(sp) # 4-byte Folded Spill
-; CHECK-CF-RV32-NEXT:    call callee_nostruct
-; CHECK-CF-RV32-NEXT:    lw ra, 12(sp) # 4-byte Folded Reload
-; CHECK-CF-RV32-NEXT:    addi sp, sp, 16
-; CHECK-CF-RV32-NEXT:    ret
+; CHECK-CF-RV32-NEXT:    tail callee_nostruct, t2
 ;
 ; CHECK-CF-RV64-LABEL: caller_struct:
 ; CHECK-CF-RV64:       # %bb.0: # %entry
 ; CHECK-CF-RV64-NEXT:    lpad 0
-; CHECK-CF-RV64-NEXT:    addi sp, sp, -16
-; CHECK-CF-RV64-NEXT:    sd ra, 8(sp) # 8-byte Folded Spill
-; CHECK-CF-RV64-NEXT:    call callee_nostruct
-; CHECK-CF-RV64-NEXT:    ld ra, 8(sp) # 8-byte Folded Reload
-; CHECK-CF-RV64-NEXT:    addi sp, sp, 16
-; CHECK-CF-RV64-NEXT:    ret
+; CHECK-CF-RV64-NEXT:    tail callee_nostruct, t2
 ;
 ; CHECK-CF-RV32-LARGE-LABEL: caller_struct:
 ; CHECK-CF-RV32-LARGE:       # %bb.0: # %entry
 ; CHECK-CF-RV32-LARGE-NEXT:    lpad 0
-; CHECK-CF-RV32-LARGE-NEXT:    addi sp, sp, -16
-; CHECK-CF-RV32-LARGE-NEXT:    sw ra, 12(sp) # 4-byte Folded Spill
 ; CHECK-CF-RV32-LARGE-NEXT:  .Lpcrel_hi15:
 ; CHECK-CF-RV32-LARGE-NEXT:    auipc a0, %pcrel_hi(.LCPI12_0)
 ; CHECK-CF-RV32-LARGE-NEXT:    lw t2, %pcrel_lo(.Lpcrel_hi15)(a0)
-; CHECK-CF-RV32-LARGE-NEXT:    jalr t2
-; CHECK-CF-RV32-LARGE-NEXT:    lw ra, 12(sp) # 4-byte Folded Reload
-; CHECK-CF-RV32-LARGE-NEXT:    addi sp, sp, 16
-; CHECK-CF-RV32-LARGE-NEXT:    ret
+; CHECK-CF-RV32-LARGE-NEXT:    jr t2
 ;
 ; CHECK-CF-RV64-LARGE-LABEL: caller_struct:
 ; CHECK-CF-RV64-LARGE:       # %bb.0: # %entry
 ; CHECK-CF-RV64-LARGE-NEXT:    lpad 0
-; CHECK-CF-RV64-LARGE-NEXT:    addi sp, sp, -16
-; CHECK-CF-RV64-LARGE-NEXT:    sd ra, 8(sp) # 8-byte Folded Spill
 ; CHECK-CF-RV64-LARGE-NEXT:  .Lpcrel_hi15:
 ; CHECK-CF-RV64-LARGE-NEXT:    auipc a0, %pcrel_hi(.LCPI12_0)
 ; CHECK-CF-RV64-LARGE-NEXT:    ld t2, %pcrel_lo(.Lpcrel_hi15)(a0)
-; CHECK-CF-RV64-LARGE-NEXT:    jalr t2
-; CHECK-CF-RV64-LARGE-NEXT:    ld ra, 8(sp) # 8-byte Folded Reload
-; CHECK-CF-RV64-LARGE-NEXT:    addi sp, sp, 16
-; CHECK-CF-RV64-LARGE-NEXT:    ret
+; CHECK-CF-RV64-LARGE-NEXT:    jr t2
 entry:
   tail call void @callee_nostruct()
   ret void

``````````

</details>


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


More information about the llvm-commits mailing list