Files
sing-box/service/usbip/linux_test.go
T
2026-06-09 10:42:26 +08:00

2832 lines
77 KiB
Go

//go:build linux
package usbip
import (
"bytes"
"context"
"encoding/binary"
"errors"
"fmt"
"io"
"net"
"os"
"os/exec"
"path/filepath"
"slices"
"sync"
"testing"
"time"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
M "github.com/sagernet/sing/common/metadata"
"github.com/stretchr/testify/require"
"golang.org/x/sys/unix"
)
type testLogWriter struct {
access sync.Mutex
buffer bytes.Buffer
}
func (w *testLogWriter) Write(p []byte) (int, error) {
w.access.Lock()
defer w.access.Unlock()
return w.buffer.Write(p)
}
func (w *testLogWriter) String() string {
w.access.Lock()
defer w.access.Unlock()
return w.buffer.String()
}
type testDialer struct{}
func (testDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
var dialer net.Dialer
return dialer.DialContext(ctx, network, destination.String())
}
func (testDialer) ListenPacket(context.Context, M.Socksaddr) (net.PacketConn, error) {
return nil, errors.New("unused")
}
type failingDialer struct {
err error
}
func (d failingDialer) DialContext(context.Context, string, M.Socksaddr) (net.Conn, error) {
return nil, d.err
}
func (d failingDialer) ListenPacket(context.Context, M.Socksaddr) (net.PacketConn, error) {
return nil, errors.New("unused")
}
type opaqueConn struct {
net.Conn
}
type wrappingDialer struct{}
func (wrappingDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
var dialer net.Dialer
conn, err := dialer.DialContext(ctx, network, destination.String())
if err != nil {
return nil, err
}
return opaqueConn{Conn: conn}, nil
}
func (wrappingDialer) ListenPacket(context.Context, M.Socksaddr) (net.PacketConn, error) {
return nil, errors.New("unused")
}
type testDeviceStore struct {
access sync.Mutex
devices map[string]sysfsDevice
statuses map[string]int
sockfds map[string]int
sockfdWrites map[string][]int
}
func newTestDeviceStore(devices ...sysfsDevice) *testDeviceStore {
store := &testDeviceStore{
devices: make(map[string]sysfsDevice),
statuses: make(map[string]int),
sockfds: make(map[string]int),
sockfdWrites: make(map[string][]int),
}
store.setDevices(devices...)
return store
}
func (s *testDeviceStore) setDevices(devices ...sysfsDevice) {
s.access.Lock()
defer s.access.Unlock()
s.devices = make(map[string]sysfsDevice, len(devices))
for _, device := range devices {
s.devices[device.BusID] = device
}
}
func (s *testDeviceStore) setStatus(busid string, status int) {
s.access.Lock()
defer s.access.Unlock()
s.statuses[busid] = status
}
func (s *testDeviceStore) listUSBDevices() ([]sysfsDevice, error) {
s.access.Lock()
defer s.access.Unlock()
out := make([]sysfsDevice, 0, len(s.devices))
for _, device := range s.devices {
out = append(out, device)
}
slices.SortFunc(out, func(left, right sysfsDevice) int {
switch {
case left.BusID < right.BusID:
return -1
case left.BusID > right.BusID:
return 1
default:
return 0
}
})
return out, nil
}
func (s *testDeviceStore) readSysfsDevice(busid, path string) (sysfsDevice, error) {
s.access.Lock()
defer s.access.Unlock()
device, ok := s.devices[busid]
if !ok {
return sysfsDevice{}, os.ErrNotExist
}
return device, nil
}
func (s *testDeviceStore) readUsbipStatus(busid string) (int, error) {
s.access.Lock()
defer s.access.Unlock()
status, ok := s.statuses[busid]
if !ok {
return 0, os.ErrNotExist
}
return status, nil
}
func (s *testDeviceStore) writeUsbipSockfd(busid string, fd int) error {
s.access.Lock()
defer s.access.Unlock()
s.sockfds[busid] = fd
s.sockfdWrites[busid] = append(s.sockfdWrites[busid], fd)
return nil
}
func (s *testDeviceStore) lastSockfd(busid string) int {
s.access.Lock()
defer s.access.Unlock()
return s.sockfds[busid]
}
func (s *testDeviceStore) hasPositiveSockfd(busid string) bool {
s.access.Lock()
defer s.access.Unlock()
for _, fd := range s.sockfdWrites[busid] {
if fd > 0 {
return true
}
}
return false
}
type testUSBEventListener struct {
closeOnce sync.Once
waitOnce sync.Once
closed chan struct{}
waitEntered chan struct{}
}
func newTestUSBEventListener() *testUSBEventListener {
return &testUSBEventListener{
closed: make(chan struct{}),
waitEntered: make(chan struct{}),
}
}
func (l *testUSBEventListener) Close() error {
l.closeOnce.Do(func() {
close(l.closed)
})
return nil
}
func (l *testUSBEventListener) WaitUSBEvent() error {
l.waitOnce.Do(func() {
close(l.waitEntered)
})
<-l.closed
return context.Canceled
}
func newTestUSBIPOps(t *testing.T) usbipOps {
t.Helper()
return usbipOps{
ensureHostDriver: func() error {
t.Fatalf("unexpected ensureHostDriver")
return nil
},
ensureVHCI: func() error {
t.Fatalf("unexpected ensureVHCI")
return nil
},
listUSBDevices: func() ([]sysfsDevice, error) {
t.Fatalf("unexpected listUSBDevices")
return nil, nil
},
readSysfsDevice: func(string, string) (sysfsDevice, error) {
t.Fatalf("unexpected readSysfsDevice")
return sysfsDevice{}, nil
},
currentDriver: func(string) (string, error) {
t.Fatalf("unexpected currentDriver")
return "", nil
},
unbindFromDriver: func(string, string) error {
t.Fatalf("unexpected unbindFromDriver")
return nil
},
bindToDriver: func(string, string) error {
t.Fatalf("unexpected bindToDriver")
return nil
},
hostMatchBusID: func(string, bool) error {
t.Fatalf("unexpected hostMatchBusID")
return nil
},
hostBind: func(string) error {
t.Fatalf("unexpected hostBind")
return nil
},
hostUnbind: func(string) error {
t.Fatalf("unexpected hostUnbind")
return nil
},
reloadHostDriver: func() error {
t.Fatalf("unexpected reloadHostDriver")
return nil
},
readUsbipStatus: func(string) (int, error) {
t.Fatalf("unexpected readUsbipStatus")
return 0, nil
},
writeUsbipSockfd: func(string, int) error {
t.Fatalf("unexpected writeUsbipSockfd")
return nil
},
newUEventListener: func() (usbEventListener, error) {
t.Fatalf("unexpected newUEventListener")
return nil, nil
},
vhciPickFreePort: func(uint32, map[int]struct{}) (int, error) {
t.Fatalf("unexpected vhciPickFreePort")
return 0, nil
},
vhciAttach: func(int, uintptr, uint32, uint32) error {
t.Fatalf("unexpected vhciAttach")
return nil
},
vhciDetach: func(int) error {
t.Fatalf("unexpected vhciDetach")
return nil
},
}
}
func newTestLogger(t testing.TB) log.ContextLogger {
t.Helper()
writer := new(testLogWriter)
factory := log.NewDefaultFactory(
context.Background(),
log.Formatter{
BaseTime: time.Now(),
DisableColors: true,
},
writer,
"",
nil,
false,
)
factory.SetLevel(log.LevelTrace)
t.Cleanup(func() {
if output := writer.String(); t.Failed() && output != "" {
t.Logf("USB/IP log:\n%s", output)
}
_ = factory.Close()
})
return factory.NewLogger("usbip")
}
func newTestDevice(busid string, vendorID, productID uint16, serial string, speed uint32) sysfsDevice {
return sysfsDevice{
BusID: busid,
Path: sysBusDevicePath(busid),
BusNum: 3,
DevNum: 9,
Speed: speed,
VendorID: vendorID,
ProductID: productID,
DeviceClass: 0,
ConfigValue: 1,
NumConfigs: 1,
NumInterfaces: 1,
Serial: serial,
Interfaces: []DeviceInterface{{
BInterfaceClass: 0xff,
}},
}
}
func startDispatchServer(t *testing.T, server *ServerService) (M.Socksaddr, func()) {
t.Helper()
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
done := make(chan struct{})
go func() {
defer close(done)
for {
conn, acceptErr := listener.Accept()
if acceptErr != nil {
return
}
go server.dispatchConn(conn)
}
}()
return M.SocksaddrFromNet(listener.Addr()), func() {
_ = listener.Close()
<-done
}
}
func duplicateConnFromFD(t *testing.T, fd uintptr, name string) net.Conn {
t.Helper()
conn, err := duplicateNetConnFromFD(fd, name)
require.NoError(t, err)
return conn
}
func duplicateNetConnFromFD(fd uintptr, name string) (net.Conn, error) {
dupFD, err := unix.Dup(int(fd))
if err != nil {
return nil, err
}
file := os.NewFile(uintptr(dupFD), name)
conn, err := net.FileConn(file)
closeErr := file.Close()
if err != nil {
return nil, err
}
if closeErr != nil {
return nil, closeErr
}
return conn, nil
}
func linuxServerControlState(server *ServerService, busid string) string {
server.controlAccess.Lock()
defer server.controlAccess.Unlock()
return server.controlState[busid].State
}
func duplicateHandoffKernelConn(t *testing.T, handoff *usbipConnHandoff) net.Conn {
t.Helper()
conn := duplicateConnFromFD(t, handoff.kernelFD(), "usbip-test-kernel")
require.NoError(t, handoff.closeKernelFD())
return conn
}
func requireConnRead(t *testing.T, conn net.Conn, expected []byte) {
t.Helper()
buffer := make([]byte, len(expected))
_, err := io.ReadFull(conn, buffer)
require.NoError(t, err)
require.Equal(t, expected, buffer)
}
func requireConnEOF(t *testing.T, conn net.Conn) {
t.Helper()
buffer := make([]byte, 1)
n, err := conn.Read(buffer)
require.Zero(t, n)
require.ErrorIs(t, err, io.EOF)
}
func setConnDeadline(t *testing.T, conn net.Conn) {
t.Helper()
require.NoError(t, conn.SetDeadline(time.Now().Add(3*time.Second)))
}
func requireStreamSocketFD(t *testing.T, fd uintptr) {
t.Helper()
socketType, err := unix.GetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_TYPE)
require.NoError(t, err)
require.Equal(t, unix.SOCK_STREAM, socketType)
}
type testUSBGadget struct {
path string
serial string
busid string
}
func requireRoot(t *testing.T) {
t.Helper()
if os.Geteuid() != 0 {
t.Skip("root required")
}
}
func requireKernelModule(t *testing.T, module string) {
t.Helper()
if _, err := os.Stat(filepath.Join("/sys/module", module)); err == nil {
return
}
modprobePath, err := findModprobePath()
if err != nil {
t.Skipf("kernel module %s unavailable: %v", module, err)
}
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
command := exec.CommandContext(ctx, modprobePath, module)
command.Env = os.Environ()
output, err := command.CombinedOutput()
if ctx.Err() != nil {
t.Skipf("modprobe %s timed out: %s", module, string(output))
}
if err != nil {
t.Skipf("kernel module %s unavailable: %v: %s", module, err, string(output))
}
}
func requireUSBIPHost(t *testing.T) {
t.Helper()
if err := ensureHostDriver(); err != nil {
t.Skipf("usbip-host unavailable: %v", err)
}
}
func requireVHCI(t *testing.T) {
t.Helper()
if err := ensureVHCI(); err != nil {
t.Skipf("vhci_hcd unavailable: %v", err)
}
}
func writeSysfsLine(path string, content string) error {
return os.WriteFile(path, []byte(content+"\n"), 0)
}
func newTestUSBGadget(t *testing.T) *testUSBGadget {
t.Helper()
requireRoot(t)
requireKernelModule(t, "configfs")
requireKernelModule(t, "libcomposite")
requireKernelModule(t, "dummy_hcd")
udcs, err := os.ReadDir("/sys/class/udc")
if err != nil {
t.Skipf("USB device controllers unavailable: %v", err)
}
if len(udcs) == 0 {
t.Skip("USB device controllers unavailable")
}
gadget := &testUSBGadget{
path: filepath.Join("/sys/kernel/config/usb_gadget", fmt.Sprintf("codex_usbip_%d", time.Now().UnixNano())),
serial: fmt.Sprintf("codex-usbip-%d", time.Now().UnixNano()),
}
require.NoError(t, os.MkdirAll(filepath.Join(gadget.path, "strings/0x409"), 0o755))
require.NoError(t, os.MkdirAll(filepath.Join(gadget.path, "configs/c.1/strings/0x409"), 0o755))
require.NoError(t, os.Mkdir(filepath.Join(gadget.path, "functions/acm.usb0"), 0o755))
require.NoError(t, writeSysfs(filepath.Join(gadget.path, "idVendor"), "0x1d6b"))
require.NoError(t, writeSysfs(filepath.Join(gadget.path, "idProduct"), "0x0104"))
require.NoError(t, writeSysfs(filepath.Join(gadget.path, "strings/0x409/serialnumber"), gadget.serial))
require.NoError(t, writeSysfs(filepath.Join(gadget.path, "strings/0x409/manufacturer"), "OpenAI"))
require.NoError(t, writeSysfs(filepath.Join(gadget.path, "strings/0x409/product"), "Codex USBIP Test"))
require.NoError(t, writeSysfs(filepath.Join(gadget.path, "configs/c.1/strings/0x409/configuration"), "config-1"))
require.NoError(t, os.Symlink(filepath.Join(gadget.path, "functions/acm.usb0"), filepath.Join(gadget.path, "configs/c.1/acm.usb0")))
require.NoError(t, writeSysfs(filepath.Join(gadget.path, "UDC"), udcs[0].Name()))
require.Eventually(t, func() bool {
devices, err := listUSBDevices()
if err != nil {
return false
}
for i := range devices {
if devices[i].VendorID == 0x1d6b &&
devices[i].ProductID == 0x0104 &&
devices[i].Serial == gadget.serial {
gadget.busid = devices[i].BusID
return true
}
}
return false
}, 5*time.Second, 100*time.Millisecond)
t.Cleanup(func() {
if gadget.busid != "" {
if driver, err := currentDriver(gadget.busid); err == nil {
switch driver {
case "usbip-host":
_ = hostUnbind(gadget.busid)
_ = hostMatchBusID(gadget.busid, false)
_ = bindToDriver(gadget.busid, "usb")
case "usb":
case "":
default:
_ = bindToDriver(gadget.busid, "usb")
}
}
}
_ = writeSysfsLine(filepath.Join(gadget.path, "UDC"), "")
_ = os.Remove(filepath.Join(gadget.path, "configs/c.1/acm.usb0"))
_ = os.Remove(filepath.Join(gadget.path, "functions/acm.usb0"))
_ = os.Remove(filepath.Join(gadget.path, "configs/c.1/strings/0x409"))
_ = os.Remove(filepath.Join(gadget.path, "configs/c.1"))
_ = os.Remove(filepath.Join(gadget.path, "strings/0x409"))
_ = os.Remove(gadget.path)
})
return gadget
}
func TestBuildTargetsDedupesFixedBusID(t *testing.T) {
t.Parallel()
client := &ClientService{
matches: []option.USBIPDeviceMatch{
{BusID: "1-1"},
{VendorID: 0x1d6b, ProductID: 0x0002},
{BusID: "1-1"},
{BusID: "1-2"},
},
}
require.Equal(t, []clientTarget{
{fixedBusID: "1-1"},
{match: option.USBIPDeviceMatch{VendorID: 0x1d6b, ProductID: 0x0002}},
{fixedBusID: "1-2"},
}, client.buildTargets())
}
func TestClientApplyRemoteExportsKeepsActiveBusIDWorker(t *testing.T) {
t.Parallel()
canceled := false
client := &ClientService{
ctx: context.Background(),
logger: newTestLogger(t),
allWorkers: map[string]*clientBusIDWorker{"1-1": {cancel: func() { canceled = true }}},
activeBusIDs: map[string]struct{}{"1-1": {}},
ops: newTestUSBIPOps(t),
}
client.applyRemoteExports(nil)
require.False(t, canceled)
require.Contains(t, client.allWorkers, "1-1")
client.setBusIDActive("1-1", false)
client.applyRemoteExports(nil)
require.True(t, canceled)
require.NotContains(t, client.allWorkers, "1-1")
}
func TestClientApplyControlDeviceStateKeepsActiveMatchedBusyBusID(t *testing.T) {
t.Parallel()
match := option.USBIPDeviceMatch{VendorID: 0x1d6b, ProductID: 0x0002}
target := clientTarget{match: match}
device := newTestDevice("1-1", 0x1d6b, 0x0002, "serial-1", SpeedHigh)
busyDevice := deviceInfoV2FromEntry(device.toDeviceEntry(), "linux-sysfs", "linux-busid:1-1", deviceStateBusy, usbipStatusUsed, "used")
worker := &clientAssignedWorker{target: target, updates: make(chan string, 1)}
client := &ClientService{
matches: []option.USBIPDeviceMatch{match},
targets: []clientTarget{target},
assigned: []string{"1-1"},
assignedWorkers: []*clientAssignedWorker{worker},
activeBusIDs: map[string]struct{}{"1-1": {}},
}
client.applyRemoteDeviceState([]DeviceInfoV2{busyDevice})
require.Equal(t, []string{"1-1"}, client.assigned)
select {
case update := <-worker.updates:
t.Fatalf("unexpected assignment update %q", update)
default:
}
idleWorker := &clientAssignedWorker{target: target, updates: make(chan string, 1)}
idleClient := &ClientService{
matches: []option.USBIPDeviceMatch{match},
targets: []clientTarget{target},
assigned: []string{""},
assignedWorkers: []*clientAssignedWorker{idleWorker},
activeBusIDs: make(map[string]struct{}),
}
idleClient.applyRemoteDeviceState([]DeviceInfoV2{busyDevice})
require.Equal(t, []string{""}, idleClient.assigned)
select {
case update := <-idleWorker.updates:
t.Fatalf("unexpected assignment update %q", update)
default:
}
}
func TestClientShouldRetryBusIDRefreshesImportAllState(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
server := &ServerService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
exports: make(map[string]serverExport),
controlSubs: make(map[uint64]*serverControlConn),
ops: newTestUSBIPOps(t),
}
serverAddr, closeServer := startDispatchServer(t, server)
defer closeServer()
canceled := false
client := &ClientService{
ctx: context.Background(),
logger: newTestLogger(t),
dialer: testDialer{},
serverAddr: serverAddr,
allWorkers: map[string]*clientBusIDWorker{"1-1": {cancel: func() { canceled = true }}},
allDesired: map[string]struct{}{"1-1": {}},
activeBusIDs: make(map[string]struct{}),
ops: newTestUSBIPOps(t),
}
require.False(t, client.shouldRetryBusID(context.Background(), "1-1"))
require.True(t, canceled)
require.NotContains(t, client.allWorkers, "1-1")
require.Empty(t, client.allDesired)
}
func TestClientShouldRetryBusIDKeepsRetryOnRefreshFailure(t *testing.T) {
t.Parallel()
expectedErr := errors.New("devlist unavailable")
canceled := false
client := &ClientService{
ctx: context.Background(),
logger: newTestLogger(t),
dialer: failingDialer{err: expectedErr},
serverAddr: M.ParseSocksaddrHostPort("127.0.0.1", 3240),
allWorkers: map[string]*clientBusIDWorker{"1-1": {cancel: func() { canceled = true }}},
allDesired: map[string]struct{}{"1-1": {}},
activeBusIDs: make(map[string]struct{}),
ops: newTestUSBIPOps(t),
}
require.True(t, client.shouldRetryBusID(context.Background(), "1-1"))
require.False(t, canceled)
require.Contains(t, client.allWorkers, "1-1")
require.Contains(t, client.allDesired, "1-1")
}
func TestAssignMatchedBusIDs(t *testing.T) {
t.Parallel()
match := option.USBIPDeviceMatch{VendorID: 0x1d6b, ProductID: 0x0002}
fixed := newTestDevice("1-1", 0x1d6b, 0x0001, "fixed", SpeedHigh)
first := newTestDevice("1-2", 0x1d6b, 0x0002, "first", SpeedHigh)
second := newTestDevice("1-3", 0x1d6b, 0x0002, "second", SpeedHigh)
entries := []DeviceEntry{
fixed.toDeviceEntry(),
first.toDeviceEntry(),
second.toDeviceEntry(),
}
require.Equal(t, []string{"1-1", "1-3", "1-2"}, assignMatchedBusIDsWithRetained(
[]clientTarget{
{fixedBusID: "1-1"},
{match: match},
{match: match},
},
[]string{"1-1", "1-3", ""},
entries,
nil,
nil,
))
}
func TestLinuxHelpers(t *testing.T) {
t.Parallel()
require.Equal(t, []vhciStatusRecord{
{hub: "hs", port: 0, state: 6},
{hub: "ss", port: 3, state: 4},
}, parseVHCIStatus("hub port sta spd dev sockfd local_busid\nhs 0 6 3 0 0 0\nignored line\nss 3 4 5 0 0 0\n"))
require.Equal(t, SpeedLow, speedCodeFromString("1.5"))
require.Equal(t, SpeedFull, speedCodeFromString("12"))
require.Equal(t, SpeedHigh, speedCodeFromString("480"))
require.Equal(t, SpeedSuper, speedCodeFromString("5000"))
require.Equal(t, SpeedSuperPlus, speedCodeFromString("10000"))
require.Equal(t, SpeedUnknown, speedCodeFromString("25"))
require.Equal(t, "hs", vhciHubForSpeed(SpeedHigh))
require.Equal(t, "ss", vhciHubForSpeed(SpeedSuper))
require.True(t, isUSBDeviceUEvent([]byte("add@/devices/platform/dummy_hcd.0/usb19/19-1\x00ACTION=add\x00SUBSYSTEM=usb\x00DEVTYPE=usb_device\x00")))
require.False(t, isUSBDeviceUEvent([]byte("add@/devices/platform/dummy_hcd.0/usb19/19-1/19-1:1.0\x00ACTION=add\x00SUBSYSTEM=usb\x00DEVTYPE=usb_interface\x00")))
require.False(t, isUSBDeviceUEvent([]byte("ACTION=add\x00SUBSYSTEM=net\x00DEVTYPE=usb_device\x00")))
}
func TestUSBIPConnHandoffDirectTCP(t *testing.T) {
t.Parallel()
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer listener.Close()
accepted := make(chan net.Conn, 1)
go func() {
conn, _ := listener.Accept()
accepted <- conn
}()
conn, err := net.Dial("tcp", listener.Addr().String())
require.NoError(t, err)
defer conn.Close()
acceptedConn := <-accepted
defer acceptedConn.Close()
handoff, err := newUSBIPConnHandoff(conn)
require.NoError(t, err)
defer handoff.Close()
require.False(t, handoff.relay())
require.Equal(t, "direct", handoff.mode())
requireStreamSocketFD(t, handoff.kernelFD())
done := handoff.startRelay(context.Background(), newTestLogger(t), "test", "direct")
_, err = conn.Write([]byte("closed"))
require.Error(t, err)
require.NoError(t, acceptedConn.Close())
select {
case <-done:
case <-time.After(time.Second):
t.Fatal("timed out waiting for direct handoff monitor")
}
}
func TestUSBIPConnHandoffRelaySocketpairCopies(t *testing.T) {
t.Parallel()
left, right := net.Pipe()
defer right.Close()
handoff, err := newUSBIPConnHandoff(opaqueConn{Conn: left})
require.NoError(t, err)
defer handoff.Close()
require.True(t, handoff.relay())
require.Equal(t, "relay", handoff.mode())
requireStreamSocketFD(t, handoff.kernelFD())
kernelConn := duplicateHandoffKernelConn(t, handoff)
defer kernelConn.Close()
setConnDeadline(t, right)
setConnDeadline(t, kernelConn)
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
done := handoff.startRelay(ctx, newTestLogger(t), "test", "relay")
_, err = right.Write([]byte("ping"))
require.NoError(t, err)
requireConnRead(t, kernelConn, []byte("ping"))
_, err = kernelConn.Write([]byte("pong"))
require.NoError(t, err)
requireConnRead(t, right, []byte("pong"))
require.NoError(t, right.Close())
require.NoError(t, kernelConn.Close())
select {
case <-done:
case <-time.After(time.Second):
t.Fatal("timed out waiting for relay handoff")
}
}
func TestServerStartRequiresHostDriver(t *testing.T) {
t.Parallel()
expectedErr := errors.New("host driver unavailable")
server := &ServerService{
ctx: context.Background(),
logger: newTestLogger(t),
exports: make(map[string]serverExport),
controlSubs: make(map[uint64]*serverControlConn),
ops: usbipOps{
ensureHostDriver: func() error { return expectedErr },
},
}
err := server.Start(adapter.StartStateStart)
require.ErrorIs(t, err, expectedErr)
}
func TestClientStartRequiresVHCI(t *testing.T) {
t.Parallel()
expectedErr := errors.New("vhci unavailable")
client := &ClientService{
ctx: context.Background(),
logger: newTestLogger(t),
ops: usbipOps{
ensureVHCI: func() error { return expectedErr },
},
}
err := client.Start(adapter.StartStateStart)
require.ErrorIs(t, err, expectedErr)
}
func TestServerReconcileExportsBindsMatchesAndSkipsHub(t *testing.T) {
t.Parallel()
regular := newTestDevice("1-1", 0x1d6b, 0x0002, "regular", SpeedHigh)
hub := newTestDevice("1-2", 0x1d6b, 0x0002, "hub", SpeedHigh)
hub.DeviceClass = 0x09
store := newTestDeviceStore(regular, hub)
ops := newTestUSBIPOps(t)
var actions []string
ops.listUSBDevices = store.listUSBDevices
ops.currentDriver = func(busid string) (string, error) {
return map[string]string{
"1-1": "usbhid",
"1-2": "hubdrv",
}[busid], nil
}
ops.unbindFromDriver = func(busid, driver string) error {
actions = append(actions, "unbind "+busid+" "+driver)
return nil
}
ops.hostMatchBusID = func(busid string, add bool) error {
actions = append(actions, "match "+busid+" "+map[bool]string{true: "add", false: "del"}[add])
return nil
}
ops.hostBind = func(busid string) error {
actions = append(actions, "hostbind "+busid)
return nil
}
ops.bindToDriver = func(busid, driver string) error {
actions = append(actions, "bind "+busid+" "+driver)
return nil
}
server := &ServerService{
ctx: context.Background(),
logger: newTestLogger(t),
matches: []option.USBIPDeviceMatch{{VendorID: 0x1d6b, ProductID: 0x0002}},
exports: make(map[string]serverExport),
controlSubs: make(map[uint64]*serverControlConn),
ops: ops,
}
changed, err := server.reconcileExports()
require.NoError(t, err)
require.True(t, changed)
require.Equal(t, []string{
"unbind 1-1 usbhid",
"match 1-1 add",
"hostbind 1-1",
}, actions)
require.Equal(t, map[string]serverExport{
"1-1": {
busid: "1-1",
managed: true,
originalDriver: "usbhid",
},
}, server.snapshotExports())
}
func TestServerBindOneRetriesAfterStaleHostMatch(t *testing.T) {
t.Parallel()
device := newTestDevice("1-1", 0x1d6b, 0x0104, "regular", SpeedHigh)
ops := newTestUSBIPOps(t)
var actions []string
bindCalls := 0
ops.currentDriver = func(busid string) (string, error) {
return "usb", nil
}
ops.unbindFromDriver = func(busid, driver string) error {
actions = append(actions, "unbind "+busid+" "+driver)
return nil
}
ops.hostMatchBusID = func(busid string, add bool) error {
actions = append(actions, "match "+busid+" "+map[bool]string{true: "add", false: "del"}[add])
return nil
}
ops.hostBind = func(busid string) error {
bindCalls++
actions = append(actions, "hostbind "+busid)
if bindCalls == 1 {
return &os.PathError{Op: "write", Path: filepath.Join(sysUsbipHostDriver, "bind"), Err: unix.ENODEV}
}
return nil
}
ops.bindToDriver = func(busid, driver string) error {
actions = append(actions, "bind "+busid+" "+driver)
return nil
}
ops.reloadHostDriver = func() error {
actions = append(actions, "reload")
return nil
}
server := &ServerService{
ctx: context.Background(),
logger: newTestLogger(t),
exports: make(map[string]serverExport),
controlSubs: make(map[uint64]*serverControlConn),
ops: ops,
}
require.NoError(t, server.bindOne(&device))
require.Equal(t, []string{
"unbind 1-1 usb",
"match 1-1 add",
"hostbind 1-1",
"match 1-1 del",
"bind 1-1 usb",
"reload",
"unbind 1-1 usb",
"match 1-1 add",
"hostbind 1-1",
}, actions)
require.Equal(t, map[string]serverExport{
"1-1": {
busid: "1-1",
managed: true,
originalDriver: "usb",
},
}, server.snapshotExports())
}
func TestServerReconcileExportsSkipsVHCIDevices(t *testing.T) {
t.Parallel()
physical := newTestDevice("1-1", 0x1d6b, 0x0002, "physical", SpeedHigh)
imported := newTestDevice("3-1", 0x1d6b, 0x0002, "imported", SpeedHigh)
imported.Path = "/sys/devices/platform/vhci_hcd.0/usb3/3-1"
store := newTestDeviceStore(physical, imported)
ops := newTestUSBIPOps(t)
var bound []string
ops.listUSBDevices = store.listUSBDevices
ops.currentDriver = func(busid string) (string, error) {
return "usb", nil
}
ops.unbindFromDriver = func(busid, driver string) error {
bound = append(bound, "unbind "+busid+" "+driver)
return nil
}
ops.hostMatchBusID = func(busid string, add bool) error {
bound = append(bound, "match "+busid)
return nil
}
ops.hostBind = func(busid string) error {
bound = append(bound, "bind "+busid)
return nil
}
server := &ServerService{
ctx: context.Background(),
logger: newTestLogger(t),
matches: []option.USBIPDeviceMatch{{VendorID: 0x1d6b, ProductID: 0x0002}},
exports: make(map[string]serverExport),
controlSubs: make(map[uint64]*serverControlConn),
ops: ops,
}
changed, err := server.reconcileExports()
require.NoError(t, err)
require.True(t, changed)
require.Equal(t, []string{
"unbind 1-1 usb",
"match 1-1",
"bind 1-1",
}, bound)
require.Equal(t, map[string]serverExport{
"1-1": {
busid: "1-1",
managed: true,
originalDriver: "usb",
},
}, server.snapshotExports())
}
func TestServerReconcileExportsReleasesRemovedExports(t *testing.T) {
t.Parallel()
device := newTestDevice("1-1", 0x1d6b, 0x0002, "regular", SpeedHigh)
store := newTestDeviceStore(device)
store.setStatus("1-1", usbipStatusUsed)
ops := newTestUSBIPOps(t)
var actions []string
ops.listUSBDevices = store.listUSBDevices
ops.readUsbipStatus = store.readUsbipStatus
ops.writeUsbipSockfd = func(busid string, fd int) error {
actions = append(actions, "sockfd "+busid)
store.setStatus(busid, usbipStatusAvailable)
return nil
}
ops.hostUnbind = func(busid string) error {
actions = append(actions, "hostunbind "+busid)
return nil
}
ops.hostMatchBusID = func(busid string, add bool) error {
actions = append(actions, "match "+busid+" "+map[bool]string{true: "add", false: "del"}[add])
return nil
}
ops.bindToDriver = func(busid, driver string) error {
actions = append(actions, "bind "+busid+" "+driver)
return nil
}
ops.readSysfsDevice = store.readSysfsDevice
server := &ServerService{
ctx: context.Background(),
logger: newTestLogger(t),
exports: map[string]serverExport{"1-1": {busid: "1-1", managed: true, originalDriver: "usbhid"}},
ops: ops,
}
changed, err := server.reconcileExports()
require.NoError(t, err)
require.True(t, changed)
require.Empty(t, server.snapshotExports())
require.Equal(t, []string{
"sockfd 1-1",
"hostunbind 1-1",
"match 1-1 del",
"bind 1-1 usbhid",
}, actions)
}
func TestServerReleaseExportLeavesCooptedSocketUntouched(t *testing.T) {
t.Parallel()
ops := newTestUSBIPOps(t)
var calls []string
ops.writeUsbipSockfd = func(busid string, fd int) error {
calls = append(calls, fmt.Sprintf("%s=%d", busid, fd))
return nil
}
server := &ServerService{
logger: newTestLogger(t),
exports: map[string]serverExport{"1-1": {busid: "1-1"}},
ops: ops,
}
err := server.releaseExport(serverExport{busid: "1-1"}, true)
require.NoError(t, err)
require.Empty(t, calls)
require.Empty(t, server.snapshotExports())
}
func TestServerReleaseExportRetainsTrackingOnFailure(t *testing.T) {
t.Parallel()
expectedErr := errors.New("host unbind failed")
export := serverExport{
busid: "1-1",
managed: true,
originalDriver: "usbhid",
}
ops := newTestUSBIPOps(t)
ops.readUsbipStatus = func(string) (int, error) {
return usbipStatusAvailable, nil
}
ops.writeUsbipSockfd = func(string, int) error {
return nil
}
ops.hostUnbind = func(string) error {
return expectedErr
}
ops.hostMatchBusID = func(string, bool) error {
return nil
}
ops.bindToDriver = func(string, string) error {
return nil
}
server := &ServerService{
logger: newTestLogger(t),
exports: map[string]serverExport{"1-1": export},
ops: ops,
}
err := server.releaseExport(export, true)
require.ErrorIs(t, err, expectedErr)
require.Equal(t, map[string]serverExport{"1-1": export}, server.snapshotExports())
}
func TestServerCloseSerializesRollbackWithActiveReconcile(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
device := newTestDevice("1-1", 0x1d6b, 0x0002, "regular", SpeedHigh)
listEntered := make(chan struct{})
releaseList := make(chan struct{})
reconcileDone := make(chan error, 1)
closeDone := make(chan error, 1)
var actionsMu sync.Mutex
var actions []string
record := func(action string) {
actionsMu.Lock()
defer actionsMu.Unlock()
actions = append(actions, action)
}
ops := newTestUSBIPOps(t)
ops.listUSBDevices = func() ([]sysfsDevice, error) {
close(listEntered)
<-releaseList
return []sysfsDevice{device}, nil
}
ops.currentDriver = func(string) (string, error) {
return "", nil
}
ops.hostMatchBusID = func(busid string, add bool) error {
if add {
record("match add " + busid)
} else {
record("match del " + busid)
}
return nil
}
ops.hostBind = func(busid string) error {
record("hostbind " + busid)
return nil
}
ops.readSysfsDevice = func(string, string) (sysfsDevice, error) {
return device, nil
}
ops.readUsbipStatus = func(string) (int, error) {
return usbipStatusAvailable, nil
}
ops.hostUnbind = func(busid string) error {
record("hostunbind " + busid)
return nil
}
server := &ServerService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
matches: []option.USBIPDeviceMatch{{BusID: "1-1"}},
exports: make(map[string]serverExport),
controlSubs: make(map[uint64]*serverControlConn),
controlState: make(map[string]DeviceInfoV2),
ops: ops,
}
go func() {
reconcileDone <- server.reconcileAndBroadcast(true)
}()
select {
case <-listEntered:
case <-time.After(3 * time.Second):
t.Fatal("timed out waiting for active reconcile")
}
go func() {
closeDone <- server.Close()
}()
close(releaseList)
select {
case err := <-reconcileDone:
require.NoError(t, err)
case <-time.After(3 * time.Second):
t.Fatal("timed out waiting for reconcile")
}
select {
case err := <-closeDone:
require.NoError(t, err)
case <-time.After(3 * time.Second):
t.Fatal("timed out waiting for close")
}
actionsMu.Lock()
defer actionsMu.Unlock()
require.Equal(t, []string{
"match add 1-1",
"hostbind 1-1",
"hostunbind 1-1",
"match del 1-1",
}, actions)
require.Empty(t, server.snapshotExports())
}
func TestServerReconcileAndBroadcastSkipsAfterCancel(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
cancel()
ops := newTestUSBIPOps(t)
server := &ServerService{
ctx: ctx,
logger: newTestLogger(t),
exports: make(map[string]serverExport),
controlSubs: make(map[uint64]*serverControlConn),
controlState: make(map[string]DeviceInfoV2),
ops: ops,
}
require.NoError(t, server.reconcileAndBroadcast(true))
}
func TestServerUEventLoopReconcilesWhenListenerStarts(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
listener := newTestUSBEventListener()
done := make(chan struct{})
t.Cleanup(func() {
cancel()
_ = listener.Close()
select {
case <-done:
case <-time.After(time.Second):
t.Error("timed out waiting for uevent loop")
}
})
device := newTestDevice("1-1", 0x1d6b, 0x0002, "startup", SpeedHigh)
store := newTestDeviceStore(device)
store.setStatus("1-1", usbipStatusAvailable)
ops := newTestUSBIPOps(t)
ops.newUEventListener = func() (usbEventListener, error) {
return listener, nil
}
ops.listUSBDevices = store.listUSBDevices
ops.currentDriver = func(string) (string, error) {
return "", nil
}
ops.hostMatchBusID = func(string, bool) error {
return nil
}
ops.hostBind = func(string) error {
return nil
}
ops.readUsbipStatus = store.readUsbipStatus
ops.readSysfsDevice = store.readSysfsDevice
server := &ServerService{
ctx: ctx,
logger: newTestLogger(t),
matches: []option.USBIPDeviceMatch{{VendorID: 0x1d6b, ProductID: 0x0002}},
exports: make(map[string]serverExport),
controlSubs: make(map[uint64]*serverControlConn),
controlState: make(map[string]DeviceInfoV2),
ops: ops,
}
go func() {
defer close(done)
server.ueventLoop()
}()
require.Eventually(t, func() bool {
_, ok := server.getExport("1-1")
return ok
}, 3*time.Second, 10*time.Millisecond)
}
func TestServerBuildDevListEntriesFiltersUnavailableAndRefreshFailures(t *testing.T) {
t.Parallel()
available := newTestDevice("1-1", 0x1d6b, 0x0002, "ok", SpeedHigh)
store := newTestDeviceStore(available)
store.setStatus("1-1", usbipStatusAvailable)
store.setStatus("1-2", usbipStatusUsed)
store.setStatus("1-3", usbipStatusAvailable)
ops := newTestUSBIPOps(t)
ops.readUsbipStatus = store.readUsbipStatus
ops.readSysfsDevice = store.readSysfsDevice
server := &ServerService{
logger: newTestLogger(t),
exports: map[string]serverExport{
"1-1": {busid: "1-1"},
"1-2": {busid: "1-2"},
"1-3": {busid: "1-3"},
},
ops: ops,
}
entries := server.buildDevListEntries()
require.Len(t, entries, 1)
require.Equal(t, "1-1", entries[0].Info.BusIDString())
require.Equal(t, "ok", entries[0].Serial)
}
func TestServerHandleImportWithOpaqueConnRelay(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
device := newTestDevice("1-1", 0x1d6b, 0x0002, "serial-1", SpeedHigh)
store := newTestDeviceStore(device)
store.setStatus("1-1", usbipStatusAvailable)
kernelConnCh := make(chan net.Conn, 1)
kernelErrCh := make(chan error, 1)
ops := newTestUSBIPOps(t)
ops.readUsbipStatus = store.readUsbipStatus
ops.readSysfsDevice = store.readSysfsDevice
ops.writeUsbipSockfd = func(busid string, fd int) error {
if fd < 0 {
store.setStatus(busid, usbipStatusAvailable)
store.writeUsbipSockfd(busid, fd)
return nil
}
if busid != "1-1" {
kernelErrCh <- fmt.Errorf("unexpected busid %s", busid)
return nil
}
store.setStatus(busid, usbipStatusUsed)
store.writeUsbipSockfd(busid, fd)
socketType, err := unix.GetsockoptInt(fd, unix.SOL_SOCKET, unix.SO_TYPE)
if err != nil {
kernelErrCh <- err
return nil
}
if socketType != unix.SOCK_STREAM {
kernelErrCh <- fmt.Errorf("unexpected socket type %d", socketType)
return nil
}
kernelConn, err := duplicateNetConnFromFD(uintptr(fd), "usbip-server-test-kernel")
if err != nil {
kernelErrCh <- err
return nil
}
kernelConnCh <- kernelConn
return nil
}
server := &ServerService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
exports: map[string]serverExport{"1-1": {busid: "1-1"}},
controlSubs: make(map[uint64]*serverControlConn),
controlState: make(map[string]DeviceInfoV2),
ops: ops,
}
serverConn, clientConn := net.Pipe()
defer clientConn.Close()
go server.dispatchConn(opaqueConn{Conn: serverConn})
setConnDeadline(t, clientConn)
require.NoError(t, WriteOpReqImport(clientConn, "1-1"))
header, err := ReadOpHeader(clientConn)
require.NoError(t, err)
require.Equal(t, OpRepImport, header.Code)
require.Equal(t, OpStatusOK, header.Status)
_, err = ReadOpRepImportBody(clientConn)
require.NoError(t, err)
require.Eventually(t, func() bool {
return linuxServerControlState(server, "1-1") == deviceStateBusy
}, time.Second, 10*time.Millisecond)
var kernelConn net.Conn
select {
case kernelConn = <-kernelConnCh:
case err = <-kernelErrCh:
require.NoError(t, err)
case <-time.After(3 * time.Second):
t.Fatal("timed out waiting for server relay kernel conn")
}
defer kernelConn.Close()
setConnDeadline(t, kernelConn)
_, err = clientConn.Write([]byte("server-in"))
require.NoError(t, err)
requireConnRead(t, kernelConn, []byte("server-in"))
_, err = kernelConn.Write([]byte("server-out"))
require.NoError(t, err)
requireConnRead(t, clientConn, []byte("server-out"))
require.NoError(t, clientConn.Close())
require.NoError(t, kernelConn.Close())
require.Eventually(t, func() bool {
return store.lastSockfd("1-1") == -1 && linuxServerControlState(server, "1-1") == deviceStateAvailable
}, time.Second, 10*time.Millisecond)
}
func TestServerHandleImportRelayClosesHandoffOnSockfdFailure(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
device := newTestDevice("1-1", 0x1d6b, 0x0002, "serial-1", SpeedHigh)
store := newTestDeviceStore(device)
store.setStatus("1-1", usbipStatusAvailable)
expectedErr := errors.New("sockfd handoff failed")
kernelConnCh := make(chan net.Conn, 1)
kernelErrCh := make(chan error, 1)
ops := newTestUSBIPOps(t)
ops.readUsbipStatus = store.readUsbipStatus
ops.readSysfsDevice = store.readSysfsDevice
ops.writeUsbipSockfd = func(busid string, fd int) error {
if fd < 0 {
return nil
}
if busid != "1-1" {
kernelErrCh <- fmt.Errorf("unexpected busid %s", busid)
return expectedErr
}
kernelConn, err := duplicateNetConnFromFD(uintptr(fd), "usbip-server-sockfd-failure-kernel")
if err != nil {
kernelErrCh <- err
} else {
kernelConnCh <- kernelConn
}
return expectedErr
}
server := &ServerService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
exports: map[string]serverExport{"1-1": {busid: "1-1"}},
controlSubs: make(map[uint64]*serverControlConn),
ops: ops,
}
serverConn, clientConn := net.Pipe()
defer clientConn.Close()
go server.dispatchConn(opaqueConn{Conn: serverConn})
setConnDeadline(t, clientConn)
require.NoError(t, WriteOpReqImport(clientConn, "1-1"))
header, err := ReadOpHeader(clientConn)
require.NoError(t, err)
require.Equal(t, OpRepImport, header.Code)
require.Equal(t, OpStatusError, header.Status)
var kernelConn net.Conn
select {
case kernelConn = <-kernelConnCh:
case err = <-kernelErrCh:
require.NoError(t, err)
case <-time.After(3 * time.Second):
t.Fatal("timed out waiting for failed server relay kernel conn")
}
defer kernelConn.Close()
setConnDeadline(t, kernelConn)
requireConnEOF(t, kernelConn)
}
func TestServerHandleImportRelayClosesHandoffOnReplyFailure(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
device := newTestDevice("1-1", 0x1d6b, 0x0002, "serial-1", SpeedHigh)
store := newTestDeviceStore(device)
store.setStatus("1-1", usbipStatusAvailable)
kernelConnCh := make(chan net.Conn, 1)
kernelErrCh := make(chan error, 1)
rollbackCh := make(chan string, 1)
allowReply := make(chan struct{})
ops := newTestUSBIPOps(t)
ops.readUsbipStatus = store.readUsbipStatus
ops.readSysfsDevice = store.readSysfsDevice
ops.writeUsbipSockfd = func(busid string, fd int) error {
if fd < 0 {
rollbackCh <- busid
return nil
}
if busid != "1-1" {
kernelErrCh <- fmt.Errorf("unexpected busid %s", busid)
<-allowReply
return nil
}
kernelConn, err := duplicateNetConnFromFD(uintptr(fd), "usbip-server-reply-failure-kernel")
if err != nil {
kernelErrCh <- err
} else {
kernelConnCh <- kernelConn
}
<-allowReply
return nil
}
server := &ServerService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
exports: map[string]serverExport{"1-1": {busid: "1-1"}},
controlSubs: make(map[uint64]*serverControlConn),
ops: ops,
}
serverConn, clientConn := net.Pipe()
go server.dispatchConn(opaqueConn{Conn: serverConn})
setConnDeadline(t, clientConn)
require.NoError(t, WriteOpReqImport(clientConn, "1-1"))
var kernelConn net.Conn
select {
case kernelConn = <-kernelConnCh:
case err := <-kernelErrCh:
require.NoError(t, err)
case <-time.After(3 * time.Second):
t.Fatal("timed out waiting for reply-failure relay kernel conn")
}
defer kernelConn.Close()
require.NoError(t, clientConn.Close())
close(allowReply)
select {
case busid := <-rollbackCh:
require.Equal(t, "1-1", busid)
case <-time.After(3 * time.Second):
t.Fatal("timed out waiting for import rollback")
}
setConnDeadline(t, kernelConn)
requireConnEOF(t, kernelConn)
}
func TestServerDispatchConnHandlesControlPingAndChanged(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
server := &ServerService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
exports: make(map[string]serverExport),
controlSubs: make(map[uint64]*serverControlConn),
ops: newTestUSBIPOps(t),
}
serverAddr, closeServer := startDispatchServer(t, server)
defer closeServer()
conn, err := net.Dial("tcp", serverAddr.String())
require.NoError(t, err)
defer conn.Close()
require.NoError(t, WriteControlPreface(conn))
require.NoError(t, WriteControlHello(conn))
ack, err := ReadControlFrame(conn)
require.NoError(t, err)
require.Equal(t, controlFrameAck, ack.Type)
require.Equal(t, controlProtocolVersion, ack.Version)
require.Equal(t, controlCapabilities, ack.Capabilities)
require.Zero(t, ack.Sequence)
snapshotMessage, err := readControlMessage(conn)
require.NoError(t, err)
require.Equal(t, controlFrameDeviceSnapshot, snapshotMessage.Frame.Type)
var snapshot controlDeviceSnapshot
require.NoError(t, unmarshalControlPayload(snapshotMessage.Payload, &snapshot))
require.Empty(t, snapshot.Devices)
require.NoError(t, WriteControlPing(conn))
pong, err := ReadControlFrame(conn)
require.NoError(t, err)
require.Equal(t, controlFramePong, pong.Type)
require.Equal(t, controlProtocolVersion, pong.Version)
server.broadcastControlState(deviceInfoV2Map(server.buildDeviceStateV2()), true)
changed, err := readControlMessage(conn)
require.NoError(t, err)
require.Equal(t, controlFrameDeviceDelta, changed.Frame.Type)
require.Equal(t, uint64(1), changed.Frame.Sequence)
var delta controlDeviceDelta
require.NoError(t, unmarshalControlPayload(changed.Payload, &delta))
require.Equal(t, uint64(1), delta.Sequence)
}
func TestServerRegisterControlConnQueuesSnapshotBeforeBroadcast(t *testing.T) {
t.Parallel()
server := &ServerService{
logger: newTestLogger(t),
exports: make(map[string]serverExport),
controlSubs: make(map[uint64]*serverControlConn),
controlState: make(map[string]DeviceInfoV2),
}
serverConn, clientConn := net.Pipe()
defer serverConn.Close()
defer clientConn.Close()
sub, seq := server.registerControlConn(serverConn, controlCapabilities)
require.Zero(t, seq)
require.Contains(t, server.controlSubs, sub.id)
require.True(t, server.broadcastControlState(map[string]DeviceInfoV2{
"1-1": {BusID: "1-1", State: deviceStateAvailable},
}, true))
first := <-sub.send
require.Equal(t, controlFrameDeviceSnapshot, first.Frame.Type)
require.Zero(t, first.Frame.Sequence)
second := <-sub.send
require.Equal(t, controlFrameDeviceDelta, second.Frame.Type)
require.Equal(t, uint64(1), second.Frame.Sequence)
}
func TestServerReconcileBroadcastsStatusOnlyDeviceDelta(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
device := newTestDevice("1-1", 0x1d6b, 0x0002, "serial-1", SpeedHigh)
store := newTestDeviceStore(device)
store.setStatus("1-1", usbipStatusUsed)
serverOps := newTestUSBIPOps(t)
serverOps.listUSBDevices = store.listUSBDevices
serverOps.readUsbipStatus = store.readUsbipStatus
serverOps.readSysfsDevice = store.readSysfsDevice
server := &ServerService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
matches: []option.USBIPDeviceMatch{{BusID: "1-1"}},
exports: map[string]serverExport{"1-1": {busid: "1-1"}},
controlSubs: make(map[uint64]*serverControlConn),
ops: serverOps,
}
server.refreshControlState()
serverAddr, closeServer := startDispatchServer(t, server)
defer closeServer()
conn, err := net.Dial("tcp", serverAddr.String())
require.NoError(t, err)
defer conn.Close()
setConnDeadline(t, conn)
require.NoError(t, WriteControlPreface(conn))
require.NoError(t, WriteControlHello(conn))
ack, err := ReadControlFrame(conn)
require.NoError(t, err)
require.Equal(t, controlFrameAck, ack.Type)
snapshotMessage, err := readControlMessage(conn)
require.NoError(t, err)
require.Equal(t, controlFrameDeviceSnapshot, snapshotMessage.Frame.Type)
var snapshot controlDeviceSnapshot
require.NoError(t, unmarshalControlPayload(snapshotMessage.Payload, &snapshot))
require.Len(t, snapshot.Devices, 1)
require.Equal(t, "1-1", snapshot.Devices[0].BusID)
require.Equal(t, deviceStateBusy, snapshot.Devices[0].State)
require.Equal(t, usbipStatusUsed, snapshot.Devices[0].StatusCode)
store.setStatus("1-1", usbipStatusAvailable)
require.NoError(t, server.reconcileAndBroadcast(true))
changed, err := readControlMessage(conn)
require.NoError(t, err)
require.Equal(t, controlFrameDeviceDelta, changed.Frame.Type)
require.Equal(t, uint64(1), changed.Frame.Sequence)
var delta controlDeviceDelta
require.NoError(t, unmarshalControlPayload(changed.Payload, &delta))
require.Equal(t, uint64(1), delta.Sequence)
require.Empty(t, delta.Added)
require.Empty(t, delta.Removed)
require.Len(t, delta.Updated, 1)
require.Equal(t, "1-1", delta.Updated[0].BusID)
require.Equal(t, deviceStateAvailable, delta.Updated[0].State)
require.Equal(t, usbipStatusAvailable, delta.Updated[0].StatusCode)
sequence := server.currentControlSequence()
require.NoError(t, server.reconcileAndBroadcast(true))
require.Equal(t, sequence, server.currentControlSequence())
require.NoError(t, conn.SetReadDeadline(time.Now().Add(100*time.Millisecond)))
_, err = readControlMessage(conn)
require.Error(t, err)
var netErr net.Error
require.ErrorAs(t, err, &netErr)
require.True(t, netErr.Timeout())
}
func TestServerControlSnapshotPreservesPendingDelta(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
device := newTestDevice("1-1", 0x1d6b, 0x0002, "serial-1", SpeedHigh)
store := newTestDeviceStore(device)
store.setStatus("1-1", usbipStatusUsed)
serverOps := newTestUSBIPOps(t)
serverOps.listUSBDevices = store.listUSBDevices
serverOps.readUsbipStatus = store.readUsbipStatus
serverOps.readSysfsDevice = store.readSysfsDevice
server := &ServerService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
matches: []option.USBIPDeviceMatch{{BusID: "1-1"}},
exports: map[string]serverExport{"1-1": {busid: "1-1"}},
controlSubs: make(map[uint64]*serverControlConn),
ops: serverOps,
}
server.refreshControlState()
serverAddr, closeServer := startDispatchServer(t, server)
defer closeServer()
firstConn, err := net.Dial("tcp", serverAddr.String())
require.NoError(t, err)
defer firstConn.Close()
setConnDeadline(t, firstConn)
require.NoError(t, WriteControlPreface(firstConn))
require.NoError(t, WriteControlHello(firstConn))
_, err = ReadControlFrame(firstConn)
require.NoError(t, err)
firstSnapshot, err := readControlMessage(firstConn)
require.NoError(t, err)
require.Equal(t, controlFrameDeviceSnapshot, firstSnapshot.Frame.Type)
store.setStatus("1-1", usbipStatusAvailable)
secondConn, err := net.Dial("tcp", serverAddr.String())
require.NoError(t, err)
defer secondConn.Close()
setConnDeadline(t, secondConn)
require.NoError(t, WriteControlPreface(secondConn))
require.NoError(t, WriteControlHello(secondConn))
_, err = ReadControlFrame(secondConn)
require.NoError(t, err)
secondSnapshotMessage, err := readControlMessage(secondConn)
require.NoError(t, err)
require.Equal(t, controlFrameDeviceSnapshot, secondSnapshotMessage.Frame.Type)
var secondSnapshot controlDeviceSnapshot
require.NoError(t, unmarshalControlPayload(secondSnapshotMessage.Payload, &secondSnapshot))
require.Len(t, secondSnapshot.Devices, 1)
require.Equal(t, deviceStateAvailable, secondSnapshot.Devices[0].State)
require.NoError(t, server.reconcileAndBroadcast(true))
changed, err := readControlMessage(firstConn)
require.NoError(t, err)
require.Equal(t, controlFrameDeviceDelta, changed.Frame.Type)
var delta controlDeviceDelta
require.NoError(t, unmarshalControlPayload(changed.Payload, &delta))
require.Len(t, delta.Updated, 1)
require.Equal(t, "1-1", delta.Updated[0].BusID)
require.Equal(t, deviceStateAvailable, delta.Updated[0].State)
}
func TestServerControlLeaseEnablesImportExt(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
device := newTestDevice("1-1", 0x1d6b, 0x0002, "serial-1", SpeedHigh)
store := newTestDeviceStore(device)
store.setStatus("1-1", usbipStatusAvailable)
serverOps := newTestUSBIPOps(t)
serverOps.readUsbipStatus = store.readUsbipStatus
serverOps.readSysfsDevice = store.readSysfsDevice
serverOps.writeUsbipSockfd = store.writeUsbipSockfd
server := &ServerService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
exports: map[string]serverExport{"1-1": {busid: "1-1"}},
controlSubs: make(map[uint64]*serverControlConn),
controlState: make(map[string]DeviceInfoV2),
leasesByBusID: make(map[string]serverImportLease),
ops: serverOps,
}
server.refreshControlState()
serverAddr, closeServer := startDispatchServer(t, server)
defer closeServer()
controlConn, err := net.Dial("tcp", serverAddr.String())
require.NoError(t, err)
defer controlConn.Close()
require.NoError(t, WriteControlPreface(controlConn))
require.NoError(t, WriteControlHello(controlConn))
ack, err := ReadControlFrame(controlConn)
require.NoError(t, err)
require.Equal(t, controlCapabilities, ack.Capabilities)
_, err = readControlMessage(controlConn)
require.NoError(t, err)
require.NoError(t, writeControlMessage(controlConn, controlFrame{
Type: controlFrameLeaseRequest,
Version: controlProtocolVersion,
}, controlLeaseRequest{BusID: "1-1", ClientNonce: 42}))
leaseMessage, err := readControlMessage(controlConn)
require.NoError(t, err)
require.Equal(t, controlFrameLeaseResponse, leaseMessage.Frame.Type)
var lease controlLeaseResponse
require.NoError(t, unmarshalControlPayload(leaseMessage.Payload, &lease))
require.Empty(t, lease.ErrorCode)
require.Equal(t, uint64(42), lease.ClientNonce)
require.NotZero(t, lease.LeaseID)
require.NoError(t, writeControlMessage(controlConn, controlFrame{
Type: controlFrameLeaseRequest,
Version: controlProtocolVersion,
}, controlLeaseRequest{BusID: "1-1", ClientNonce: 43}))
busyMessage, err := readControlMessage(controlConn)
require.NoError(t, err)
var busy controlLeaseResponse
require.NoError(t, unmarshalControlPayload(busyMessage.Payload, &busy))
require.Equal(t, "busy", busy.ErrorCode)
importConn, err := net.Dial("tcp", serverAddr.String())
require.NoError(t, err)
require.NoError(t, WriteOpReqImportExt(importConn, ImportExtRequest{
BusID: "1-1",
LeaseID: lease.LeaseID,
ClientNonce: lease.ClientNonce,
}))
header, err := ReadOpHeader(importConn)
require.NoError(t, err)
require.Equal(t, OpRepImportExt, header.Code)
require.Equal(t, OpStatusOK, header.Status)
info, err := ReadOpRepImportBody(importConn)
require.NoError(t, err)
require.Equal(t, "1-1", info.BusIDString())
require.NoError(t, importConn.Close())
require.True(t, store.hasPositiveSockfd("1-1"))
reuseConn, err := net.Dial("tcp", serverAddr.String())
require.NoError(t, err)
defer reuseConn.Close()
require.NoError(t, WriteOpReqImportExt(reuseConn, ImportExtRequest{
BusID: "1-1",
LeaseID: lease.LeaseID,
ClientNonce: lease.ClientNonce,
}))
header, err = ReadOpHeader(reuseConn)
require.NoError(t, err)
require.Equal(t, OpRepImportExt, header.Code)
require.Equal(t, OpStatusError, header.Status)
}
func TestClientAttemptAttachUsesImportReplyAndVHCIAttach(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
device := newTestDevice("1-1", 0x1d6b, 0x0002, "serial-1", SpeedSuper)
device.BusNum = 7
device.DevNum = 11
store := newTestDeviceStore(device)
store.setStatus("1-1", usbipStatusAvailable)
serverOps := newTestUSBIPOps(t)
serverOps.readUsbipStatus = store.readUsbipStatus
serverOps.readSysfsDevice = store.readSysfsDevice
serverOps.writeUsbipSockfd = store.writeUsbipSockfd
server := &ServerService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
exports: map[string]serverExport{"1-1": {busid: "1-1"}},
controlSubs: make(map[uint64]*serverControlConn),
ops: serverOps,
}
serverAddr, closeServer := startDispatchServer(t, server)
defer closeServer()
clientOps := newTestUSBIPOps(t)
var attachedPort int
var attachedDevID uint32
var attachedSpeed uint32
clientOps.vhciPickFreePort = func(speed uint32, skip map[int]struct{}) (int, error) {
require.Empty(t, skip)
require.Equal(t, SpeedSuper, speed)
return 7, nil
}
clientOps.vhciAttach = func(port int, _ uintptr, devid uint32, speed uint32) error {
attachedPort = port
attachedDevID = devid
attachedSpeed = speed
return nil
}
client := &ClientService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
dialer: testDialer{},
serverAddr: serverAddr,
ops: clientOps,
}
port, done, err := client.attemptAttach(ctx, "1-1")
require.NoError(t, err)
require.NotNil(t, done)
require.Equal(t, 7, port)
require.Equal(t, 7, attachedPort)
info := device.toProtocol()
require.Equal(t, info.DevID(), attachedDevID)
require.Equal(t, SpeedSuper, attachedSpeed)
require.Positive(t, store.lastSockfd("1-1"))
}
func TestClientAttemptAttachUsesImportExtLease(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
controlClient, controlServer := net.Pipe()
defer controlClient.Close()
defer controlServer.Close()
controlSession := newClientControlSession(controlClient, controlCapabilities)
controlErrCh := make(chan error, 1)
go func() {
message, err := readControlMessage(controlServer)
if err != nil {
controlErrCh <- err
return
}
if message.Frame.Type != controlFrameLeaseRequest {
controlErrCh <- fmt.Errorf("unexpected control frame %d", message.Frame.Type)
return
}
var request controlLeaseRequest
if err := unmarshalControlPayload(message.Payload, &request); err != nil {
controlErrCh <- err
return
}
if request.BusID != "1-1" {
controlErrCh <- fmt.Errorf("unexpected lease busid %s", request.BusID)
return
}
controlErrCh <- writeControlMessage(controlServer, controlFrame{
Type: controlFrameLeaseResponse,
Version: controlProtocolVersion,
}, controlLeaseResponse{
BusID: request.BusID,
LeaseID: 55,
ClientNonce: request.ClientNonce,
Generation: 2,
TTLMillis: int64(importLeaseTTL / time.Millisecond),
})
}()
deliverErrCh := make(chan error, 1)
go func() {
message, err := readControlMessage(controlClient)
if err != nil {
deliverErrCh <- err
return
}
if message.Frame.Type != controlFrameLeaseResponse {
deliverErrCh <- fmt.Errorf("unexpected control response %d", message.Frame.Type)
return
}
var response controlLeaseResponse
if err := unmarshalControlPayload(message.Payload, &response); err != nil {
deliverErrCh <- err
return
}
controlSession.deliverLeaseResponse(response)
deliverErrCh <- nil
}()
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer listener.Close()
device := newTestDevice("1-1", 0x1d6b, 0x0002, "serial-1", SpeedHigh)
serverErrCh := make(chan error, 1)
go func() {
conn, acceptErr := listener.Accept()
if acceptErr != nil {
serverErrCh <- acceptErr
return
}
defer conn.Close()
header, readErr := ReadOpHeader(conn)
if readErr != nil {
serverErrCh <- readErr
return
}
if header.Code != OpReqImportExt {
serverErrCh <- fmt.Errorf("unexpected request code 0x%04x", header.Code)
return
}
request, readErr := ReadOpReqImportExtBody(conn)
if readErr != nil {
serverErrCh <- readErr
return
}
if request.BusID != "1-1" || request.LeaseID != 55 || request.ClientNonce != 1 {
serverErrCh <- fmt.Errorf("unexpected import-ext request %+v", request)
return
}
info := device.toProtocol()
serverErrCh <- WriteOpRepImportExt(conn, OpStatusOK, &info)
}()
ops := newTestUSBIPOps(t)
ops.vhciPickFreePort = func(speed uint32, skip map[int]struct{}) (int, error) {
require.Empty(t, skip)
require.Equal(t, SpeedHigh, speed)
return 4, nil
}
ops.vhciAttach = func(port int, _ uintptr, devid uint32, speed uint32) error {
require.Equal(t, 4, port)
info := device.toProtocol()
require.Equal(t, info.DevID(), devid)
require.Equal(t, SpeedHigh, speed)
return nil
}
client := &ClientService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
dialer: testDialer{},
serverAddr: M.SocksaddrFromNet(listener.Addr()),
ops: ops,
}
client.setControlSession(controlSession)
defer client.clearControlSession(controlSession, errClientControlSessionClosed)
port, done, err := client.attemptAttach(ctx, "1-1")
require.NoError(t, err)
require.NotNil(t, done)
require.Equal(t, 4, port)
require.NoError(t, <-controlErrCh)
require.NoError(t, <-deliverErrCh)
require.NoError(t, <-serverErrCh)
}
func TestClientAttemptAttachWithOpaqueConnRelay(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer listener.Close()
device := newTestDevice("1-1", 0x1d6b, 0x0002, "serial-1", SpeedHigh)
serverConnCh := make(chan net.Conn, 1)
serverErrCh := make(chan error, 1)
serverDone := make(chan struct{})
go func() {
conn, acceptErr := listener.Accept()
if acceptErr != nil {
serverErrCh <- acceptErr
return
}
header, readErr := ReadOpHeader(conn)
if readErr != nil {
_ = conn.Close()
serverErrCh <- readErr
return
}
if header.Code != OpReqImport {
_ = conn.Close()
serverErrCh <- fmt.Errorf("unexpected request code 0x%04x", header.Code)
return
}
busid, readErr := ReadOpReqImportBody(conn)
if readErr != nil {
_ = conn.Close()
serverErrCh <- readErr
return
}
if busid != "1-1" {
_ = conn.Close()
serverErrCh <- fmt.Errorf("unexpected busid %s", busid)
return
}
info := device.toProtocol()
if writeErr := WriteOpRepImport(conn, OpStatusOK, &info); writeErr != nil {
_ = conn.Close()
serverErrCh <- writeErr
return
}
serverConnCh <- conn
<-serverDone
_ = conn.Close()
serverErrCh <- nil
}()
defer close(serverDone)
kernelConnCh := make(chan net.Conn, 1)
ops := newTestUSBIPOps(t)
ops.vhciPickFreePort = func(speed uint32, skip map[int]struct{}) (int, error) {
require.Empty(t, skip)
require.Equal(t, SpeedHigh, speed)
return 4, nil
}
ops.vhciAttach = func(port int, fd uintptr, devid uint32, speed uint32) error {
require.Equal(t, 4, port)
requireStreamSocketFD(t, fd)
info := device.toProtocol()
require.Equal(t, info.DevID(), devid)
require.Equal(t, SpeedHigh, speed)
kernelConnCh <- duplicateConnFromFD(t, fd, "usbip-client-test-kernel")
return nil
}
client := &ClientService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
dialer: wrappingDialer{},
serverAddr: M.SocksaddrFromNet(listener.Addr()),
ops: ops,
}
port, done, err := client.attemptAttach(ctx, "1-1")
require.NoError(t, err)
require.NotNil(t, done)
require.Equal(t, 4, port)
var serverConn net.Conn
select {
case serverConn = <-serverConnCh:
case err = <-serverErrCh:
require.NoError(t, err)
case <-time.After(3 * time.Second):
t.Fatal("timed out waiting for server conn")
}
var kernelConn net.Conn
select {
case kernelConn = <-kernelConnCh:
case <-time.After(3 * time.Second):
t.Fatal("timed out waiting for client relay kernel conn")
}
defer kernelConn.Close()
setConnDeadline(t, serverConn)
setConnDeadline(t, kernelConn)
_, err = serverConn.Write([]byte("client-in"))
require.NoError(t, err)
requireConnRead(t, kernelConn, []byte("client-in"))
_, err = kernelConn.Write([]byte("client-out"))
require.NoError(t, err)
requireConnRead(t, serverConn, []byte("client-out"))
require.NoError(t, kernelConn.Close())
select {
case <-done:
case <-time.After(time.Second):
t.Fatal("timed out waiting for client relay handoff")
}
}
func TestClientAttemptAttachRetriesNextPortOnEBUSY(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer listener.Close()
device := newTestDevice("1-1", 0x1d6b, 0x0002, "serial-1", SpeedHigh)
serverErrCh := make(chan error, 1)
go func() {
conn, acceptErr := listener.Accept()
if acceptErr != nil {
serverErrCh <- acceptErr
return
}
defer conn.Close()
header, readErr := ReadOpHeader(conn)
if readErr != nil {
serverErrCh <- readErr
return
}
if header.Code != OpReqImport {
serverErrCh <- fmt.Errorf("unexpected request code 0x%04x", header.Code)
return
}
busid, readErr := ReadOpReqImportBody(conn)
if readErr != nil {
serverErrCh <- readErr
return
}
if busid != "1-1" {
serverErrCh <- fmt.Errorf("unexpected busid %s", busid)
return
}
info := device.toProtocol()
if writeErr := WriteOpRepImport(conn, OpStatusOK, &info); writeErr != nil {
serverErrCh <- writeErr
return
}
buffer := make([]byte, 1)
n, readErr := conn.Read(buffer)
if n != 0 {
serverErrCh <- fmt.Errorf("unexpected server read bytes after relay close: %d", n)
return
}
if !errors.Is(readErr, io.EOF) {
serverErrCh <- readErr
return
}
serverErrCh <- nil
}()
ops := newTestUSBIPOps(t)
var pickCalls int
var attachedPorts []int
ops.vhciPickFreePort = func(speed uint32, skip map[int]struct{}) (int, error) {
require.Equal(t, SpeedHigh, speed)
pickCalls++
if _, skipped := skip[4]; skipped {
return 5, nil
}
return 4, nil
}
ops.vhciAttach = func(port int, fd uintptr, devid uint32, speed uint32) error {
requireStreamSocketFD(t, fd)
info := device.toProtocol()
require.Equal(t, info.DevID(), devid)
require.Equal(t, SpeedHigh, speed)
attachedPorts = append(attachedPorts, port)
if port == 4 {
return unix.EBUSY
}
if port == 5 {
return nil
}
return fmt.Errorf("unexpected vhci port %d", port)
}
client := &ClientService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
dialer: wrappingDialer{},
serverAddr: M.SocksaddrFromNet(listener.Addr()),
ops: ops,
}
port, done, err := client.attemptAttach(ctx, "1-1")
require.NoError(t, err)
require.NotNil(t, done)
require.Equal(t, 5, port)
require.Equal(t, []int{4, 5}, attachedPorts)
require.Equal(t, 2, pickCalls)
select {
case <-done:
case <-time.After(time.Second):
t.Fatal("timed out waiting for client relay handoff")
}
select {
case err = <-serverErrCh:
require.NoError(t, err)
case <-time.After(3 * time.Second):
t.Fatal("timed out waiting for server side close")
}
client.portsAccess.Lock()
_, firstReserved := client.ports[4]
_, secondReserved := client.ports[5]
client.portsAccess.Unlock()
require.False(t, firstReserved)
require.True(t, secondReserved)
}
func TestClientAttemptAttachRelayClosesHandoffOnVHCIAttachFailure(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer listener.Close()
device := newTestDevice("1-1", 0x1d6b, 0x0002, "serial-1", SpeedHigh)
serverErrCh := make(chan error, 1)
go func() {
conn, acceptErr := listener.Accept()
if acceptErr != nil {
serverErrCh <- acceptErr
return
}
defer conn.Close()
header, readErr := ReadOpHeader(conn)
if readErr != nil {
serverErrCh <- readErr
return
}
if header.Code != OpReqImport {
serverErrCh <- fmt.Errorf("unexpected request code 0x%04x", header.Code)
return
}
busid, readErr := ReadOpReqImportBody(conn)
if readErr != nil {
serverErrCh <- readErr
return
}
if busid != "1-1" {
serverErrCh <- fmt.Errorf("unexpected busid %s", busid)
return
}
info := device.toProtocol()
if writeErr := WriteOpRepImport(conn, OpStatusOK, &info); writeErr != nil {
serverErrCh <- writeErr
return
}
buffer := make([]byte, 1)
n, readErr := conn.Read(buffer)
if n != 0 {
serverErrCh <- fmt.Errorf("unexpected server read bytes after attach failure: %d", n)
return
}
if !errors.Is(readErr, io.EOF) {
serverErrCh <- readErr
return
}
serverErrCh <- nil
}()
expectedErr := errors.New("vhci attach failed")
kernelConnCh := make(chan net.Conn, 1)
ops := newTestUSBIPOps(t)
ops.vhciPickFreePort = func(speed uint32, skip map[int]struct{}) (int, error) {
require.Empty(t, skip)
require.Equal(t, SpeedHigh, speed)
return 4, nil
}
ops.vhciAttach = func(port int, fd uintptr, devid uint32, speed uint32) error {
require.Equal(t, 4, port)
requireStreamSocketFD(t, fd)
info := device.toProtocol()
require.Equal(t, info.DevID(), devid)
require.Equal(t, SpeedHigh, speed)
kernelConnCh <- duplicateConnFromFD(t, fd, "usbip-client-vhci-failure-kernel")
return expectedErr
}
client := &ClientService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
dialer: wrappingDialer{},
serverAddr: M.SocksaddrFromNet(listener.Addr()),
ops: ops,
}
port, done, err := client.attemptAttach(ctx, "1-1")
require.Equal(t, -1, port)
require.Nil(t, done)
require.ErrorIs(t, err, expectedErr)
var kernelConn net.Conn
select {
case kernelConn = <-kernelConnCh:
case <-time.After(3 * time.Second):
t.Fatal("timed out waiting for failed client relay kernel conn")
}
defer kernelConn.Close()
setConnDeadline(t, kernelConn)
requireConnEOF(t, kernelConn)
select {
case err = <-serverErrCh:
require.NoError(t, err)
case <-time.After(3 * time.Second):
t.Fatal("timed out waiting for server side close")
}
client.portsAccess.Lock()
_, reserved := client.ports[4]
client.portsAccess.Unlock()
require.False(t, reserved)
}
func TestClientFetchDevListRejectsUnexpectedReplyVersion(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer listener.Close()
serverErr := make(chan error, 1)
go func() {
conn, acceptErr := listener.Accept()
if acceptErr != nil {
serverErr <- acceptErr
return
}
defer conn.Close()
header, readErr := ReadOpHeader(conn)
if readErr != nil {
serverErr <- readErr
return
}
if header.Code != OpReqDevList {
serverErr <- fmt.Errorf("unexpected request code 0x%04x", header.Code)
return
}
if writeErr := binary.Write(conn, binary.BigEndian, OpHeader{
Version: ProtocolVersion + 1,
Code: OpRepDevList,
Status: OpStatusOK,
}); writeErr != nil {
serverErr <- writeErr
return
}
if writeErr := binary.Write(conn, binary.BigEndian, uint32(0)); writeErr != nil {
serverErr <- writeErr
return
}
serverErr <- nil
}()
client := &ClientService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
dialer: testDialer{},
serverAddr: M.SocksaddrFromNet(listener.Addr()),
ops: newTestUSBIPOps(t),
}
entries, err := client.fetchDevList(ctx)
require.Nil(t, entries)
require.ErrorContains(t, err, "unexpected reply version")
require.NoError(t, <-serverErr)
}
func TestClientFetchDevListReturnsOnContextCancelWhileServerStalls(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer listener.Close()
requestReady := make(chan struct{})
serverErr := make(chan error, 1)
go func() {
conn, acceptErr := listener.Accept()
if acceptErr != nil {
serverErr <- acceptErr
return
}
defer conn.Close()
header, readErr := ReadOpHeader(conn)
if readErr != nil {
serverErr <- readErr
return
}
if header.Code != OpReqDevList {
serverErr <- fmt.Errorf("unexpected request code 0x%04x", header.Code)
return
}
close(requestReady)
var buf [1]byte
_, readErr = conn.Read(buf[:])
if readErr == nil {
serverErr <- errors.New("expected client close after cancellation")
return
}
serverErr <- nil
}()
client := &ClientService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
dialer: testDialer{},
serverAddr: M.SocksaddrFromNet(listener.Addr()),
ops: newTestUSBIPOps(t),
}
fetchErr := make(chan error, 1)
go func() {
_, fetchErrValue := client.fetchDevList(ctx)
fetchErr <- fetchErrValue
}()
select {
case <-requestReady:
case <-time.After(3 * time.Second):
t.Fatal("fetchDevList did not reach stalled read path")
}
cancel()
select {
case err = <-fetchErr:
require.Error(t, err)
case <-time.After(3 * time.Second):
t.Fatal("fetchDevList did not exit after cancellation")
}
require.NoError(t, <-serverErr)
}
func TestClientSyncRemoteStateAndResetControlStateRebuildsV2Map(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
device := newTestDevice("1-1", 0x1d6b, 0x0002, "serial-1", SpeedHigh)
entry := device.toDeviceEntry()
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer listener.Close()
serverErr := make(chan error, 1)
go func() {
conn, acceptErr := listener.Accept()
if acceptErr != nil {
serverErr <- acceptErr
return
}
defer conn.Close()
header, readErr := ReadOpHeader(conn)
if readErr != nil {
serverErr <- readErr
return
}
if header.Code != OpReqDevList {
serverErr <- fmt.Errorf("unexpected request code 0x%04x", header.Code)
return
}
serverErr <- WriteOpRepDevList(conn, []DeviceEntry{entry})
}()
client := &ClientService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
dialer: testDialer{},
serverAddr: M.SocksaddrFromNet(listener.Addr()),
matches: []option.USBIPDeviceMatch{{BusID: "unused"}},
ops: newTestUSBIPOps(t),
remoteDevicesV2: map[string]DeviceInfoV2{"stale": {BusID: "stale", State: deviceStateAvailable}},
}
require.NoError(t, client.syncRemoteStateAndResetControlState(ctx))
require.NoError(t, <-serverErr)
client.remoteAccess.Lock()
devices := client.remoteDevicesV2
client.remoteAccess.Unlock()
require.Len(t, devices, 1)
require.Contains(t, devices, "1-1")
require.Equal(t, deviceStateAvailable, devices["1-1"].State)
require.Equal(t, uint16(0x1d6b), devices["1-1"].VendorID)
}
func TestClientAttemptAttachRejectsUnexpectedReplyVersion(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer listener.Close()
device := newTestDevice("1-1", 0x1d6b, 0x0002, "serial-1", SpeedHigh)
info := device.toProtocol()
serverErr := make(chan error, 1)
go func() {
conn, acceptErr := listener.Accept()
if acceptErr != nil {
serverErr <- acceptErr
return
}
defer conn.Close()
header, readErr := ReadOpHeader(conn)
if readErr != nil {
serverErr <- readErr
return
}
if header.Code != OpReqImport {
serverErr <- fmt.Errorf("unexpected request code 0x%04x", header.Code)
return
}
busid, readErr := ReadOpReqImportBody(conn)
if readErr != nil {
serverErr <- readErr
return
}
if busid != "1-1" {
serverErr <- fmt.Errorf("unexpected busid %s", busid)
return
}
if writeErr := binary.Write(conn, binary.BigEndian, OpHeader{
Version: ProtocolVersion + 1,
Code: OpRepImport,
Status: OpStatusOK,
}); writeErr != nil {
serverErr <- writeErr
return
}
if writeErr := binary.Write(conn, binary.BigEndian, &info); writeErr != nil {
serverErr <- writeErr
return
}
serverErr <- nil
}()
ops := newTestUSBIPOps(t)
ops.vhciPickFreePort = func(uint32, map[int]struct{}) (int, error) {
return -1, errors.New("unexpected vhci attach path")
}
client := &ClientService{
ctx: ctx,
cancel: cancel,
logger: newTestLogger(t),
dialer: testDialer{},
serverAddr: M.SocksaddrFromNet(listener.Addr()),
ops: ops,
}
port, done, err := client.attemptAttach(ctx, "1-1")
require.Equal(t, -1, port)
require.Nil(t, done)
require.ErrorContains(t, err, "unexpected reply version")
require.NoError(t, <-serverErr)
}
func TestClientRunControlSessionSyncsAssignmentsOnChanged(t *testing.T) {
t.Parallel()
serverCtx, serverCancel := context.WithCancel(context.Background())
defer serverCancel()
initialDevice := newTestDevice("1-1", 0x1d6b, 0x0002, "first", SpeedHigh)
updatedDevice := newTestDevice("1-2", 0x1d6b, 0x0002, "second", SpeedHigh)
store := newTestDeviceStore(initialDevice)
store.setStatus("1-1", usbipStatusAvailable)
store.setStatus("1-2", usbipStatusAvailable)
serverOps := newTestUSBIPOps(t)
serverOps.readUsbipStatus = store.readUsbipStatus
serverOps.readSysfsDevice = store.readSysfsDevice
server := &ServerService{
ctx: serverCtx,
cancel: serverCancel,
logger: newTestLogger(t),
exports: map[string]serverExport{"1-1": {busid: "1-1"}},
controlSubs: make(map[uint64]*serverControlConn),
ops: serverOps,
}
server.refreshControlState()
serverAddr, closeServer := startDispatchServer(t, server)
defer closeServer()
clientCtx, clientCancel := context.WithCancel(context.Background())
defer clientCancel()
match := option.USBIPDeviceMatch{VendorID: 0x1d6b, ProductID: 0x0002}
client := &ClientService{
ctx: clientCtx,
cancel: clientCancel,
logger: newTestLogger(t),
dialer: testDialer{},
serverAddr: serverAddr,
matches: []option.USBIPDeviceMatch{match},
targets: []clientTarget{{match: match}},
assigned: make([]string, 1),
ops: newTestUSBIPOps(t),
}
errCh := make(chan error, 1)
go func() {
errCh <- client.runControlSession()
}()
require.Eventually(t, func() bool {
client.stateAccess.Lock()
defer client.stateAccess.Unlock()
return client.assigned[0] == "1-1"
}, 3*time.Second, 10*time.Millisecond)
store.setDevices(updatedDevice)
server.deleteExport("1-1")
server.setExport(serverExport{busid: "1-2"})
server.broadcastControlState(deviceInfoV2Map(server.buildDeviceStateV2()), true)
require.Eventually(t, func() bool {
client.stateAccess.Lock()
defer client.stateAccess.Unlock()
return client.assigned[0] == "1-2"
}, 3*time.Second, 10*time.Millisecond)
clientCancel()
select {
case <-errCh:
case <-time.After(3 * time.Second):
t.Fatal("runControlSession did not exit after cancellation")
}
}
func TestUSBIPLinuxSmoke(t *testing.T) {
requireRoot(t)
requireUSBIPHost(t)
requireVHCI(t)
gadget := newTestUSBGadget(t)
device, err := readSysfsDevice(gadget.busid, sysBusDevicePath(gadget.busid))
require.NoError(t, err)
require.Equal(t, gadget.busid, device.BusID)
server := &ServerService{
ctx: context.Background(),
logger: newTestLogger(t),
exports: make(map[string]serverExport),
controlSubs: make(map[uint64]*serverControlConn),
ops: systemUSBIPOps,
}
require.NoError(t, server.bindOne(&device))
_, ok := server.snapshotExports()[gadget.busid]
require.True(t, ok)
driver, err := currentDriver(gadget.busid)
require.NoError(t, err)
require.Equal(t, "usbip-host", driver)
status, err := readUsbipStatus(gadget.busid)
require.NoError(t, err)
require.Equal(t, usbipStatusAvailable, status)
require.NoError(t, hostUnbind(gadget.busid))
require.NoError(t, hostMatchBusID(gadget.busid, false))
require.NoError(t, bindToDriver(gadget.busid, "usb"))
server.deleteExport(gadget.busid)
driver, err = currentDriver(gadget.busid)
require.NoError(t, err)
require.Equal(t, "usb", driver)
}