BPF_CGROUP_RUN_PROG_XXX() macros guard __cgroup_bpf_run_filter_XXX() with cgroup_bpf_enabled(). However, even when no SOCK_OPS prog is attached, callers still initialise struct bpf_sock_ops_kern (memset(), etc.) or evaluate BPF_SOCK_OPS_TEST_FLAG(), which loads tp->bpf_sock_ops_cb_flags from a cold cacheline near the end of struct tcp_sock. Similar to bpf_tcp_ops, let's check cgroup_bpf_enabled() before BPF_SOCK_OPS_TEST_FLAG() and struct bpf_sock_ops_kern setup, and rename BPF_CGROUP_RUN_PROG_SOCK_OPS{,_SK}() with __ prefix. Since all direct callers of tcp_call_bpf() pass 0 and NULL for nargs and args, they can be folded into the new tcp_call_bpf() macro. All callers of tcp_call_bpf_{2,3}arg() check BPF_SOCK_OPS_TEST_FLAG() and do not need the return value. These are replaced with the new tcp_call_bpf_flag() macro. Signed-off-by: Kuniyuki Iwashima --- include/linux/bpf-cgroup.h | 36 ++++++++++-------------- include/net/tcp.h | 56 ++++++++++++++++---------------------- net/ipv4/af_inet.c | 2 +- net/ipv4/tcp.c | 3 +- net/ipv4/tcp_input.c | 21 ++++++++------ net/ipv4/tcp_nv.c | 2 +- net/ipv4/tcp_output.c | 31 +++++++++------------ net/ipv4/tcp_timer.c | 6 ++-- 8 files changed, 69 insertions(+), 88 deletions(-) diff --git a/include/linux/bpf-cgroup.h b/include/linux/bpf-cgroup.h index 8a75a6cd7309..02ef7f899b3d 100644 --- a/include/linux/bpf-cgroup.h +++ b/include/linux/bpf-cgroup.h @@ -345,27 +345,19 @@ static inline bool cgroup_bpf_sock_enabled(struct sock *sk, * calling bpf_setsockopt on listener-sk will not make sense anyway, * so passing 'sock_ops->sk == req_sk' to the bpf prog is appropriate here. */ -#define BPF_CGROUP_RUN_PROG_SOCK_OPS_SK(sock_ops, sk) \ -({ \ - int __ret = 0; \ - if (cgroup_bpf_enabled(CGROUP_SOCK_OPS)) \ - __ret = __cgroup_bpf_run_filter_sock_ops(sk, \ - sock_ops, \ - CGROUP_SOCK_OPS); \ - __ret; \ -}) - -#define BPF_CGROUP_RUN_PROG_SOCK_OPS(sock_ops) \ -({ \ - int __ret = 0; \ - if (cgroup_bpf_enabled(CGROUP_SOCK_OPS) && (sock_ops)->sk) { \ - typeof(sk) __sk = sk_to_full_sk((sock_ops)->sk); \ - if (__sk && sk_fullsock(__sk)) \ - __ret = __cgroup_bpf_run_filter_sock_ops(__sk, \ - sock_ops, \ - CGROUP_SOCK_OPS); \ - } \ - __ret; \ +#define __BPF_CGROUP_RUN_PROG_SOCK_OPS_SK(sock_ops, sk) \ + __cgroup_bpf_run_filter_sock_ops(sk, sock_ops, \ + CGROUP_SOCK_OPS) + +#define __BPF_CGROUP_RUN_PROG_SOCK_OPS(sock_ops) \ +({ \ + int __ret = 0; \ + typeof(sk) __sk = sk_to_full_sk((sock_ops)->sk); \ + if (__sk && sk_fullsock(__sk)) \ + __ret = __cgroup_bpf_run_filter_sock_ops(__sk, \ + sock_ops, \ + CGROUP_SOCK_OPS); \ + __ret; \ }) #define BPF_CGROUP_RUN_PROG_DEVICE_CGROUP(atype, major, minor, access) \ @@ -529,7 +521,7 @@ static inline int cgroup_bpf_struct_ops_attach(struct bpf_map *map, #define BPF_CGROUP_RUN_PROG_UDP4_RECVMSG_LOCK(sk, uaddr, uaddrlen) ({ 0; }) #define BPF_CGROUP_RUN_PROG_UDP6_RECVMSG_LOCK(sk, uaddr, uaddrlen) ({ 0; }) #define BPF_CGROUP_RUN_PROG_UNIX_RECVMSG_LOCK(sk, uaddr, uaddrlen) ({ 0; }) -#define BPF_CGROUP_RUN_PROG_SOCK_OPS(sock_ops) ({ 0; }) +#define __BPF_CGROUP_RUN_PROG_SOCK_OPS(sock_ops) ({ 0; }) #define BPF_CGROUP_RUN_PROG_DEVICE_CGROUP(atype, major, minor, access) ({ 0; }) #define BPF_CGROUP_RUN_PROG_SYSCTL(head,table,write,buf,count,pos) ({ 0; }) #define BPF_CGROUP_RUN_PROG_GETSOCKOPT(sock, level, optname, optval, \ diff --git a/include/net/tcp.h b/include/net/tcp.h index b7c0f1a8797a..85b4bfe963d3 100644 --- a/include/net/tcp.h +++ b/include/net/tcp.h @@ -2891,7 +2891,7 @@ static inline void bpf_skops_init_skb(struct bpf_sock_ops_kern *skops, * program loaded). */ #ifdef CONFIG_BPF -static inline int tcp_call_bpf(struct sock *sk, int op, u32 nargs, u32 *args) +static inline int __tcp_call_bpf(struct sock *sk, int op, u32 nargs, u32 *args) { struct bpf_sock_ops_kern sock_ops; int ret; @@ -2908,7 +2908,7 @@ static inline int tcp_call_bpf(struct sock *sk, int op, u32 nargs, u32 *args) if (nargs > 0) memcpy(sock_ops.args, args, nargs * sizeof(*args)); - ret = BPF_CGROUP_RUN_PROG_SOCK_OPS(&sock_ops); + ret = __BPF_CGROUP_RUN_PROG_SOCK_OPS(&sock_ops); if (ret == 0) ret = sock_ops.reply; else @@ -2916,20 +2916,23 @@ static inline int tcp_call_bpf(struct sock *sk, int op, u32 nargs, u32 *args) return ret; } -static inline int tcp_call_bpf_2arg(struct sock *sk, int op, u32 arg1, u32 arg2) -{ - u32 args[2] = {arg1, arg2}; - - return tcp_call_bpf(sk, op, 2, args); -} - -static inline int tcp_call_bpf_3arg(struct sock *sk, int op, u32 arg1, u32 arg2, - u32 arg3) -{ - u32 args[3] = {arg1, arg2, arg3}; +#define tcp_call_bpf(sk, op) \ +({ \ + int __ret = 0; \ + if (cgroup_bpf_enabled(CGROUP_SOCK_OPS)) { \ + __ret = __tcp_call_bpf(sk, op, 0, NULL); \ + } \ + __ret; \ +}) - return tcp_call_bpf(sk, op, 3, args); -} +#define tcp_call_bpf_flag(sk, op, ...) \ +do { \ + if (cgroup_bpf_enabled(CGROUP_SOCK_OPS) && \ + BPF_SOCK_OPS_TEST_FLAG(tcp_sk(sk), op ## _FLAG)) { \ + u32 __args[] = { __VA_ARGS__ }; \ + __tcp_call_bpf(sk, op, ARRAY_SIZE(__args), __args); \ + } \ +} while (0) static inline void tcp_clear_sock_ops_cb_flags(struct sock *sk) { @@ -2938,21 +2941,12 @@ static inline void tcp_clear_sock_ops_cb_flags(struct sock *sk) } #else -static inline int tcp_call_bpf(struct sock *sk, int op, u32 nargs, u32 *args) -{ - return -EPERM; -} - -static inline int tcp_call_bpf_2arg(struct sock *sk, int op, u32 arg1, u32 arg2) +static inline int tcp_call_bpf(struct sock *sk, int op) { return -EPERM; } -static inline int tcp_call_bpf_3arg(struct sock *sk, int op, u32 arg1, u32 arg2, - u32 arg3) -{ - return -EPERM; -} +#define tcp_call_bpf_flag(sk, op, ...) do { } while (0) static inline void tcp_clear_sock_ops_cb_flags(struct sock *sk) { @@ -3136,7 +3130,7 @@ static inline u32 tcp_timeout_init(struct sock *sk) { int timeout; - timeout = tcp_call_bpf(sk, BPF_SOCK_OPS_TIMEOUT_INIT, 0, NULL); + timeout = tcp_call_bpf(sk, BPF_SOCK_OPS_TIMEOUT_INIT); timeout = bpf_tcp_ops_call_int(timeout_init, timeout, sk); if (timeout <= 0) timeout = TCP_TIMEOUT_INIT; @@ -3147,7 +3141,7 @@ static inline u32 tcp_rwnd_init_bpf(struct sock *sk) { int rwnd; - rwnd = tcp_call_bpf(sk, BPF_SOCK_OPS_RWND_INIT, 0, NULL); + rwnd = tcp_call_bpf(sk, BPF_SOCK_OPS_RWND_INIT); rwnd = bpf_tcp_ops_call_int(rwnd_init, rwnd, sk); if (rwnd < 0) rwnd = 0; @@ -3156,14 +3150,12 @@ static inline u32 tcp_rwnd_init_bpf(struct sock *sk) static inline bool tcp_bpf_ca_needs_ecn(struct sock *sk) { - return (tcp_call_bpf(sk, BPF_SOCK_OPS_NEEDS_ECN, 0, NULL) == 1); + return (tcp_call_bpf(sk, BPF_SOCK_OPS_NEEDS_ECN) == 1); } static inline void tcp_bpf_rtt(struct sock *sk, long mrtt, u32 srtt) { - if (BPF_SOCK_OPS_TEST_FLAG(tcp_sk(sk), BPF_SOCK_OPS_RTT_CB_FLAG)) - tcp_call_bpf_2arg(sk, BPF_SOCK_OPS_RTT_CB, mrtt, srtt); - + tcp_call_bpf_flag(sk, BPF_SOCK_OPS_RTT_CB, mrtt, srtt); bpf_tcp_ops_call_flag(rtt, RTT, sk, mrtt, srtt); } diff --git a/net/ipv4/af_inet.c b/net/ipv4/af_inet.c index cdcfc7d3c6d2..035c83338e0f 100644 --- a/net/ipv4/af_inet.c +++ b/net/ipv4/af_inet.c @@ -226,7 +226,7 @@ int __inet_listen_sk(struct sock *sk, int backlog) if (err) return err; - tcp_call_bpf(sk, BPF_SOCK_OPS_TCP_LISTEN_CB, 0, NULL); + tcp_call_bpf(sk, BPF_SOCK_OPS_TCP_LISTEN_CB); bpf_tcp_ops_call(listen, sk); } return 0; diff --git a/net/ipv4/tcp.c b/net/ipv4/tcp.c index fa69961c47d3..5d9d3bcde8f7 100644 --- a/net/ipv4/tcp.c +++ b/net/ipv4/tcp.c @@ -2994,8 +2994,7 @@ void tcp_set_state(struct sock *sk, int state) */ BTF_TYPE_EMIT_ENUM(BPF_TCP_ESTABLISHED); - if (BPF_SOCK_OPS_TEST_FLAG(tcp_sk(sk), BPF_SOCK_OPS_STATE_CB_FLAG)) - tcp_call_bpf_2arg(sk, BPF_SOCK_OPS_STATE_CB, oldstate, state); + tcp_call_bpf_flag(sk, BPF_SOCK_OPS_STATE_CB, oldstate, state); bpf_tcp_ops_call(set_state, sk, state); switch (state) { diff --git a/net/ipv4/tcp_input.c b/net/ipv4/tcp_input.c index 4478d3f3d4b0..f374257013b1 100644 --- a/net/ipv4/tcp_input.c +++ b/net/ipv4/tcp_input.c @@ -146,14 +146,16 @@ EXPORT_SYMBOL_GPL(clean_acked_data_flush); #ifdef CONFIG_CGROUP_BPF static void bpf_skops_parse_hdr(struct sock *sk, struct sk_buff *skb) { - bool unknown_opt = tcp_sk(sk)->rx_opt.saw_unknown && - BPF_SOCK_OPS_TEST_FLAG(tcp_sk(sk), - BPF_SOCK_OPS_PARSE_UNKNOWN_HDR_OPT_CB_FLAG); - bool parse_all_opt = BPF_SOCK_OPS_TEST_FLAG(tcp_sk(sk), - BPF_SOCK_OPS_PARSE_ALL_HDR_OPT_CB_FLAG); struct bpf_sock_ops_kern sock_ops; + const struct tcp_sock *tp; + + if (!cgroup_bpf_enabled(CGROUP_SOCK_OPS)) + return; - if (likely(!unknown_opt && !parse_all_opt)) + tp = tcp_sk(sk); + if (!(tp->rx_opt.saw_unknown && + BPF_SOCK_OPS_TEST_FLAG(tp, BPF_SOCK_OPS_PARSE_UNKNOWN_HDR_OPT_CB_FLAG)) && + !BPF_SOCK_OPS_TEST_FLAG(tp, BPF_SOCK_OPS_PARSE_ALL_HDR_OPT_CB_FLAG)) return; /* The skb will be handled in the @@ -176,7 +178,7 @@ static void bpf_skops_parse_hdr(struct sock *sk, struct sk_buff *skb) sock_ops.sk = sk; bpf_skops_init_skb(&sock_ops, skb, tcp_hdrlen(skb)); - BPF_CGROUP_RUN_PROG_SOCK_OPS(&sock_ops); + __BPF_CGROUP_RUN_PROG_SOCK_OPS(&sock_ops); } static void bpf_skops_established(struct sock *sk, int bpf_op, @@ -184,6 +186,9 @@ static void bpf_skops_established(struct sock *sk, int bpf_op, { struct bpf_sock_ops_kern sock_ops; + if (!cgroup_bpf_enabled(CGROUP_SOCK_OPS)) + return; + sock_owned_by_me(sk); memset(&sock_ops, 0, offsetof(struct bpf_sock_ops_kern, temp)); @@ -195,7 +200,7 @@ static void bpf_skops_established(struct sock *sk, int bpf_op, if (skb) bpf_skops_init_skb(&sock_ops, skb, tcp_hdrlen(skb)); - BPF_CGROUP_RUN_PROG_SOCK_OPS(&sock_ops); + __BPF_CGROUP_RUN_PROG_SOCK_OPS(&sock_ops); } #else static void bpf_skops_parse_hdr(struct sock *sk, struct sk_buff *skb) diff --git a/net/ipv4/tcp_nv.c b/net/ipv4/tcp_nv.c index f345897a68df..7b0dae23d9aa 100644 --- a/net/ipv4/tcp_nv.c +++ b/net/ipv4/tcp_nv.c @@ -146,7 +146,7 @@ static void tcpnv_init(struct sock *sk) * within a datacenter, where we have reasonable estimates of * RTTs */ - base_rtt = tcp_call_bpf(sk, BPF_SOCK_OPS_BASE_RTT, 0, NULL); + base_rtt = tcp_call_bpf(sk, BPF_SOCK_OPS_BASE_RTT); if (base_rtt > 0) { ca->nv_base_rtt = base_rtt; ca->nv_lower_bound_rtt = (base_rtt * 205) >> 8; /* 80% */ diff --git a/net/ipv4/tcp_output.c b/net/ipv4/tcp_output.c index 8770f3084efe..db4fd1825d99 100644 --- a/net/ipv4/tcp_output.c +++ b/net/ipv4/tcp_output.c @@ -476,8 +476,9 @@ static u32 bpf_skops_hdr_opt_len(struct sock *sk, struct sk_buff *skb, struct bpf_sock_ops_kern sock_ops; int err; - if (likely(!BPF_SOCK_OPS_TEST_FLAG(tcp_sk(sk), - BPF_SOCK_OPS_WRITE_HDR_OPT_CB_FLAG)) || + if (!cgroup_bpf_enabled(CGROUP_SOCK_OPS) || + !BPF_SOCK_OPS_TEST_FLAG(tcp_sk(sk), + BPF_SOCK_OPS_WRITE_HDR_OPT_CB_FLAG)|| !remaining) return remaining; @@ -518,7 +519,7 @@ static u32 bpf_skops_hdr_opt_len(struct sock *sk, struct sk_buff *skb, if (skb) bpf_skops_init_skb(&sock_ops, skb, 0); - err = BPF_CGROUP_RUN_PROG_SOCK_OPS_SK(&sock_ops, sk); + err = __BPF_CGROUP_RUN_PROG_SOCK_OPS_SK(&sock_ops, sk); if (err || sock_ops.remaining_opt_len == remaining) return remaining; @@ -543,7 +544,8 @@ static void bpf_skops_write_hdr_opt(struct sock *sk, struct sk_buff *skb, first_opt_off = tcp_hdrlen(skb) - max_opt_len; - if (BPF_SOCK_OPS_TEST_FLAG(tcp_sk(sk), + if (cgroup_bpf_enabled(CGROUP_SOCK_OPS) && + BPF_SOCK_OPS_TEST_FLAG(tcp_sk(sk), BPF_SOCK_OPS_WRITE_HDR_OPT_CB_FLAG)) { struct bpf_sock_ops_kern sock_ops; int err; @@ -567,7 +569,7 @@ static void bpf_skops_write_hdr_opt(struct sock *sk, struct sk_buff *skb, sock_ops.remaining_opt_len = max_opt_len; bpf_skops_init_skb(&sock_ops, skb, first_opt_off); - err = BPF_CGROUP_RUN_PROG_SOCK_OPS_SK(&sock_ops, sk); + err = __BPF_CGROUP_RUN_PROG_SOCK_OPS_SK(&sock_ops, sk); if (!err) nr_written = max_opt_len - sock_ops.remaining_opt_len; } @@ -1279,16 +1281,10 @@ static unsigned int tcp_established_options(struct sock *sk, struct sk_buff *skb } } - if (unlikely(BPF_SOCK_OPS_TEST_FLAG(tp, - BPF_SOCK_OPS_WRITE_HDR_OPT_CB_FLAG))) { - remaining = MAX_TCP_OPTION_SPACE - size; - remaining = bpf_skops_hdr_opt_len(sk, skb, NULL, NULL, 0, opts, - remaining); - - size = MAX_TCP_OPTION_SPACE - remaining; - } - remaining = MAX_TCP_OPTION_SPACE - size; + + remaining = bpf_skops_hdr_opt_len(sk, skb, NULL, NULL, 0, opts, + remaining); remaining = bpf_tcp_ops_hdr_opt_len(sk, skb, NULL, NULL, 0, opts, remaining); @@ -3725,9 +3721,8 @@ int __tcp_retransmit_skb(struct sock *sk, struct sk_buff *skb, int segs) err = tcp_transmit_skb(sk, skb, 1, GFP_ATOMIC); } - if (BPF_SOCK_OPS_TEST_FLAG(tp, BPF_SOCK_OPS_RETRANS_CB_FLAG)) - tcp_call_bpf_3arg(sk, BPF_SOCK_OPS_RETRANS_CB, - TCP_SKB_CB(skb)->seq, segs, err); + tcp_call_bpf_flag(sk, BPF_SOCK_OPS_RETRANS_CB, + TCP_SKB_CB(skb)->seq, segs, err); bpf_tcp_ops_call(retrans, sk, skb, err); if (unlikely(err) && err != -EBUSY) @@ -4355,7 +4350,7 @@ int tcp_connect(struct sock *sk) struct sk_buff *buff; int err; - tcp_call_bpf(sk, BPF_SOCK_OPS_TCP_CONNECT_CB, 0, NULL); + tcp_call_bpf(sk, BPF_SOCK_OPS_TCP_CONNECT_CB); bpf_tcp_ops_call(connect, sk); #if defined(CONFIG_TCP_MD5SIG) && defined(CONFIG_TCP_AO) diff --git a/net/ipv4/tcp_timer.c b/net/ipv4/tcp_timer.c index 3d49adc51766..00debcba122b 100644 --- a/net/ipv4/tcp_timer.c +++ b/net/ipv4/tcp_timer.c @@ -286,10 +286,8 @@ static int tcp_write_timeout(struct sock *sk) tcp_fastopen_active_detect_blackhole(sk, expired); mptcp_active_detect_blackhole(sk, expired); - if (BPF_SOCK_OPS_TEST_FLAG(tp, BPF_SOCK_OPS_RTO_CB_FLAG)) - tcp_call_bpf_3arg(sk, BPF_SOCK_OPS_RTO_CB, - icsk->icsk_retransmits, - icsk->icsk_rto, (int)expired); + tcp_call_bpf_flag(sk, BPF_SOCK_OPS_RTO_CB, + icsk->icsk_retransmits, icsk->icsk_rto, (int)expired); bpf_tcp_ops_call(rto, sk); if (expired) { -- 2.56.0.360.g66cac248cb-goog