[PATCH v2] net/smc: hold a reference on net_device returned by pnet_find_base_ndev()
flat view
HOTtoday
From: Atharva Vartak <hidden>
Date: 2026-10-07 04:00:22
Also in:
linux-rdma, linux-s390, lkml
Subsystem:
networking [general], shared memory communications (smc) sockets, the rest · Maintainers:
"David S. Miller", Eric Dumazet, Jakub Kicinski, Paolo Abeni, D. Wythe, Dust Li, Sidraya Jayagond, Mahanta Jambigi, Linus Torvalds
pnet_find_base_ndev() resolves the base net_device for stacked devices
under RTNL, then drops the lock and returns the raw pointer without
taking a reference. All three callers dereference the pointer after
RTNL has been released, creating a use-after-free window if the device
is concurrently unregistered.
Take a reference with netdev_hold() before dropping RTNL and add the
matching netdev_put() in every return path of the three callers:
- smc_pnet_add_eth()
- smc_pnet_find_roce_by_pnetid()
- smc_pnet_find_ism_by_pnetid()
Fixes: 0afff91c6f5e ("net/smc: add pnetid support")
Fixes: 1619f770589a ("net/smc: add pnetid support for SMC-D and ISM")
Cc: "D. Wythe" <alibuda@linux.alibaba.com>
Cc: Dust Li <dust.li@linux.alibaba.com>
Cc: Sidraya Jayagond <sidraya@linux.ibm.com>
Cc: Mahanta Jambigi <mjambigi@linux.ibm.com>
Cc: Tony Lu <tonylu@linux.alibaba.com>
Cc: Wen Gu <guwen@linux.alibaba.com>
Cc: Willy Tarreau <w@1wt.eu>
Cc: linux-rdma@vger.kernel.org
Cc: linux-s390@vger.kernel.org
Cc: netdev@vger.kernel.org
Signed-off-by: Atharva Vartak <redacted>
---
v2: use netdev_hold()/netdev_put() with a netdevice_tracker instead of
the deprecated dev_hold()/dev_put() (Jakub Kicinski).
v1: https://lore.kernel.org/netdev/6aba080c.b933486d.297b72.5042@mx.google.com/ (local)
net/smc/smc_pnet.c | 29 ++++++++++++++++++++++-------
1 file changed, 22 insertions(+), 7 deletions(-)
diff --git a/net/smc/smc_pnet.c b/net/smc/smc_pnet.c
index ff9c9c35cc2f..61ef1d697eee 100644
--- a/net/smc/smc_pnet.c
+++ b/net/smc/smc_pnet.c@@ -30,7 +30,8 @@ #include "smc_core.h" static struct net_device *__pnet_find_base_ndev(struct net_device *ndev); -static struct net_device *pnet_find_base_ndev(struct net_device *ndev); +static struct net_device *pnet_find_base_ndev(struct net_device *ndev, + netdevice_tracker *tracker); static const struct nla_policy smc_pnet_policy[SMC_PNETID_MAX + 1] = { [SMC_PNETID_NAME] = {
@@ -356,6 +357,7 @@ static int smc_pnet_add_eth(struct smc_pnettable *pnettable, struct net *net, struct smc_pnetentry *tmp_pe, *new_pe; struct net_device *ndev, *base_ndev; u8 ndev_pnetid[SMC_MAX_PNETID_LEN]; + netdevice_tracker base_tracker; bool new_netdev; int rc;
@@ -365,10 +367,13 @@ static int smc_pnet_add_eth(struct smc_pnettable *pnettable, struct net *net, rc = -EEXIST; ndev = dev_get_by_name(net, eth_name); /* dev_hold() */ if (ndev) { - base_ndev = pnet_find_base_ndev(ndev); + base_ndev = pnet_find_base_ndev(ndev, &base_tracker); if (!smc_pnetid_by_dev_port(base_ndev->dev.parent, - base_ndev->dev_port, ndev_pnetid)) + base_ndev->dev_port, ndev_pnetid)) { + netdev_put(base_ndev, &base_tracker); goto out_put; + } + netdev_put(base_ndev, &base_tracker); } /* add a new netdev entry to the pnet table if there isn't one */
@@ -945,10 +950,13 @@ static struct net_device *__pnet_find_base_ndev(struct net_device *ndev) * (for instance with bonding slaves), just the first device * is used to reach a base device. */ -static struct net_device *pnet_find_base_ndev(struct net_device *ndev) +static struct net_device *pnet_find_base_ndev(struct net_device *ndev, + netdevice_tracker *tracker) { rtnl_lock(); ndev = __pnet_find_base_ndev(ndev); + /* keep ndev alive after dropping RTNL, callers must netdev_put() */ + netdev_hold(ndev, tracker, GFP_KERNEL); rtnl_unlock(); return ndev; }
@@ -1085,17 +1093,20 @@ static void smc_pnet_find_roce_by_pnetid(struct net_device *ndev, { u8 ndev_pnetid[SMC_MAX_PNETID_LEN]; struct net_device *base_ndev; + netdevice_tracker base_tracker; struct net *net; - base_ndev = pnet_find_base_ndev(ndev); + base_ndev = pnet_find_base_ndev(ndev, &base_tracker); net = dev_net(ndev); if (smc_pnetid_by_dev_port(base_ndev->dev.parent, base_ndev->dev_port, ndev_pnetid) && smc_pnet_find_ndev_pnetid_by_table(base_ndev, ndev_pnetid) && smc_pnet_find_ndev_pnetid_by_table(ndev, ndev_pnetid)) { smc_pnet_find_rdma_dev(base_ndev, ini); + netdev_put(base_ndev, &base_tracker); return; /* pnetid could not be determined */ } + netdev_put(base_ndev, &base_tracker); _smc_pnet_find_roce_by_pnetid(ndev_pnetid, ini, NULL, net); }
@@ -1103,13 +1114,16 @@ static void smc_pnet_find_ism_by_pnetid(struct net_device *ndev, struct smc_init_info *ini) { u8 ndev_pnetid[SMC_MAX_PNETID_LEN]; + netdevice_tracker base_tracker; struct smcd_dev *ismdev; - ndev = pnet_find_base_ndev(ndev); + ndev = pnet_find_base_ndev(ndev, &base_tracker); if (smc_pnetid_by_dev_port(ndev->dev.parent, ndev->dev_port, ndev_pnetid) && - smc_pnet_find_ndev_pnetid_by_table(ndev, ndev_pnetid)) + smc_pnet_find_ndev_pnetid_by_table(ndev, ndev_pnetid)) { + netdev_put(ndev, &base_tracker); return; /* pnetid could not be determined */ + } mutex_lock(&smcd_dev_list.mutex); list_for_each_entry(ismdev, &smcd_dev_list.list, list) {
@@ -1123,6 +1137,7 @@ static void smc_pnet_find_ism_by_pnetid(struct net_device *ndev, } } mutex_unlock(&smcd_dev_list.mutex); + netdev_put(ndev, &base_tracker); } /* PNET table analysis for a given sock:
--
2.56.0