[PATCH 08/16] NTB: ntb_transport: Add opt-in direct-DMA channel reservation

Koichiro Den <[email protected]>
Newsgroups dev.linux.lists.ntb,org.kernel.vger.linux-kernel,org.kernel.vger.netdev
Message-ID <[email protected]>
Direct DMA is opt-in. Add 'use_direct_dma' module parameter to request
it. Some endpoint integrations expose NTB and DMA through sibling PCI
functions, so add direct_dma_func to select the DMA function when the
NTB device cannot return the DMA device directly.

When use_direct_dma is set and negotiation scratchpads are available,
locate the DMA device and reserve one suitable channel for every QP
before registering clients. This avoids negotiating direct TX before a
later-created QP discovers that no channel is available. If the full set
is unavailable, leave direct TX disabled. Note that direct RX does not
require these channels.

Signed-off-by: Koichiro Den <[email protected]>
---
 drivers/ntb/ntb_transport.c | 120 ++++++++++++++++++++++++++++++++++++
 1 file changed, 120 insertions(+)

diff --git a/drivers/ntb/ntb_transport.c b/drivers/ntb/ntb_transport.c
index f873eacd532b..6051d8abfc0d 100644
--- a/drivers/ntb/ntb_transport.c
+++ b/drivers/ntb/ntb_transport.c
@@ -98,6 +98,16 @@ static bool use_dma;
 module_param(use_dma, bool, 0644);
 MODULE_PARM_DESC(use_dma, "Use DMA engine to perform large data copy");
 
+static bool use_direct_dma;
+module_param(use_direct_dma, bool, 0644);
+MODULE_PARM_DESC(use_direct_dma,
+		 "Use PCI endpoint DMA to transfer directly to peer RX buffers");
+
+static unsigned int direct_dma_func;
+module_param(direct_dma_func, uint, 0644);
+MODULE_PARM_DESC(direct_dma_func,
+		 "PCI function number of a sibling endpoint DMA function");
+
 static unsigned int direct_dma_ring_entries = 64;
 module_param(direct_dma_ring_entries, uint, 0644);
 MODULE_PARM_DESC(direct_dma_ring_entries,
@@ -147,6 +157,7 @@ struct ntb_transport_qp {
 	void *cb_data;
 	struct dma_chan *tx_dma_chan;
 	struct dma_chan *rx_dma_chan;
+	struct dma_chan *direct_dma_chan;
 
 	bool client_ready;
 	bool link_is_up;
@@ -269,6 +280,7 @@ struct ntb_transport_ctx {
 	struct work_struct link_cleanup;
 
 	struct dentry *debugfs_node_dir;
+	struct device *direct_dma_dev;
 	u32 direct_features;
 	u32 peer_direct_features;
 	unsigned int direct_ring_entries;
@@ -352,6 +364,97 @@ static inline u32 ntb_direct_ring_used(u32 head, u32 tail)
 	return head - tail;
 }
 
+static bool ntb_direct_dma_filter_fn(struct dma_chan *chan, void *data)
+{
+	struct dma_slave_caps caps;
+
+	if (chan->device->dev != data || dma_get_slave_caps(chan, &caps))
+		return false;
+
+	/* Payload and completion descriptors must complete in order. */
+	return caps.cmd_terminate &&
+	       !dma_has_cap(DMA_COMPLETION_NO_ORDER, chan->device->cap_mask) &&
+	       (caps.directions & BIT(DMA_MEM_TO_DEV));
+}
+
+static struct device *ntb_direct_get_dma_dev(struct ntb_dev *ndev)
+{
+	struct pci_dev *pdev = ndev->pdev;
+	struct pci_dev *dma_pdev;
+	struct device *dev;
+
+	if (ndev->ops->get_dma_dev) {
+		dev = ntb_get_dma_dev(ndev);
+		return dev ? get_device(dev) : ERR_PTR(-ENODEV);
+	}
+
+	if (!pdev || !pdev->bus || direct_dma_func > 7)
+		return ERR_PTR(-ENODEV);
+
+	dma_pdev = pci_get_domain_bus_and_slot(pci_domain_nr(pdev->bus),
+					       pdev->bus->number,
+					       PCI_DEVFN(PCI_SLOT(pdev->devfn),
+							 direct_dma_func));
+	if (!dma_pdev)
+		return ERR_PTR(-ENODEV);
+
+	return &dma_pdev->dev;
+}
+
+static void ntb_direct_dma_release_channels(struct ntb_transport_ctx *nt)
+{
+	unsigned int i;
+
+	for (i = 0; i < nt->qp_count; i++) {
+		if (!nt->qp_vec[i].direct_dma_chan)
+			continue;
+
+		dma_release_channel(nt->qp_vec[i].direct_dma_chan);
+		nt->qp_vec[i].direct_dma_chan = NULL;
+	}
+}
+
+static void ntb_direct_dma_release(struct ntb_transport_ctx *nt)
+{
+	ntb_direct_dma_release_channels(nt);
+	if (nt->direct_dma_dev) {
+		put_device(nt->direct_dma_dev);
+		nt->direct_dma_dev = NULL;
+	}
+}
+
+static void ntb_direct_dma_init(struct ntb_transport_ctx *nt)
+{
+	dma_cap_mask_t mask;
+	unsigned int i;
+
+	if (!use_direct_dma)
+		return;
+
+	nt->direct_dma_dev = ntb_direct_get_dma_dev(nt->ndev);
+	if (IS_ERR(nt->direct_dma_dev)) {
+		dev_info(&nt->ndev->dev, "direct DMA device unavailable: %pe\n",
+			 nt->direct_dma_dev);
+		nt->direct_dma_dev = NULL;
+		return;
+	}
+
+	dma_cap_zero(mask);
+	dma_cap_set(DMA_SLAVE, mask);
+
+	for (i = 0; i < nt->qp_count; i++) {
+		nt->qp_vec[i].direct_dma_chan =
+			dma_request_channel(mask, ntb_direct_dma_filter_fn,
+					    nt->direct_dma_dev);
+		if (!nt->qp_vec[i].direct_dma_chan) {
+			dev_info(&nt->ndev->dev,
+				 "not enough direct DMA channels for all QPs\n");
+			ntb_direct_dma_release_channels(nt);
+			return;
+		}
+	}
+}
+
 struct ntb_payload_header {
 	unsigned int ver;
 	unsigned int len;
@@ -1476,6 +1579,16 @@ static int ntb_transport_probe(struct ntb_client *self, struct ntb_dev *ndev)
 		return -ENOMEM;
 
 	nt->ndev = ndev;
+	if (use_direct_dma &&
+	    (direct_dma_ring_entries < 2 ||
+	     direct_dma_ring_entries > NTB_DIRECT_MAX_RING_ENTRIES ||
+	     !is_power_of_2(direct_dma_ring_entries))) {
+		dev_err(&ndev->dev,
+			"direct DMA ring entries must be a power of two between 2 and %u\n",
+			NTB_DIRECT_MAX_RING_ENTRIES);
+		rc = -EINVAL;
+		goto err;
+	}
 	nt->direct_ring_entries = direct_dma_ring_entries;
 
 	/*
@@ -1582,6 +1695,11 @@ static int ntb_transport_probe(struct ntb_client *self, struct ntb_dev *ndev)
 		if (rc)
 			goto err2;
 	}
+	if (ntb_direct_spads_available(nt))
+		ntb_direct_dma_init(nt);
+	else if (use_direct_dma)
+		dev_info(&ndev->dev,
+			 "not enough scratchpads for direct DMA negotiation\n");
 
 	mutex_init(&nt->link_event_lock);
 	INIT_DELAYED_WORK(&nt->link_work, ntb_transport_link_work);
@@ -1605,6 +1723,7 @@ static int ntb_transport_probe(struct ntb_client *self, struct ntb_dev *ndev)
 err3:
 	ntb_clear_ctx(ndev);
 err2:
+	ntb_direct_dma_release(nt);
 	kfree(nt->qp_vec);
 err1:
 	while (i--) {
@@ -1648,6 +1767,7 @@ static void ntb_transport_free(struct ntb_client *self, struct ntb_dev *ndev)
 		iounmap(nt->mw_vec[i].vbase);
 	}
 
+	ntb_direct_dma_release(nt);
 	kfree(nt->qp_vec);
 	kfree(nt->mw_vec);
 	kfree(nt);
-- 
2.51.0
lmpx.com only provides a reader for public news (NNTP) servers. It is not affiliated with the servers or forums shown here and is not responsible for the content of articles, which is written by their respective authors.