[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