From: Bobby Eshleman <bobbyeshle...@meta.com>

Add the ability to isolate vsock flows using namespaces.

The VM, via the vhost_vsock struct, inherits its namespace from the
process that opens the vhost-vsock device. vhost_vsock lookup functions
are modified to take into account the mode (e.g., if CIDs are matching
but modes don't align, then return NULL).

vhost_vsock now acquires a reference to the namespace.

Signed-off-by: Bobby Eshleman <bobbyeshle...@meta.com>

---
Changes in v5:
- respect pid namespaces when assigning namespace to vhost_vsock
---
 drivers/vhost/vsock.c | 26 ++++++++++++++++++--------
 1 file changed, 18 insertions(+), 8 deletions(-)

diff --git a/drivers/vhost/vsock.c b/drivers/vhost/vsock.c
index 34adf0cf9124..f7405bb27aab 100644
--- a/drivers/vhost/vsock.c
+++ b/drivers/vhost/vsock.c
@@ -46,6 +46,8 @@ static DEFINE_READ_MOSTLY_HASHTABLE(vhost_vsock_hash, 8);
 struct vhost_vsock {
        struct vhost_dev dev;
        struct vhost_virtqueue vqs[2];
+       struct net *net;
+       netns_tracker ns_tracker;
 
        /* Link to global vhost_vsock_hash, writes use vhost_vsock_mutex */
        struct hlist_node hash;
@@ -67,7 +69,7 @@ static u32 vhost_transport_get_local_cid(void)
 /* Callers that dereference the return value must hold vhost_vsock_mutex or the
  * RCU read lock.
  */
-static struct vhost_vsock *vhost_vsock_get(u32 guest_cid)
+static struct vhost_vsock *vhost_vsock_get(u32 guest_cid, struct net *net)
 {
        struct vhost_vsock *vsock;
 
@@ -78,9 +80,8 @@ static struct vhost_vsock *vhost_vsock_get(u32 guest_cid)
                if (other_cid == 0)
                        continue;
 
-               if (other_cid == guest_cid)
+               if (other_cid == guest_cid && vsock_net_check_mode(net, 
vsock->net))
                        return vsock;
-
        }
 
        return NULL;
@@ -272,13 +273,14 @@ static int
 vhost_transport_send_pkt(struct sk_buff *skb)
 {
        struct virtio_vsock_hdr *hdr = virtio_vsock_hdr(skb);
+       struct net *net = virtio_vsock_skb_net(skb);
        struct vhost_vsock *vsock;
        int len = skb->len;
 
        rcu_read_lock();
 
        /* Find the vhost_vsock according to guest context id  */
-       vsock = vhost_vsock_get(le64_to_cpu(hdr->dst_cid));
+       vsock = vhost_vsock_get(le64_to_cpu(hdr->dst_cid), net);
        if (!vsock) {
                rcu_read_unlock();
                kfree_skb(skb);
@@ -305,7 +307,7 @@ vhost_transport_cancel_pkt(struct vsock_sock *vsk)
        rcu_read_lock();
 
        /* Find the vhost_vsock according to guest context id  */
-       vsock = vhost_vsock_get(vsk->remote_addr.svm_cid);
+       vsock = vhost_vsock_get(vsk->remote_addr.svm_cid, 
sock_net(sk_vsock(vsk)));
        if (!vsock)
                goto out;
 
@@ -462,11 +464,12 @@ static struct virtio_transport vhost_transport = {
 
 static bool vhost_transport_seqpacket_allow(struct vsock_sock *vsk, u32 
remote_cid)
 {
+       struct net *net = sock_net(sk_vsock(vsk));
        struct vhost_vsock *vsock;
        bool seqpacket_allow = false;
 
        rcu_read_lock();
-       vsock = vhost_vsock_get(remote_cid);
+       vsock = vhost_vsock_get(remote_cid, net);
 
        if (vsock)
                seqpacket_allow = vsock->seqpacket_allow;
@@ -526,6 +529,7 @@ static void vhost_vsock_handle_tx_kick(struct vhost_work 
*work)
                        continue;
                }
 
+               virtio_vsock_skb_set_net(skb, vsock->net);
                total_len += sizeof(*hdr) + skb->len;
 
                /* Deliver to monitoring devices all received packets */
@@ -652,10 +656,14 @@ static void vhost_vsock_free(struct vhost_vsock *vsock)
 
 static int vhost_vsock_dev_open(struct inode *inode, struct file *file)
 {
+
        struct vhost_virtqueue **vqs;
        struct vhost_vsock *vsock;
+       struct net *net;
        int ret;
 
+       net = current->nsproxy->net_ns;
+
        /* This struct is large and allocation could fail, fall back to vmalloc
         * if there is no other way.
         */
@@ -669,6 +677,7 @@ static int vhost_vsock_dev_open(struct inode *inode, struct 
file *file)
                goto out;
        }
 
+       vsock->net = get_net_track(net, &vsock->ns_tracker, GFP_KERNEL);
        vsock->guest_cid = 0; /* no CID assigned yet */
        vsock->seqpacket_allow = false;
 
@@ -708,7 +717,7 @@ static void vhost_vsock_reset_orphans(struct sock *sk)
         */
 
        /* If the peer is still valid, no need to reset connection */
-       if (vhost_vsock_get(vsk->remote_addr.svm_cid))
+       if (vhost_vsock_get(vsk->remote_addr.svm_cid, sock_net(sk)))
                return;
 
        /* If the close timeout is pending, let it expire.  This avoids races
@@ -753,6 +762,7 @@ static int vhost_vsock_dev_release(struct inode *inode, 
struct file *file)
        virtio_vsock_skb_queue_purge(&vsock->send_pkt_queue);
 
        vhost_dev_cleanup(&vsock->dev);
+       put_net_track(vsock->net, &vsock->ns_tracker);
        kfree(vsock->dev.vqs);
        vhost_vsock_free(vsock);
        return 0;
@@ -779,7 +789,7 @@ static int vhost_vsock_set_cid(struct vhost_vsock *vsock, 
u64 guest_cid)
 
        /* Refuse if CID is already in use */
        mutex_lock(&vhost_vsock_mutex);
-       other = vhost_vsock_get(guest_cid);
+       other = vhost_vsock_get(guest_cid, vsock->net);
        if (other && other != vsock) {
                mutex_unlock(&vhost_vsock_mutex);
                return -EADDRINUSE;

-- 
2.47.3


Reply via email to