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

1205 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
removeOnce 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() {
records, err := readVHCIStatus()
if err != nil {
return
}
for _, record := range records {
if record.state == 6 {
_ = vhciDetach(record.port)
}
}
}
func allVHCIPortsIdle() bool {
records, err := readVHCIStatus()
if err != nil {
return true
}
for _, record := range records {
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 {
used, err := vhciPortUsed(port)
return err == nil && !used
}, 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 {
_ = writeUsbipSockfd(busid, -1)
_ = waitForUSBIPHostAvailable(busid)
}
if driver, err := currentDriver(busid); err == nil && driver == "usbip-host" {
_ = hostUnbind(busid)
_ = hostMatchBusID(busid, false)
_ = 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)
_ = bindToDriver(device.BusID, "usb")
}
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.listen.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()
records, err := readVHCIStatus()
require.NoError(t, err)
ports := make(map[int]struct{})
for _, record := range records {
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 {
records, err := readVHCIStatus()
if err != nil {
return false
}
for _, record := range records {
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 openBinaryDevice(t *testing.T, path string) *os.File {
t.Helper()
file, err := os.OpenFile(path, os.O_RDWR, 0)
require.NoError(t, err)
return file
}
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(sysBusDevicePath(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.currentExports()) == 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())
}