Same scheme as x86: have the arm64 BPF JIT open-code rcu_read_lock_trace()
and rcu_read_unlock_trace() in the trampoline, one reader from after the
callee-saved registers are stored to just before the original function
is called (covering __bpf_tramp_enter() and the fentry/fmod_ret programs)
and a second from just after it returns to just before those registers
are restored (covering the fexit programs and __bpf_tramp_exit()), with
the original function itself outside both and pinned by im->pcref.  The
second reader is entered before ip_after_call and the fmod_ret cbnz
lands past it still holding the first, so exactly one is held on every
path; trampolines without an original call get a single reader.

The sequence mirrors entry-ftrace.S: current via sp_el0,
trc_reader_nesting bumped, and for the outermost reader the SRCU-fast
per-CPU counter incremented LL/SC and the counter pointer stashed in
trc_reader_scp (plus the dmb when CONFIG_TASKS_TRACE_RCU_NO_MB is not
set).  x10-x15 are scratch at every emission point.  Nothing is emitted
on other configurations.

Suggested-by: Alexei Starovoitov <[email protected]>
Assisted-by: LLM
Signed-off-by: Josef Bacik <[email protected]>
---
 arch/arm64/net/bpf_jit_comp.c | 90 +++++++++++++++++++++++++++++++++++++++++++
 1 file changed, 90 insertions(+)

diff --git a/arch/arm64/net/bpf_jit_comp.c b/arch/arm64/net/bpf_jit_comp.c
index c18e005a41db..87f9d53caf2c 100644
--- a/arch/arm64/net/bpf_jit_comp.c
+++ b/arch/arm64/net/bpf_jit_comp.c
@@ -14,6 +14,7 @@
 #include <linux/filter.h>
 #include <linux/memory.h>
 #include <linux/printk.h>
+#include <linux/rcupdate_trace.h>
 #include <linux/slab.h>
 
 #include <asm/asm-extable.h>
@@ -2591,6 +2592,74 @@ static void emit_arena_arg_conv(struct jit_ctx *ctx, u8 
dst, u8 src, bool nullab
        emit(A64_SUB(0, dst, src, base_lo), ctx);
 }
 
+/*
+ * Open-coded rcu_read_lock_trace() / rcu_read_unlock_trace() for the
+ * trampoline, see CONFIG_HAVE_RCU_TRAMPOLINE_READERS and the equivalent macros
+ * in arch/arm64/kernel/entry-ftrace.S.  The SRCU-fast per-CPU increment is an
+ * LL/SC add on this CPU's counter; being migrated between reading the per-CPU
+ * offset and the store-exclusive only means another CPU's counter is
+ * incremented atomically instead, which SRCU sums over anyway.  Uses x10-x15,
+ * which are scratch at every emission point, and the flags.
+ */
+static void emit_trace_rcu_reader(struct jit_ctx *ctx, bool lock)
+{
+#ifdef CONFIG_TASKS_RCU_TRAMPOLINE_READERS
+       const int nesting = offsetof(struct task_struct, trc_reader_nesting);
+       const int scp = offsetof(struct task_struct, trc_reader_scp);
+       const bool mb = !IS_ENABLED(CONFIG_TASKS_TRACE_RCU_NO_MB);
+       const u8 tsk = A64_R(10), n = A64_R(11), tmp = A64_R(12);
+       const u8 ctr = A64_R(13), addr = A64_R(14), val = A64_R(15);
+
+       BUILD_BUG_ON(IS_ENABLED(CONFIG_NEED_SRCU_NMI_SAFE));
+       /* LDR/STR (immediate, unsigned offset) ranges */
+       BUILD_BUG_ON((nesting & 3) || nesting >= SZ_16K || (scp & 7) || scp >= 
SZ_32K);
+
+       emit(A64_MRS_SP_EL0(tsk), ctx);                 /* current */
+       emit(A64_LDR32I(n, tsk, nesting), ctx);
+       if (lock) {
+               emit(A64_ADD_I(0, tmp, n, 1), ctx);
+               emit(A64_STR32I(tmp, tsk, nesting), ctx);
+               /* interrupted a reader: done */
+               emit(A64_CBNZ(0, n, 12 + mb), ctx);
+               /* scp = rcu_tasks_trace_srcu_struct.srcu_ctrp; 
current->trc_reader_scp = scp */
+               emit_addr_mov_i64(ctr, 
(u64)&rcu_tasks_trace_srcu_struct.srcu_ctrp, ctx);
+               emit(A64_LDR64I(ctr, ctr, 0), ctx);
+               emit(A64_STR64I(ctr, tsk, scp), ctx);
+       } else {
+               emit(A64_SUB_I(0, n, n, 1), ctx);
+               /* still nested: just store the count */
+               emit(A64_CBNZ(0, n, 11 + mb), ctx);
+               /* outermost: pick up scp before an interrupt can see nesting 
== 0 */
+               emit(A64_LDR64I(ctr, tsk, scp), ctx);
+               emit(A64_STR32I(A64_ZR, tsk, nesting), ctx);
+               if (mb)
+                       emit(A64_DMB_ISH, ctx);
+       }
+       /* this_cpu_inc(scp->srcu_locks / srcu_unlocks) */
+       if (cpus_have_cap(ARM64_HAS_VIRT_HOST_EXTN))
+               emit(A64_MRS_TPIDR_EL2(addr), ctx);
+       else
+               emit(A64_MRS_TPIDR_EL1(addr), ctx);
+       emit(A64_ADD(1, addr, addr, ctr), ctx);
+       if (!lock)
+               emit(A64_ADD_I(1, addr, addr, offsetof(struct srcu_ctr, 
srcu_unlocks)), ctx);
+       emit(A64_LDXR(1, val, addr), ctx);
+       emit(A64_ADD_I(1, val, val, 1), ctx);
+       emit(A64_STXR(1, val, addr, tmp), ctx);
+       emit(A64_CBNZ(0, tmp, -3), ctx);
+       if (lock) {
+               if (mb)
+                       emit(A64_DMB_ISH, ctx);
+               /* 1: */
+       } else {
+               emit(A64_B(2), ctx);
+               /* 2: */
+               emit(A64_STR32I(n, tsk, nesting), ctx);
+               /* 3: */
+       }
+#endif
+}
+
 static void save_args(struct jit_ctx *ctx, int bargs_off, int oargs_off,
                      const struct btf_func_model *m, const struct arg_aux *a,
                      bool for_call_origin, bool is_struct_ops, u64 arena_base)
@@ -2854,6 +2923,16 @@ static int prepare_trampoline(struct jit_ctx *ctx, 
struct bpf_tramp_image *im,
        emit(A64_STR64I(A64_R(19), A64_SP, regs_off), ctx);
        emit(A64_STR64I(A64_R(20), A64_SP, regs_off + 8), ctx);
 
+       /*
+        * Tasks RCU keeps this image alive only while we are a Tasks Trace
+        * reader; the instructions before this point (and after the final
+        * unlock) are covered by the irq-exit IP check.  One reader spans
+        * __bpf_tramp_enter() and the fentry/fmod_ret progs, a second one the
+        * fexit progs and __bpf_tramp_exit(); the original function runs
+        * outside both, with the image pinned by im->pcref instead.
+        */
+       emit_trace_rcu_reader(ctx, true);
+
        if (flags & BPF_TRAMP_F_CALL_ORIG) {
                /* for the first pass, assume the worst case */
                if (!ctx->image)
@@ -2898,12 +2977,20 @@ static int prepare_trampoline(struct jit_ctx *ctx, 
struct bpf_tramp_image *im,
        if (flags & BPF_TRAMP_F_CALL_ORIG) {
                /* the original func takes kernel addresses, never converted 
ones */
                save_args(ctx, bargs_off, oargs_off, m, a, true, is_struct_ops, 
0);
+               emit_trace_rcu_reader(ctx, false);
                /* call original func */
                emit(A64_LDR64I(A64_R(10), A64_SP, retaddr_off), ctx);
                emit(A64_ADR(A64_LR, AARCH64_INSN_SIZE * 2), ctx);
                emit(A64_RET(A64_R(10)), ctx);
                /* store return value */
                emit(A64_STR64I(A64_R(0), A64_SP, retval_off), ctx);
+               /*
+                * Second reader.  Taken before ip_after_call so that the branch
+                * to the epilogue patched in at teardown is inside it too; the
+                * fmod_ret early exit lands past this still holding the first
+                * reader, so either way exactly one is held.
+                */
+               emit_trace_rcu_reader(ctx, true);
                /* reserve a nop for bpf_tramp_image_put */
                im->ip_after_call = ctx->ro_image + ctx->idx;
                emit(A64_NOP, ctx);
@@ -2945,6 +3032,9 @@ static int prepare_trampoline(struct jit_ctx *ctx, struct 
bpf_tramp_image *im,
        if (flags & BPF_TRAMP_F_RESTORE_REGS)
                restore_args(ctx, bargs_off, a->regs_for_args);
 
+       /* Remaining instructions are covered by the irq-exit IP check. */
+       emit_trace_rcu_reader(ctx, false);
+
        /* restore callee saved register x19 and x20 */
        emit(A64_LDR64I(A64_R(19), A64_SP, regs_off), ctx);
        emit(A64_LDR64I(A64_R(20), A64_SP, regs_off + 8), ctx);

-- 
2.55.0


Reply via email to