virtio_vdpa relies on the core's synchronize_rcu() fallback for callback
synchronization. That covers IRQ handlers and sections with interrupts,
preemption or bottom halves disabled. virtio_vdpa does not enforce any
of those contexts when a vDPA driver invokes a callback.

Add an SRCU domain per device and enter it around the virtqueue and
config callbacks. Implement synchronize_cbs() with synchronize_srcu()
so it covers the callbacks regardless of the calling context.

SRCU uses per-CPU reader accounting, avoiding a single callback lock
shared by all queues. Cover the config callback too, as virtio-pci
does for its config vector.

Assisted-by: LLM
Signed-off-by: Karl Mehltretter <[email protected]>
---
 drivers/virtio/virtio_vdpa.c | 34 +++++++++++++++++++++++++++++++---
 1 file changed, 31 insertions(+), 3 deletions(-)

diff --git a/drivers/virtio/virtio_vdpa.c b/drivers/virtio/virtio_vdpa.c
index 6bcf4567a929..f9559b6b2e0b 100644
--- a/drivers/virtio/virtio_vdpa.c
+++ b/drivers/virtio/virtio_vdpa.c
@@ -12,6 +12,7 @@
 #include <linux/device.h>
 #include <linux/kernel.h>
 #include <linux/slab.h>
+#include <linux/srcu.h>
 #include <linux/uuid.h>
 #include <linux/group_cpus.h>
 #include <linux/virtio.h>
@@ -27,6 +28,7 @@
 struct virtio_vdpa_device {
        struct virtio_device vdev;
        struct vdpa_device *vdpa;
+       struct srcu_struct cb_srcu;
        u64 features;
 };
 
@@ -114,8 +116,11 @@ static bool virtio_vdpa_notify_with_data(struct virtqueue 
*vq)
 static irqreturn_t virtio_vdpa_config_cb(void *private)
 {
        struct virtio_vdpa_device *vd_dev = private;
+       int idx;
 
+       idx = srcu_read_lock(&vd_dev->cb_srcu);
        virtio_config_changed(&vd_dev->vdev);
+       srcu_read_unlock(&vd_dev->cb_srcu, idx);
 
        return IRQ_HANDLED;
 }
@@ -123,8 +128,22 @@ static irqreturn_t virtio_vdpa_config_cb(void *private)
 static irqreturn_t virtio_vdpa_virtqueue_cb(void *private)
 {
        struct virtqueue *vq = private;
+       struct virtio_vdpa_device *vd_dev = to_virtio_vdpa_device(vq->vdev);
+       irqreturn_t ret;
+       int idx;
+
+       idx = srcu_read_lock(&vd_dev->cb_srcu);
+       ret = vring_interrupt(0, vq);
+       srcu_read_unlock(&vd_dev->cb_srcu, idx);
+
+       return ret;
+}
+
+static void virtio_vdpa_synchronize_cbs(struct virtio_device *vdev)
+{
+       struct virtio_vdpa_device *vd_dev = to_virtio_vdpa_device(vdev);
 
-       return vring_interrupt(0, vq);
+       synchronize_srcu(&vd_dev->cb_srcu);
 }
 
 static struct virtqueue *
@@ -440,6 +459,7 @@ static const struct virtio_config_ops 
virtio_vdpa_config_ops = {
        .reset          = virtio_vdpa_reset,
        .find_vqs       = virtio_vdpa_find_vqs,
        .del_vqs        = virtio_vdpa_del_vqs,
+       .synchronize_cbs = virtio_vdpa_synchronize_cbs,
        .get_features   = virtio_vdpa_get_features,
        .finalize_features = virtio_vdpa_finalize_features,
        .bus_name       = virtio_vdpa_bus_name,
@@ -454,6 +474,7 @@ static void virtio_vdpa_release_dev(struct device *_d)
        struct virtio_vdpa_device *vd_dev =
               container_of(vdev, struct virtio_vdpa_device, vdev);
 
+       cleanup_srcu_struct(&vd_dev->cb_srcu);
        kfree(vd_dev);
 }
 
@@ -467,6 +488,11 @@ static int virtio_vdpa_probe(struct vdpa_device *vdpa)
        if (!vd_dev)
                return -ENOMEM;
 
+       if (init_srcu_struct(&vd_dev->cb_srcu)) {
+               kfree(vd_dev);
+               return -ENOMEM;
+       }
+
        vd_dev->vdev.dev.parent = vdpa->map ? &vdpa->dev :
                                  vdpa_get_map(vdpa).dma_dev;
        vd_dev->vdev.dev.release = virtio_vdpa_release_dev;
@@ -489,10 +515,12 @@ static int virtio_vdpa_probe(struct vdpa_device *vdpa)
        return 0;
 
 err:
-       if (reg_dev)
+       if (reg_dev) {
                put_device(&vd_dev->vdev.dev);
-       else
+       } else {
+               cleanup_srcu_struct(&vd_dev->cb_srcu);
                kfree(vd_dev);
+       }
        return ret;
 }
 
-- 
2.39.5 (Apple Git-154)

Reply via email to