Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
99 changes: 15 additions & 84 deletions lib/hypervisor/qemu/vsock.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,9 +5,9 @@ package qemu
import (
"context"
"fmt"
"io"
"log/slog"
"net"
"os"
"time"

"golang.org/x/sys/unix"
Expand Down Expand Up @@ -130,67 +130,41 @@ 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
remotePort uint32
}

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,
remotePort: remotePort,
}, 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}
}
Expand All @@ -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
Expand Down
121 changes: 121 additions & 0 deletions lib/hypervisor/qemu/vsock_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
Loading