bee231b9ff
ServerService used to own three mutexes coordinating six pieces of mutable state (exports, busy, controlState, controlSeq, controlSubs, LeaseManager). The lock-ordering rule that fixed the import/attach race in 0c892411a lived only as a comment. Two-mutex (fast/slow) ledger absorbs all of it; the two mutexes are never held simultaneously, so the ordering rule becomes a property of the type instead of a discipline callers must observe. TryReserveForImport closes the window where two concurrent OpReqImports could both pass the busy check before either marked busy.
1205 lines
31 KiB
Go
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.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()
|
|
|
|
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.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())
|
|
}
|