9c8147f234
- Route every reserved-state mutation through withInventoryWrite; the lease insert path now broadcasts so extended subscribers see the busy transition immediately instead of waiting for the next mutation. Folds cleanupExpiredLocked changes into the broadcast decision on every IssueLease early-reject path. Rename the field to inventory and force read/write/write-quiet sites through dedicated accessors so future callers cannot bypass the invariant. - Mirror the linux clone-then-swap pattern in darwinExportHost.Reconcile via cloneDarwinExport, so the ledger's unlocked Snapshot / LeaseCheck reads never observe a half-mutated stale flag. Documented as docs/adr/0001-export-pointer-immutability.md and on the Export interface; both hosts now share the applyStaleClones helper. - Rewrite 27 if (_, )?err := …; err != nil sites to assign-then-check per .claude/rules/go-syntax.md. - Delete parse / builder / tautological tests forbidden by .claude/rules/code-test.md (option/usbip_test.go, iso_scheduler_test.go, usbhost_darwin_status_test.go). - Tag nine more usbip files with linux || (darwin && cgo); fixes pre-existing windows / android lint failures because the protocol types were only consumed by tagged code.
1190 lines
31 KiB
Go
1190 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 {
|
|
_, err := os.Stat(path)
|
|
if 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 != "" {
|
|
_, err := os.Stat(path)
|
|
if 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 {
|
|
err := writeSysfs(functionPath+"/protocol", "0")
|
|
if err != nil {
|
|
return err
|
|
}
|
|
err = writeSysfs(functionPath+"/subclass", "0")
|
|
if err != nil {
|
|
return err
|
|
}
|
|
err = writeSysfs(functionPath+"/report_length", "8")
|
|
if 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())
|
|
}
|