diff --git a/lib/hypervisor/qemu/vsock.go b/lib/hypervisor/qemu/vsock.go index 9d828df38..284768e67 100644 --- a/lib/hypervisor/qemu/vsock.go +++ b/lib/hypervisor/qemu/vsock.go @@ -5,9 +5,9 @@ package qemu import ( "context" "fmt" - "io" "log/slog" "net" + "os" "time" "golang.org/x/sys/unix" @@ -130,21 +130,20 @@ func (d *VsockDialer) DialVsock(ctx context.Context, port int) (net.Conn, error) } } - // Set back to blocking mode for normal I/O - if err := unix.SetNonblock(fd, false); err != nil { - unix.Close(fd) - return nil, fmt.Errorf("set blocking: %w", err) - } - slog.DebugContext(ctx, "vsock connection established", "cid", d.cid, "port", port) // Wrap the file descriptor in a net.Conn - return newVsockConn(fd, d.cid, uint32(port)) + conn, err := newVsockConn(fd, d.cid, uint32(port)) + if err != nil { + return nil, err + } + return conn, nil } -// vsockConn wraps a vsock file descriptor as a net.Conn +// vsockConn wraps a vsock file descriptor as a net.Conn. The embedded +// os.File owns the descriptor: close, in-flight I/O, and deadlines. type vsockConn struct { - fd int + *os.File localCID uint32 localPort uint32 remoteCID uint32 @@ -152,8 +151,13 @@ type vsockConn struct { } func newVsockConn(fd int, remoteCID, remotePort uint32) (*vsockConn, error) { + // os.NewFile only registers a nonblocking descriptor with the poller. + if err := unix.SetNonblock(fd, true); err != nil { + unix.Close(fd) + return nil, fmt.Errorf("set non-blocking: %w", err) + } return &vsockConn{ - fd: fd, + File: os.NewFile(uintptr(fd), "vsock"), localCID: unix.VMADDR_CID_HOST, localPort: 0, // ephemeral remoteCID: remoteCID, @@ -161,36 +165,6 @@ func newVsockConn(fd int, remoteCID, remotePort uint32) (*vsockConn, error) { }, nil } -func (c *vsockConn) Read(b []byte) (int, error) { - n, err := unix.Read(c.fd, b) - // Ensure we never return negative n (violates io.Reader contract) - // This can happen when the vsock fd becomes invalid (VM died) - if n < 0 { - if err == nil { - err = io.EOF - } - return 0, err - } - return n, err -} - -func (c *vsockConn) Write(b []byte) (int, error) { - n, err := unix.Write(c.fd, b) - // Ensure we never return negative n (violates io.Writer contract) - // This can happen when the vsock fd becomes invalid (VM died) - if n < 0 { - if err == nil { - err = io.ErrClosedPipe - } - return 0, err - } - return n, err -} - -func (c *vsockConn) Close() error { - return unix.Close(c.fd) -} - func (c *vsockConn) LocalAddr() net.Addr { return &vsockAddr{cid: c.localCID, port: c.localPort} } @@ -199,49 +173,6 @@ func (c *vsockConn) RemoteAddr() net.Addr { return &vsockAddr{cid: c.remoteCID, port: c.remotePort} } -func (c *vsockConn) SetDeadline(t time.Time) error { - if t.IsZero() { - // Clear deadlines - if err := c.SetReadDeadline(t); err != nil { - return err - } - return c.SetWriteDeadline(t) - } - timeout := time.Until(t) - if timeout < 0 { - timeout = 0 - } - tv := unix.NsecToTimeval(timeout.Nanoseconds()) - if err := unix.SetsockoptTimeval(c.fd, unix.SOL_SOCKET, unix.SO_RCVTIMEO, &tv); err != nil { - return err - } - return unix.SetsockoptTimeval(c.fd, unix.SOL_SOCKET, unix.SO_SNDTIMEO, &tv) -} - -func (c *vsockConn) SetReadDeadline(t time.Time) error { - var tv unix.Timeval - if !t.IsZero() { - timeout := time.Until(t) - if timeout < 0 { - timeout = 0 - } - tv = unix.NsecToTimeval(timeout.Nanoseconds()) - } - return unix.SetsockoptTimeval(c.fd, unix.SOL_SOCKET, unix.SO_RCVTIMEO, &tv) -} - -func (c *vsockConn) SetWriteDeadline(t time.Time) error { - var tv unix.Timeval - if !t.IsZero() { - timeout := time.Until(t) - if timeout < 0 { - timeout = 0 - } - tv = unix.NsecToTimeval(timeout.Nanoseconds()) - } - return unix.SetsockoptTimeval(c.fd, unix.SOL_SOCKET, unix.SO_SNDTIMEO, &tv) -} - // vsockAddr implements net.Addr for vsock addresses type vsockAddr struct { cid uint32 diff --git a/lib/hypervisor/qemu/vsock_test.go b/lib/hypervisor/qemu/vsock_test.go new file mode 100644 index 000000000..02fa4f1aa --- /dev/null +++ b/lib/hypervisor/qemu/vsock_test.go @@ -0,0 +1,121 @@ +//go:build linux + +package qemu + +import ( + "io" + "os" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" + "golang.org/x/sys/unix" +) + +func vsockPair(t *testing.T) (*vsockConn, *os.File, int) { + t.Helper() + fds, err := unix.Socketpair(unix.AF_UNIX, unix.SOCK_STREAM|unix.SOCK_CLOEXEC, 0) + require.NoError(t, err) + peer := os.NewFile(uintptr(fds[1]), "peer") + t.Cleanup(func() { _ = peer.Close() }) + conn, err := newVsockConn(fds[0], 42, 1024) + require.NoError(t, err) + t.Cleanup(func() { _ = conn.Close() }) + return conn, peer, fds[0] +} + +func TestVsockConnDoesNotCloseReusedDescriptor(t *testing.T) { + conn, _, fd := vsockPair(t) + directory, err := os.Open("/proc/self/fd") + require.NoError(t, err) + defer directory.Close() + require.NoError(t, conn.Close()) + reused, err := unix.FcntlInt(directory.Fd(), unix.F_DUPFD_CLOEXEC, fd) + require.NoError(t, err) + victim := os.NewFile(uintptr(reused), "proc-fds") + defer victim.Close() + if reused != fd { + t.Skip("another goroutine reused the descriptor first") + } + + var wg sync.WaitGroup + for range 8 { + wg.Go(func() { _ = conn.Close() }) + } + wg.Wait() + _, err = victim.ReadDir(-1) + require.NoError(t, err, "repeated close must not close an unrelated directory") + + n, err := conn.Read(make([]byte, 1)) + require.Zero(t, n) + require.ErrorIs(t, err, os.ErrClosed) + n, err = conn.Write([]byte("x")) + require.Zero(t, n) + require.ErrorIs(t, err, os.ErrClosed) + require.Error(t, conn.SetDeadline(time.Now())) +} + +func TestVsockConnReadWrite(t *testing.T) { + conn, peer, _ := vsockPair(t) + require.NoError(t, conn.SetDeadline(time.Now().Add(time.Second))) + _, err := peer.Write([]byte("in")) + require.NoError(t, err) + buf := make([]byte, 2) + _, err = io.ReadFull(conn, buf) + require.NoError(t, err) + require.Equal(t, "in", string(buf)) + _, err = conn.Write([]byte("out")) + require.NoError(t, err) + buf = make([]byte, 3) + _, err = io.ReadFull(peer, buf) + require.NoError(t, err) + require.Equal(t, "out", string(buf)) + require.NoError(t, peer.Close()) + n, err := conn.Read(buf) + require.Zero(t, n) + require.ErrorIs(t, err, io.EOF) +} + +func TestVsockConnCloseInterruptsRead(t *testing.T) { + conn, _, _ := vsockPair(t) + started := make(chan struct{}) + result := make(chan error, 1) + go func() { + close(started) + _, err := conn.Read(make([]byte, 1)) + result <- err + }() + <-started + // Give the reader time to block before closing the connection. + select { + case err := <-result: + t.Fatalf("read returned before close: %v", err) + case <-time.After(10 * time.Millisecond): + } + require.NoError(t, conn.Close()) + select { + case err := <-result: + require.ErrorIs(t, err, os.ErrClosed) + case <-time.After(time.Second): + t.Fatal("close did not interrupt read") + } +} + +func TestVsockConnDeadlines(t *testing.T) { + conn, peer, _ := vsockPair(t) + require.NoError(t, conn.SetReadDeadline(time.Now().Add(-time.Second))) + _, err := conn.Read(make([]byte, 1)) + require.ErrorIs(t, err, os.ErrDeadlineExceeded) + require.NoError(t, conn.SetReadDeadline(time.Time{})) + _, err = peer.Write([]byte("x")) + require.NoError(t, err) + _, err = conn.Read(make([]byte, 1)) + require.NoError(t, err) + require.NoError(t, conn.SetWriteDeadline(time.Now().Add(-time.Second))) + _, err = conn.Write([]byte("x")) + require.ErrorIs(t, err, os.ErrDeadlineExceeded) + require.NoError(t, conn.SetDeadline(time.Time{})) + _, err = conn.Write([]byte("x")) + require.NoError(t, err) +}