Introduce vsock_unconnected_reset() and adapt vsock_connect().

Change transport assignment life cycle. On connect(), socket gets a
transport assigned. If connection fails (init went wrong, peer
misbehaviour, time out, signal), transport is de-assigned and socket state
is re-initialized. Once the connection is established, transport remains
assigned until close().

Signed-off-by: Michal Luczaj <[email protected]>
---
 net/vmw_vsock/af_vsock.c | 74 ++++++++++++++++++++++++++++++------------------
 1 file changed, 47 insertions(+), 27 deletions(-)

diff --git a/net/vmw_vsock/af_vsock.c b/net/vmw_vsock/af_vsock.c
index 5c8e7e7d35b4..20181ddde114 100644
--- a/net/vmw_vsock/af_vsock.c
+++ b/net/vmw_vsock/af_vsock.c
@@ -1679,6 +1679,42 @@ static int vsock_transport_cancel_pkt(struct vsock_sock 
*vsk)
        return transport->cancel_pkt(vsk);
 }
 
+static void vsock_unconnected_reset(struct sock *sk)
+{
+       struct vsock_sock *vsk = vsock_sk(sk);
+
+       sock_owned_by_me(sk);
+
+       /*
+        * Only connected socks may have peer_shutdown or SOCK_DONE set.
+        *
+        * Once established (TCP_ESTABLISHED, TCP_CLOSING), a socket can be
+        * de-assigned only on close(). But we can narrow the check down to
+        * states we actually expect (TCP_SYN_SENT, TCP_CLOSE).
+        */
+       if (WARN_ON_ONCE(vsk->peer_shutdown) ||
+           WARN_ON_ONCE(sock_flag(sk, SOCK_DONE)) ||
+           WARN_ON_ONCE(sk->sk_state != TCP_SYN_SENT &&
+                        sk->sk_state != TCP_CLOSE))
+               return;
+
+       /*
+        * Try to cancel a VIRTIO_VSOCK_OP_REQUEST skb that may have been sent
+        * out by transport->connect().
+        */
+       vsock_transport_cancel_pkt(vsk);
+
+       /*
+        * No need to invoke transport->release() for unconnected connectible
+        * sockets. Go straight for transport deassign.
+        */
+       vsock_deassign_transport(vsk);
+
+       /* Revert socket to initial state. Keep sk_err. */
+       WRITE_ONCE(sk->sk_state, TCP_CLOSE);
+       sk->sk_socket->state = SS_UNCONNECTED;
+}
+
 static void vsock_connect_timeout(struct work_struct *work)
 {
        struct sock *sk;
@@ -1690,11 +1726,9 @@ static void vsock_connect_timeout(struct work_struct 
*work)
        lock_sock(sk);
        if (sk->sk_state == TCP_SYN_SENT &&
            (sk->sk_shutdown != SHUTDOWN_MASK)) {
-               sk->sk_state = TCP_CLOSE;
-               sk->sk_socket->state = SS_UNCONNECTED;
                sk->sk_err = ETIMEDOUT;
                sk_error_report(sk);
-               vsock_transport_cancel_pkt(vsk);
+               vsock_unconnected_reset(sk);
        }
        release_sock(sk);
 
@@ -1760,7 +1794,7 @@ static int vsock_connect(struct socket *sock, struct 
sockaddr_unsized *addr,
                if (!transport->stream_allow(vsk, remote_addr->svm_cid,
                                             remote_addr->svm_port)) {
                        err = -ENETUNREACH;
-                       goto out;
+                       goto out_reset;
                }
 
                if (vsock_msgzerocopy_allow(transport)) {
@@ -1771,18 +1805,18 @@ static int vsock_connect(struct socket *sock, struct 
sockaddr_unsized *addr,
                         * feature is supported here.
                         */
                        err = -EOPNOTSUPP;
-                       goto out;
+                       goto out_reset;
                }
 
                err = vsock_auto_bind(vsk);
                if (err)
-                       goto out;
+                       goto out_reset;
 
                sk->sk_state = TCP_SYN_SENT;
 
                err = transport->connect(vsk);
                if (err < 0)
-                       goto out;
+                       goto out_reset;
 
                /* sk_err might have been set as a result of an earlier
                 * (failed) connect attempt.
@@ -1825,8 +1859,9 @@ static int vsock_connect(struct socket *sock, struct 
sockaddr_unsized *addr,
                                             timeout))
                                sock_put(sk);
 
+                       finish_wait(sk_sleep(sk), &wait);
                        /* Skip ahead to preserve error code set above. */
-                       goto out_wait;
+                       goto out;
                }
 
                release_sock(sk);
@@ -1844,12 +1879,7 @@ static int vsock_connect(struct socket *sock, struct 
sockaddr_unsized *addr,
                        break;
 
                /* If connection was _not_ established and a signal/timeout came
-                * to be, we want the socket's state reset. User space may want
-                * to retry.
-                *
-                * sk_state != TCP_ESTABLISHED implies that socket is not on
-                * vsock_connected_table. We keep the binding and the transport
-                * assigned.
+                * to be, we want the socket's state reset. We keep the binding.
                 */
                if (signal_pending(current) || timeout == 0) {
                        err = timeout == 0 ? -ETIMEDOUT : 
sock_intr_errno(timeout);
@@ -1859,14 +1889,6 @@ static int vsock_connect(struct socket *sock, struct 
sockaddr_unsized *addr,
                         * sk_state == TCP_SYN_SENT, which hereby we break.
                         * In such case VIRTIO_VSOCK_OP_RST will follow.
                         */
-                       sk->sk_state = TCP_CLOSE;
-                       sock->state = SS_UNCONNECTED;
-
-                       /* Try to cancel VIRTIO_VSOCK_OP_REQUEST skb sent out by
-                        * transport->connect().
-                        */
-                       vsock_transport_cancel_pkt(vsk);
-
                        goto out_wait;
                }
 
@@ -1874,13 +1896,11 @@ static int vsock_connect(struct socket *sock, struct 
sockaddr_unsized *addr,
        }
 
        err = sock_error(sk);
-       if (err) {
-               sk->sk_state = TCP_CLOSE;
-               sock->state = SS_UNCONNECTED;
-       }
-
 out_wait:
        finish_wait(sk_sleep(sk), &wait);
+out_reset:
+       if (err)
+               vsock_unconnected_reset(sk);
 out:
        release_sock(sk);
        return err;

-- 
2.55.0


Reply via email to