udp: Set NULL to sk->sk_prot->h.udp_table.
authorKuniyuki Iwashima <kuniyu@amazon.com>
Mon, 14 Nov 2022 21:57:54 +0000 (13:57 -0800)
committerDavid S. Miller <davem@davemloft.net>
Wed, 16 Nov 2022 09:43:35 +0000 (09:43 +0000)
We will soon introduce an optional per-netns hash table
for UDP.

This means we cannot use the global sk->sk_prot->h.udp_table
to fetch a UDP hash table.

Instead, set NULL to sk->sk_prot->h.udp_table for UDP and get
a proper table from net->ipv4.udp_table.

Note that we still need sk->sk_prot->h.udp_table for UDP LITE.

Signed-off-by: Kuniyuki Iwashima <kuniyu@amazon.com>
Signed-off-by: David S. Miller <davem@davemloft.net>
include/net/netns/ipv4.h
net/ipv4/udp.c
net/ipv6/udp.c

index 25f90bb..e4cc4d3 100644 (file)
@@ -43,6 +43,7 @@ struct tcp_fastopen_context;
 
 struct netns_ipv4 {
        struct inet_timewait_death_row tcp_death_row;
+       struct udp_table *udp_table;
 
 #ifdef CONFIG_SYSCTL
        struct ctl_table_header *forw_hdr;
index a34de26..6206c27 100644 (file)
@@ -131,6 +131,11 @@ EXPORT_PER_CPU_SYMBOL_GPL(udp_memory_per_cpu_fw_alloc);
 #define MAX_UDP_PORTS 65536
 #define PORTS_PER_CHAIN (MAX_UDP_PORTS / UDP_HTABLE_SIZE_MIN)
 
+static struct udp_table *udp_get_table_prot(struct sock *sk)
+{
+       return sk->sk_prot->h.udp_table ? : sock_net(sk)->ipv4.udp_table;
+}
+
 static int udp_lib_lport_inuse(struct net *net, __u16 num,
                               const struct udp_hslot *hslot,
                               unsigned long *bitmap,
@@ -232,7 +237,7 @@ static int udp_reuseport_add_sock(struct sock *sk, struct udp_hslot *hslot)
 int udp_lib_get_port(struct sock *sk, unsigned short snum,
                     unsigned int hash2_nulladdr)
 {
-       struct udp_table *udptable = sk->sk_prot->h.udp_table;
+       struct udp_table *udptable = udp_get_table_prot(sk);
        struct udp_hslot *hslot, *hslot2;
        struct net *net = sock_net(sk);
        int error = 1;
@@ -1999,7 +2004,7 @@ EXPORT_SYMBOL(udp_disconnect);
 void udp_lib_unhash(struct sock *sk)
 {
        if (sk_hashed(sk)) {
-               struct udp_table *udptable = sk->sk_prot->h.udp_table;
+               struct udp_table *udptable = udp_get_table_prot(sk);
                struct udp_hslot *hslot, *hslot2;
 
                hslot  = udp_hashslot(udptable, sock_net(sk),
@@ -2030,7 +2035,7 @@ EXPORT_SYMBOL(udp_lib_unhash);
 void udp_lib_rehash(struct sock *sk, u16 newhash)
 {
        if (sk_hashed(sk)) {
-               struct udp_table *udptable = sk->sk_prot->h.udp_table;
+               struct udp_table *udptable = udp_get_table_prot(sk);
                struct udp_hslot *hslot, *hslot2, *nhslot2;
 
                hslot2 = udp_hashslot2(udptable, udp_sk(sk)->udp_portaddr_hash);
@@ -2967,7 +2972,7 @@ struct proto udp_prot = {
        .sysctl_wmem_offset     = offsetof(struct net, ipv4.sysctl_udp_wmem_min),
        .sysctl_rmem_offset     = offsetof(struct net, ipv4.sysctl_udp_rmem_min),
        .obj_size               = sizeof(struct udp_sock),
-       .h.udp_table            = &udp_table,
+       .h.udp_table            = NULL,
        .diag_destroy           = udp_abort,
 };
 EXPORT_SYMBOL(udp_prot);
@@ -3280,6 +3285,8 @@ EXPORT_SYMBOL(udp_flow_hashrnd);
 
 static int __net_init udp_sysctl_init(struct net *net)
 {
+       net->ipv4.udp_table = &udp_table;
+
        net->ipv4.sysctl_udp_rmem_min = PAGE_SIZE;
        net->ipv4.sysctl_udp_wmem_min = PAGE_SIZE;
 
index 727de67..bbd6dc3 100644 (file)
@@ -1774,7 +1774,7 @@ struct proto udpv6_prot = {
        .sysctl_wmem_offset     = offsetof(struct net, ipv4.sysctl_udp_wmem_min),
        .sysctl_rmem_offset     = offsetof(struct net, ipv4.sysctl_udp_rmem_min),
        .obj_size               = sizeof(struct udp6_sock),
-       .h.udp_table            = &udp_table,
+       .h.udp_table            = NULL,
        .diag_destroy           = udp_abort,
 };