Files
sing-box/service/usbip/linux_interop_test.go
T
世界 8799b10b9a usbip: defer darwin RET_UNLINK until completion; route VHCI through primary; close lease-consume race
Three independent correctness bugs that all violated "decide before
publishing": each surface (wire reply, sysfs path, broadcast state)
committed to a state the code later had to contradict.

darwin: switch IOUSBHostPipe abort to IOUSBHostAbortOptionSynchronous and
add the EP-0 path via abortDeviceRequestsWithOption (previously a silent
no-op for control transfers). After abort returns, await a per-pending
drained channel before writing RET_UNLINK so the late RetSubmit
suppression in finishSubmit has actually taken effect — the drained
channel is allocated lazily by markSubmitUnlinked, so the hot submit
path is unchanged.

linux: collapse the per-controller VHCI abstraction. The kernel exposes
attach/detach/status/status.N only on vhci_hcd.0 with globally unique
port numbers; the previous code wrote vhci_hcd.N/attach for N>0 and
silently failed past the primary controller's port range. Glob status*
on the primary, key reservations on bare port int, and carry the
status-suffix as a diagnostic-only secondary index for Description().

export_ledger: ConsumeLeaseAndReserve no longer publishes busy=true and
then conditionally rolls it back. Read seq under fast first, then make
every decision (including the lease.Generation check) inside one slow
critical section before any busy mutation. This removes the window
where a concurrent BroadcastIfChanged could observe transient busy and
poison l.state with no corrective broadcast on rollback.
2026-06-09 10:42:30 +08:00

1185 lines
31 KiB
Go

//go:build linux
package usbip
import (
"bytes"
"context"
"fmt"
"io"
"net"
"net/netip"
"os"
"os/exec"
"path/filepath"
"sort"
"strconv"
"strings"
"sync"
"testing"
"time"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing/common/json/badoption"
M "github.com/sagernet/sing/common/metadata"
"github.com/stretchr/testify/require"
"golang.org/x/term"
)
const (
testVendorID uint16 = 0x1d6b
testACMProductID uint16 = 0x0104
testHIDProductID uint16 = 0x0105
testUDCCount = 2
testUSBIPTeardownTimeout = 20 * time.Second
testUSBIPTeardownPollInterval = 100 * time.Millisecond
)
var testHIDReportDescriptor = []byte{
0x06, 0x00, 0xff,
0x09, 0x01,
0xa1, 0x01,
0x15, 0x00,
0x26, 0xff, 0x00,
0x75, 0x08,
0x95, 0x08,
0x09, 0x01,
0x81, 0x02,
0x95, 0x08,
0x09, 0x01,
0x91, 0x02,
0xc0,
}
type testUSBIPTools struct {
usbip string
usbipd string
}
type testVirtualFunction struct {
name string
instance string
nodePattern string
configure func(functionPath string) error
}
type testVirtualGadget struct {
path string
serial string
busid string
functions []testVirtualFunction
nodes map[string]string
closeOnce sync.Once
udcName string
}
type testACMGadget struct {
*testVirtualGadget
ttyPath string
}
type testHIDGadget struct {
*testVirtualGadget
hidPath string
}
type rawFile struct {
file *os.File
state *term.State
}
type readResult struct {
data []byte
err error
}
var (
testUDCMu sync.Mutex
testAllocatedUDC = make(map[string]struct{})
)
func requireUSBIPTools(t *testing.T) testUSBIPTools {
t.Helper()
requireRoot(t)
usbipPath, usbipErr := exec.LookPath("usbip")
usbipdPath, usbipdErr := exec.LookPath("usbipd")
if usbipErr != nil || usbipdErr != nil {
t.Skip("usbip and usbipd are required")
}
requireRunnableUSBIPTool(t, usbipPath, "version")
requireRunnableUSBIPTool(t, usbipdPath, "--version")
return testUSBIPTools{
usbip: usbipPath,
usbipd: usbipdPath,
}
}
func requireRunnableUSBIPTool(t *testing.T, path string, args ...string) {
t.Helper()
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
command := exec.CommandContext(ctx, path, args...)
command.Env = os.Environ()
output, err := command.CombinedOutput()
if ctx.Err() != nil {
t.Skipf("%s %s timed out", path, strings.Join(args, " "))
}
if err != nil {
t.Skipf("%s is unavailable: %v\n%s", path, err, strings.TrimSpace(string(output)))
}
}
func currentUDCNames() []string {
entries, err := os.ReadDir("/sys/class/udc")
if err != nil {
return nil
}
names := make([]string, 0, len(entries))
for _, entry := range entries {
names = append(names, entry.Name())
}
sort.Strings(names)
return names
}
func ensureTestUDCs(t *testing.T, minCount int) []string {
t.Helper()
requireKernelModule(t, "configfs")
requireKernelModule(t, "libcomposite")
udcs := currentUDCNames()
if len(udcs) >= minCount {
return udcs
}
modprobePath, err := findModprobePath()
if err != nil {
t.Skipf("dummy_hcd unavailable: %v", err)
}
command := exec.Command(modprobePath, "-r", "dummy_hcd")
command.Env = os.Environ()
_, _ = command.CombinedOutput()
command = exec.Command(modprobePath, "dummy_hcd", "num="+strconv.Itoa(minCount))
command.Env = os.Environ()
output, err := command.CombinedOutput()
if err != nil {
t.Skipf("dummy_hcd with %d UDCs unavailable: %v\n%s", minCount, err, string(output))
}
deadline := time.Now().Add(5 * time.Second)
for time.Now().Before(deadline) {
if udcs = currentUDCNames(); len(udcs) >= minCount {
return udcs
}
time.Sleep(100 * time.Millisecond)
}
t.Skipf("dummy_hcd provided %d UDCs, need %d", len(currentUDCNames()), minCount)
return nil
}
func reserveTestUDC(t *testing.T) string {
t.Helper()
testUDCMu.Lock()
defer testUDCMu.Unlock()
udcs := ensureTestUDCs(t, testUDCCount)
for _, udc := range udcs {
if _, inUse := testAllocatedUDC[udc]; inUse {
continue
}
testAllocatedUDC[udc] = struct{}{}
return udc
}
t.Fatal("no free test UDC available")
return ""
}
func releaseTestUDC(name string) {
if name == "" {
return
}
testUDCMu.Lock()
delete(testAllocatedUDC, name)
testUDCMu.Unlock()
}
func waitForUSBIPTeardown(condition func() bool) bool {
deadline := time.Now().Add(testUSBIPTeardownTimeout)
for {
if condition() {
return true
}
if time.Now().After(deadline) {
return false
}
time.Sleep(testUSBIPTeardownPollInterval)
}
}
func detachUsedVHCIPorts() {
for _, record := range readAllVHCIStatus() {
if record.state == 6 {
_ = writeSysfs(filepath.Join(sysVHCIControllerV0, "detach"), strconv.Itoa(record.port))
}
}
}
func allVHCIPortsIdle() bool {
for _, record := range readAllVHCIStatus() {
if record.state == 6 {
return false
}
}
return true
}
func waitForAllVHCIPortsIdle(t *testing.T) {
t.Helper()
require.Eventually(t, allVHCIPortsIdle, testUSBIPTeardownTimeout, testUSBIPTeardownPollInterval)
}
func waitForVHCIPortIdle(t *testing.T, port int) {
t.Helper()
require.Eventually(t, func() bool {
for _, record := range readAllVHCIStatus() {
if record.port == port && record.state == 6 {
return false
}
}
return true
}, testUSBIPTeardownTimeout, testUSBIPTeardownPollInterval)
}
func waitForUSBIPHostAvailable(busid string) bool {
return waitForUSBIPTeardown(func() bool {
status, err := readUsbipStatus(busid)
if err != nil {
return os.IsNotExist(err) || isMissingUSBDeviceError(err)
}
return status == usbipStatusAvailable
})
}
func waitForDriverAway(busid string, driver string) bool {
return waitForUSBIPTeardown(func() bool {
current, err := currentDriver(busid)
if err != nil {
return os.IsNotExist(err) || isMissingUSBDeviceError(err)
}
return current != driver
})
}
func waitForSysfsPathGone(path string) bool {
return waitForUSBIPTeardown(func() bool {
_, err := os.Stat(path)
return os.IsNotExist(err)
})
}
func waitForGadgetNodesGone(nodes map[string]string) bool {
return waitForUSBIPTeardown(func() bool {
for _, path := range nodes {
if _, err := os.Stat(path); err == nil {
return false
}
}
return true
})
}
func shutdownUSBIPHostDevice(busid string) {
status, err := readUsbipStatus(busid)
if err == nil && status == usbipStatusUsed {
_ = writeSysfs(filepath.Join(sysBusUSBDevices, busid, "usbip_sockfd"), "-1")
_ = waitForUSBIPHostAvailable(busid)
}
if driver, err := currentDriver(busid); err == nil && driver == "usbip-host" {
_ = writeSysfs(filepath.Join(sysUsbipHostDriver, "unbind"), busid)
_ = writeSysfs(filepath.Join(sysUsbipHostDriver, "match_busid"), "del "+busid)
_ = waitForDriverAway(busid, "usbip-host")
}
}
func resetUSBIPInteropState(t *testing.T) {
t.Helper()
requireRoot(t)
detachUsedVHCIPorts()
waitForAllVHCIPortsIdle(t)
devices, err := listUSBDevices()
if err != nil {
return
}
for _, device := range devices {
if !strings.HasPrefix(device.Serial, "codex-usbip-") {
continue
}
shutdownUSBIPHostDevice(device.BusID)
_ = writeSysfs("/sys/bus/usb/drivers/usb/bind", device.BusID)
}
paths, _ := filepath.Glob("/sys/kernel/config/usb_gadget/codex_usbip_*")
for _, path := range paths {
_ = writeSysfsLine(filepath.Join(path, "UDC"), "")
links, _ := filepath.Glob(filepath.Join(path, "configs", "*", "*"))
for _, link := range links {
info, err := os.Lstat(link)
if err == nil && info.Mode()&os.ModeSymlink != 0 {
_ = os.Remove(link)
}
}
functions, _ := filepath.Glob(filepath.Join(path, "functions", "*"))
for _, functionPath := range functions {
_ = os.RemoveAll(functionPath)
}
_ = os.RemoveAll(filepath.Join(path, "configs"))
_ = os.RemoveAll(filepath.Join(path, "strings"))
_ = os.RemoveAll(path)
}
require.Eventually(t, func() bool {
paths, _ := filepath.Glob("/sys/kernel/config/usb_gadget/codex_usbip_*")
return len(paths) == 0
}, testUSBIPTeardownTimeout, testUSBIPTeardownPollInterval)
require.Eventually(t, func() bool {
return len(importedNodeSnapshot("/dev/ttyACM*")) == 0 && len(importedNodeSnapshot("/dev/hidraw*")) == 0
}, testUSBIPTeardownTimeout, testUSBIPTeardownPollInterval)
testUDCMu.Lock()
testAllocatedUDC = make(map[string]struct{})
testUDCMu.Unlock()
}
func loopbackListenAddr() *badoption.Addr {
addr := badoption.Addr(netip.MustParseAddr("127.0.0.1"))
return &addr
}
func pickFreeTCPPort(t *testing.T) uint16 {
t.Helper()
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer listener.Close()
return uint16(listener.Addr().(*net.TCPAddr).Port)
}
func startRealUSBIPServer(t *testing.T, devices []option.USBIPDeviceMatch) (*ServerService, M.Socksaddr) {
t.Helper()
requireUSBIPHost(t)
serviceInstance, err := NewServerService(context.Background(), newTestLogger(t), "usbip-server-test", option.USBIPServerServiceOptions{
ListenOptions: option.ListenOptions{
Listen: loopbackListenAddr(),
ListenPort: pickFreeTCPPort(t),
},
Devices: devices,
})
require.NoError(t, err)
server := serviceInstance.(*ServerService)
require.NoError(t, server.Start(adapter.StartStateStart))
t.Cleanup(func() {
_ = server.Close()
})
return server, M.SocksaddrFromNet(server.listener.TCPListener().Addr())
}
func startRealUSBIPClient(t *testing.T, destination M.Socksaddr, devices []option.USBIPDeviceMatch) *ClientService {
t.Helper()
requireVHCI(t)
serviceInstance, err := NewClientService(context.Background(), newTestLogger(t), "usbip-client-test", option.USBIPClientServiceOptions{
ServerOptions: option.ServerOptions{
Server: destination.AddrString(),
ServerPort: destination.Port,
},
Devices: devices,
})
require.NoError(t, err)
client := serviceInstance.(*ClientService)
require.NoError(t, client.Start(adapter.StartStateStart))
t.Cleanup(func() {
_ = client.Close()
})
return client
}
func runCommand(t *testing.T, name string, args ...string) string {
t.Helper()
command := exec.Command(name, args...)
command.Env = os.Environ()
output, err := command.CombinedOutput()
require.NoErrorf(t, err, "%s %s\n%s", name, strings.Join(args, " "), string(output))
return string(output)
}
func runUSBIP(t *testing.T, tools testUSBIPTools, args ...string) string {
t.Helper()
return runCommand(t, tools.usbip, args...)
}
func startUSBIPD(t *testing.T, tools testUSBIPTools, port uint16) {
t.Helper()
ctx, cancel := context.WithCancel(context.Background())
command := exec.CommandContext(ctx, tools.usbipd, "--debug", "--tcp-port", strconv.Itoa(int(port)))
var output bytes.Buffer
command.Stdout = &output
command.Stderr = &output
command.Env = os.Environ()
require.NoError(t, command.Start())
waitForTCPPort(t, port)
t.Cleanup(func() {
cancel()
done := make(chan error, 1)
go func() {
done <- command.Wait()
}()
select {
case <-time.After(5 * time.Second):
_ = command.Process.Kill()
<-done
case <-done:
}
})
}
func waitForTCPPort(t *testing.T, port uint16) {
t.Helper()
address := net.JoinHostPort("127.0.0.1", strconv.Itoa(int(port)))
require.Eventually(t, func() bool {
conn, err := net.DialTimeout("tcp", address, 200*time.Millisecond)
if err != nil {
return false
}
_ = conn.Close()
return true
}, 5*time.Second, 100*time.Millisecond)
}
func snapshotPaths(pattern string) map[string]struct{} {
paths, _ := filepath.Glob(pattern)
snapshot := make(map[string]struct{}, len(paths))
for _, path := range paths {
snapshot[path] = struct{}{}
}
return snapshot
}
func newPaths(pattern string, before map[string]struct{}) []string {
paths, _ := filepath.Glob(pattern)
var out []string
for _, path := range paths {
if _, found := before[path]; found {
continue
}
out = append(out, path)
}
sort.Strings(out)
return out
}
func waitForNewPath(t *testing.T, pattern string, before map[string]struct{}) string {
t.Helper()
var found string
require.Eventually(t, func() bool {
paths := newPaths(pattern, before)
if len(paths) == 0 {
return false
}
found = paths[0]
return true
}, 5*time.Second, 100*time.Millisecond)
return found
}
func importedNodeSnapshot(pattern string) map[string]struct{} {
paths, _ := filepath.Glob(pattern)
snapshot := make(map[string]struct{}, len(paths))
for _, path := range paths {
if isVHCINode(path) {
snapshot[path] = struct{}{}
}
}
return snapshot
}
func isVHCINode(path string) bool {
base := filepath.Base(path)
var sysfsPath string
switch {
case strings.HasPrefix(base, "ttyACM"):
sysfsPath = filepath.Join("/sys/class/tty", base, "device")
case strings.HasPrefix(base, "hidraw"):
sysfsPath = filepath.Join("/sys/class/hidraw", base, "device")
default:
return false
}
realPath, err := filepath.EvalSymlinks(sysfsPath)
if err != nil {
return false
}
return strings.Contains(realPath, "vhci_hcd")
}
func waitForNewImportedNode(t *testing.T, pattern string, before map[string]struct{}) string {
t.Helper()
var found string
require.Eventually(t, func() bool {
paths, _ := filepath.Glob(pattern)
var candidates []string
for _, path := range paths {
if !isVHCINode(path) {
continue
}
if _, present := before[path]; present {
continue
}
candidates = append(candidates, path)
}
if len(candidates) == 0 {
return false
}
sort.Strings(candidates)
found = candidates[0]
return true
}, 20*time.Second, 100*time.Millisecond)
return found
}
func waitForImportedNodePresent(t *testing.T, pattern string, path string) string {
t.Helper()
if path != "" {
if _, err := os.Stat(path); err == nil && isVHCINode(path) {
return path
}
}
var found string
require.Eventually(t, func() bool {
paths, _ := filepath.Glob(pattern)
var candidates []string
for _, candidate := range paths {
if !isVHCINode(candidate) {
continue
}
candidates = append(candidates, candidate)
}
if len(candidates) == 0 {
return false
}
sort.Strings(candidates)
found = candidates[0]
return true
}, 10*time.Second, 100*time.Millisecond)
return found
}
func waitForPathGone(t *testing.T, path string) {
t.Helper()
require.Eventually(t, func() bool {
_, err := os.Stat(path)
return os.IsNotExist(err)
}, 10*time.Second, 100*time.Millisecond)
}
func ensureNoNewImportedNode(t *testing.T, pattern string, before map[string]struct{}, duration time.Duration) {
t.Helper()
deadline := time.Now().Add(duration)
for time.Now().Before(deadline) {
paths, _ := filepath.Glob(pattern)
for _, path := range paths {
if !isVHCINode(path) {
continue
}
if _, present := before[path]; !present {
t.Fatalf("unexpected imported node %s", path)
}
}
time.Sleep(100 * time.Millisecond)
}
}
func usedVHCIPorts(t *testing.T) map[int]struct{} {
t.Helper()
ports := make(map[int]struct{})
for _, record := range readAllVHCIStatus() {
if record.state == 6 {
ports[record.port] = struct{}{}
}
}
return ports
}
func waitForNewUsedVHCIPort(t *testing.T, before map[int]struct{}) int {
t.Helper()
var port int
require.Eventually(t, func() bool {
for _, record := range readAllVHCIStatus() {
if record.state != 6 {
continue
}
if _, found := before[record.port]; found {
continue
}
port = record.port
return true
}
return false
}, 10*time.Second, 100*time.Millisecond)
return port
}
func readExactlyAsync(reader io.Reader, size int) <-chan readResult {
results := make(chan readResult, 1)
go func() {
buffer := make([]byte, size)
_, err := io.ReadFull(reader, buffer)
results <- readResult{
data: buffer,
err: err,
}
}()
return results
}
func requireRead(t *testing.T, results <-chan readResult, expected []byte) {
t.Helper()
select {
case result := <-results:
require.NoError(t, result.err)
require.Equal(t, expected, result.data)
case <-time.After(5 * time.Second):
t.Fatal("timed out waiting for device I/O")
}
}
func readExactlyWithin(reader io.Reader, size int, timeout time.Duration) ([]byte, error) {
results := readExactlyAsync(reader, size)
select {
case result := <-results:
return result.data, result.err
case <-time.After(timeout):
return nil, context.DeadlineExceeded
}
}
func openRawTTY(t *testing.T, path string) *rawFile {
t.Helper()
var lastErr error
deadline := time.Now().Add(testUSBIPTeardownTimeout)
for {
file, err := os.OpenFile(path, os.O_RDWR, 0)
if err == nil {
state, err := term.MakeRaw(int(file.Fd()))
if err == nil {
return &rawFile{
file: file,
state: state,
}
}
lastErr = err
_ = file.Close()
} else {
lastErr = err
}
if time.Now().After(deadline) {
require.NoErrorf(t, lastErr, "open raw tty %s", path)
}
time.Sleep(testUSBIPTeardownPollInterval)
}
}
func (r *rawFile) Close() {
if r == nil || r.file == nil {
return
}
_ = term.Restore(int(r.file.Fd()), r.state)
_ = r.file.Close()
}
func newTestVirtualGadget(t *testing.T, productID uint16, productName string, functions []testVirtualFunction) *testVirtualGadget {
t.Helper()
requireRoot(t)
requireKernelModule(t, "configfs")
requireKernelModule(t, "libcomposite")
suffix := strconv.FormatInt(time.Now().UnixNano(), 10)
resolvedFunctions := make([]testVirtualFunction, len(functions))
for i, function := range functions {
resolvedFunctions[i] = function
typeName, _, hasInstance := strings.Cut(function.name, ".")
if hasInstance {
resolvedFunctions[i].instance = typeName + ".codex" + suffix
} else {
resolvedFunctions[i].instance = function.name + "codex" + suffix
}
}
snapshots := make(map[string]map[string]struct{})
for _, function := range resolvedFunctions {
if function.nodePattern == "" {
continue
}
snapshots[function.name] = snapshotPaths(function.nodePattern)
}
gadget := &testVirtualGadget{
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()),
functions: resolvedFunctions,
nodes: make(map[string]string, len(resolvedFunctions)),
udcName: reserveTestUDC(t),
}
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, writeSysfs(filepath.Join(gadget.path, "idVendor"), fmt.Sprintf("0x%04x", testVendorID)))
require.NoError(t, writeSysfs(filepath.Join(gadget.path, "idProduct"), fmt.Sprintf("0x%04x", productID)))
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"), productName))
require.NoError(t, writeSysfs(filepath.Join(gadget.path, "configs/c.1/strings/0x409/configuration"), "config-1"))
for _, function := range resolvedFunctions {
functionPath := filepath.Join(gadget.path, "functions", function.instance)
require.NoError(t, os.Mkdir(functionPath, 0o755))
if function.configure != nil {
require.NoError(t, function.configure(functionPath))
}
require.NoError(t, os.Symlink(functionPath, filepath.Join(gadget.path, "configs/c.1", function.instance)))
}
require.NoError(t, writeSysfs(filepath.Join(gadget.path, "UDC"), gadget.udcName))
require.Eventually(t, func() bool {
devices, err := listUSBDevices()
if err != nil {
return false
}
for i := range devices {
if devices[i].VendorID == testVendorID &&
devices[i].ProductID == productID &&
devices[i].Serial == gadget.serial {
gadget.busid = devices[i].BusID
return true
}
}
return false
}, 10*time.Second, 100*time.Millisecond)
for _, function := range functions {
if function.nodePattern == "" {
continue
}
gadget.nodes[function.name] = waitForNewPath(t, function.nodePattern, snapshots[function.name])
}
t.Cleanup(func() {
gadget.Close()
})
return gadget
}
func (g *testVirtualGadget) Close() {
g.closeOnce.Do(func() {
defer releaseTestUDC(g.udcName)
if g.busid != "" {
shutdownUSBIPHostDevice(g.busid)
}
_ = writeSysfsLine(filepath.Join(g.path, "UDC"), "")
if g.busid != "" {
_ = waitForSysfsPathGone(filepath.Join(sysBusUSBDevices, g.busid))
}
_ = waitForGadgetNodesGone(g.nodes)
for _, function := range g.functions {
_ = os.Remove(filepath.Join(g.path, "configs/c.1", function.instance))
}
for _, function := range g.functions {
_ = os.RemoveAll(filepath.Join(g.path, "functions", function.instance))
}
_ = os.RemoveAll(filepath.Join(g.path, "configs/c.1/strings/0x409"))
_ = os.RemoveAll(filepath.Join(g.path, "configs/c.1"))
_ = os.RemoveAll(filepath.Join(g.path, "strings/0x409"))
_ = os.RemoveAll(g.path)
_ = waitForSysfsPathGone(g.path)
})
}
func newTestACMGadget(t *testing.T) *testACMGadget {
t.Helper()
gadget := newTestVirtualGadget(t, testACMProductID, "Codex USBIP ACM", []testVirtualFunction{{
name: "acm.usb0",
nodePattern: "/dev/ttyGS*",
}})
return &testACMGadget{
testVirtualGadget: gadget,
ttyPath: gadget.nodes["acm.usb0"],
}
}
func newTestHIDGadget(t *testing.T) *testHIDGadget {
t.Helper()
gadget := newTestVirtualGadget(t, testHIDProductID, "Codex USBIP HID", []testVirtualFunction{{
name: "hid.usb0",
nodePattern: "/dev/hidg*",
configure: func(functionPath string) error {
if err := writeSysfs(functionPath+"/protocol", "0"); err != nil {
return err
}
if err := writeSysfs(functionPath+"/subclass", "0"); err != nil {
return err
}
if err := writeSysfs(functionPath+"/report_length", "8"); err != nil {
return err
}
return os.WriteFile(functionPath+"/report_desc", testHIDReportDescriptor, 0o644)
},
}})
return &testHIDGadget{
testVirtualGadget: gadget,
hidPath: gadget.nodes["hid.usb0"],
}
}
func (g *testACMGadget) exerciseImportedIO(t *testing.T, importedTTY string) {
t.Helper()
gadgetTTY := openRawTTY(t, g.ttyPath)
imported := openRawTTY(t, importedTTY)
defer gadgetTTY.Close()
defer imported.Close()
gadgetToHost := []byte("acm-g2h!")
hostToGadget := []byte("acm-h2g!")
hostRead := readExactlyAsync(imported.file, len(gadgetToHost))
_, err := gadgetTTY.file.Write(gadgetToHost)
require.NoError(t, err)
requireRead(t, hostRead, gadgetToHost)
gadgetRead := readExactlyAsync(gadgetTTY.file, len(hostToGadget))
_, err = imported.file.Write(hostToGadget)
require.NoError(t, err)
requireRead(t, gadgetRead, hostToGadget)
}
func (g *testHIDGadget) exerciseImportedIO(t *testing.T, importedHID string) {
t.Helper()
gadgetToHost := []byte{1, 2, 3, 4, 5, 6, 7, 8}
hostToGadget := []byte{8, 7, 6, 5, 4, 3, 2, 1}
require.Eventually(t, func() bool {
gadgetHID, err := os.OpenFile(g.hidPath, os.O_RDWR, 0)
if err != nil {
return false
}
defer gadgetHID.Close()
imported, err := os.OpenFile(importedHID, os.O_RDWR, 0)
if err != nil {
return false
}
defer imported.Close()
if _, err = gadgetHID.Write(gadgetToHost); err != nil {
return false
}
readBack, err := readExactlyWithin(imported, len(gadgetToHost), time.Second)
if err != nil || !bytes.Equal(readBack, gadgetToHost) {
return false
}
if _, err = imported.Write(hostToGadget); err != nil {
return false
}
readBack, err = readExactlyWithin(gadgetHID, len(hostToGadget), time.Second)
return err == nil && bytes.Equal(readBack, hostToGadget)
}, 10*time.Second, 100*time.Millisecond)
}
func bindWithOfficialUSBIP(t *testing.T, tools testUSBIPTools, busid string) {
t.Helper()
requireUSBIPHost(t)
if driver, err := currentDriver(busid); err == nil && driver == "usbip-host" {
return
}
runUSBIP(t, tools, "bind", "--busid="+busid)
require.Eventually(t, func() bool {
driver, err := currentDriver(busid)
return err == nil && driver == "usbip-host"
}, 5*time.Second, 100*time.Millisecond)
}
func unbindWithOfficialUSBIP(t *testing.T, tools testUSBIPTools, busid string) {
t.Helper()
runUSBIP(t, tools, "unbind", "--busid="+busid)
require.Eventually(t, func() bool {
driver, err := currentDriver(busid)
return err == nil && driver != "usbip-host"
}, 5*time.Second, 100*time.Millisecond)
}
func TestUSBIPInteropOurServerWithOfficialClientACM(t *testing.T) {
requireRoot(t)
resetUSBIPInteropState(t)
tools := requireUSBIPTools(t)
requireVHCI(t)
gadget := newTestACMGadget(t)
server, address := startRealUSBIPServer(t, []option.USBIPDeviceMatch{{Serial: gadget.serial}})
beforePorts := usedVHCIPorts(t)
beforeTTY := importedNodeSnapshot("/dev/ttyACM*")
listOutput := runUSBIP(t, tools, "--tcp-port", strconv.Itoa(int(address.Port)), "list", "--remote=127.0.0.1")
require.Contains(t, listOutput, gadget.busid)
runUSBIP(t, tools, "--tcp-port", strconv.Itoa(int(address.Port)), "attach", "--remote=127.0.0.1", "--busid="+gadget.busid)
port := waitForNewUsedVHCIPort(t, beforePorts)
importedTTY := waitForNewImportedNode(t, "/dev/ttyACM*", beforeTTY)
gadget.exerciseImportedIO(t, importedTTY)
portOutput := runUSBIP(t, tools, "port")
require.Contains(t, portOutput, fmt.Sprintf("Port %02d", port))
runUSBIP(t, tools, "detach", "--port="+strconv.Itoa(port))
waitForVHCIPortIdle(t, port)
waitForPathGone(t, importedTTY)
require.NoError(t, server.Close())
_ = server
}
func TestUSBIPInteropOurServerWithOfficialClientHID(t *testing.T) {
requireRoot(t)
resetUSBIPInteropState(t)
tools := requireUSBIPTools(t)
requireVHCI(t)
gadget := newTestHIDGadget(t)
server, address := startRealUSBIPServer(t, []option.USBIPDeviceMatch{{Serial: gadget.serial}})
beforePorts := usedVHCIPorts(t)
beforeHID := importedNodeSnapshot("/dev/hidraw*")
listOutput := runUSBIP(t, tools, "--tcp-port", strconv.Itoa(int(address.Port)), "list", "--remote=127.0.0.1")
require.Contains(t, listOutput, gadget.busid)
runUSBIP(t, tools, "--tcp-port", strconv.Itoa(int(address.Port)), "attach", "--remote=127.0.0.1", "--busid="+gadget.busid)
port := waitForNewUsedVHCIPort(t, beforePorts)
importedHID := waitForNewImportedNode(t, "/dev/hidraw*", beforeHID)
gadget.exerciseImportedIO(t, importedHID)
portOutput := runUSBIP(t, tools, "port")
require.Contains(t, portOutput, fmt.Sprintf("Port %02d", port))
runUSBIP(t, tools, "detach", "--port="+strconv.Itoa(port))
waitForVHCIPortIdle(t, port)
waitForPathGone(t, importedHID)
require.NoError(t, server.Close())
}
func TestUSBIPInteropOurClientWithOfficialServerACM(t *testing.T) {
requireRoot(t)
resetUSBIPInteropState(t)
tools := requireUSBIPTools(t)
requireVHCI(t)
gadget := newTestACMGadget(t)
bindWithOfficialUSBIP(t, tools, gadget.busid)
t.Cleanup(func() {
unbindWithOfficialUSBIP(t, tools, gadget.busid)
})
port := pickFreeTCPPort(t)
startUSBIPD(t, tools, port)
beforeTTY := importedNodeSnapshot("/dev/ttyACM*")
client := startRealUSBIPClient(t, M.ParseSocksaddrHostPort("127.0.0.1", port), []option.USBIPDeviceMatch{{
VendorID: option.USBIPHexUint16(testVendorID),
ProductID: option.USBIPHexUint16(testACMProductID),
}})
_ = client
importedTTY := waitForNewImportedNode(t, "/dev/ttyACM*", beforeTTY)
gadget.exerciseImportedIO(t, importedTTY)
require.NoError(t, client.Close())
waitForAllVHCIPortsIdle(t)
waitForPathGone(t, importedTTY)
}
func TestUSBIPInteropOurClientWithOfficialServerHID(t *testing.T) {
requireRoot(t)
resetUSBIPInteropState(t)
tools := requireUSBIPTools(t)
requireVHCI(t)
gadget := newTestHIDGadget(t)
bindWithOfficialUSBIP(t, tools, gadget.busid)
t.Cleanup(func() {
unbindWithOfficialUSBIP(t, tools, gadget.busid)
})
port := pickFreeTCPPort(t)
startUSBIPD(t, tools, port)
beforeHID := importedNodeSnapshot("/dev/hidraw*")
client := startRealUSBIPClient(t, M.ParseSocksaddrHostPort("127.0.0.1", port), []option.USBIPDeviceMatch{{
VendorID: option.USBIPHexUint16(testVendorID),
ProductID: option.USBIPHexUint16(testHIDProductID),
}})
_ = client
importedHID := waitForNewImportedNode(t, "/dev/hidraw*", beforeHID)
gadget.exerciseImportedIO(t, importedHID)
require.NoError(t, client.Close())
waitForAllVHCIPortsIdle(t)
waitForPathGone(t, importedHID)
}
func TestUSBIPOfficialServerHasStaticDiscoveryOnly(t *testing.T) {
requireRoot(t)
resetUSBIPInteropState(t)
tools := requireUSBIPTools(t)
requireVHCI(t)
first := newTestACMGadget(t)
bindWithOfficialUSBIP(t, tools, first.busid)
t.Cleanup(func() {
unbindWithOfficialUSBIP(t, tools, first.busid)
})
port := pickFreeTCPPort(t)
startUSBIPD(t, tools, port)
beforeTTY := importedNodeSnapshot("/dev/ttyACM*")
client := startRealUSBIPClient(t, M.ParseSocksaddrHostPort("127.0.0.1", port), nil)
_ = client
importedTTY := waitForNewImportedNode(t, "/dev/ttyACM*", beforeTTY)
first.exerciseImportedIO(t, importedTTY)
second := newTestHIDGadget(t)
bindWithOfficialUSBIP(t, tools, second.busid)
t.Cleanup(func() {
unbindWithOfficialUSBIP(t, tools, second.busid)
})
beforeHID := importedNodeSnapshot("/dev/hidraw*")
ensureNoNewImportedNode(t, "/dev/hidraw*", beforeHID, 3*time.Second)
require.NoError(t, client.Close())
waitForAllVHCIPortsIdle(t)
waitForPathGone(t, importedTTY)
}
func TestUSBIPControlHotplugACMReattach(t *testing.T) {
requireRoot(t)
resetUSBIPInteropState(t)
requireVHCI(t)
ensureTestUDCs(t, testUDCCount)
server, address := startRealUSBIPServer(t, []option.USBIPDeviceMatch{{
VendorID: option.USBIPHexUint16(testVendorID),
ProductID: option.USBIPHexUint16(testACMProductID),
}})
client := startRealUSBIPClient(t, address, []option.USBIPDeviceMatch{{
VendorID: option.USBIPHexUint16(testVendorID),
ProductID: option.USBIPHexUint16(testACMProductID),
}})
_ = client
beforeTTY := importedNodeSnapshot("/dev/ttyACM*")
first := newTestACMGadget(t)
firstImportedTTY := waitForNewImportedNode(t, "/dev/ttyACM*", beforeTTY)
first.exerciseImportedIO(t, firstImportedTTY)
first.Close()
waitForPathGone(t, firstImportedTTY)
require.Eventually(t, func() bool {
return len(server.ledger.AvailableExports()) == 0
}, 5*time.Second, 100*time.Millisecond)
secondBefore := importedNodeSnapshot("/dev/ttyACM*")
second := newTestACMGadget(t)
secondImportedTTY := waitForNewImportedNode(t, "/dev/ttyACM*", secondBefore)
second.exerciseImportedIO(t, secondImportedTTY)
require.NoError(t, client.Close())
waitForAllVHCIPortsIdle(t)
waitForPathGone(t, secondImportedTTY)
require.NoError(t, server.Close())
}
func TestUSBIPControlImportAllACMAndHID(t *testing.T) {
requireRoot(t)
resetUSBIPInteropState(t)
requireVHCI(t)
ensureTestUDCs(t, testUDCCount)
server, address := startRealUSBIPServer(t, []option.USBIPDeviceMatch{
{VendorID: option.USBIPHexUint16(testVendorID), ProductID: option.USBIPHexUint16(testACMProductID)},
{VendorID: option.USBIPHexUint16(testVendorID), ProductID: option.USBIPHexUint16(testHIDProductID)},
})
client := startRealUSBIPClient(t, address, nil)
_ = client
beforeTTY := importedNodeSnapshot("/dev/ttyACM*")
beforeHID := importedNodeSnapshot("/dev/hidraw*")
acm := newTestACMGadget(t)
hid := newTestHIDGadget(t)
importedTTY := waitForNewImportedNode(t, "/dev/ttyACM*", beforeTTY)
importedHID := waitForNewImportedNode(t, "/dev/hidraw*", beforeHID)
importedTTY = waitForImportedNodePresent(t, "/dev/ttyACM*", importedTTY)
importedHID = waitForImportedNodePresent(t, "/dev/hidraw*", importedHID)
acm.exerciseImportedIO(t, importedTTY)
hid.exerciseImportedIO(t, importedHID)
require.NoError(t, client.Close())
waitForAllVHCIPortsIdle(t)
waitForPathGone(t, importedTTY)
waitForPathGone(t, importedHID)
require.NoError(t, server.Close())
}