diff --git a/service/usbip/client_linux.go b/service/usbip/client_linux.go index d45d23042..cb3c7cc89 100644 --- a/service/usbip/client_linux.go +++ b/service/usbip/client_linux.go @@ -74,6 +74,9 @@ type ClientService struct { portsMu sync.Mutex ports map[int]struct{} + + activeMu sync.Mutex + activeBusIDs map[string]struct{} } func NewClientService(ctx context.Context, logger log.ContextLogger, tag string, options option.USBIPClientServiceOptions) (adapter.Service, error) { @@ -94,16 +97,17 @@ func NewClientService(ctx context.Context, logger log.ContextLogger, tag string, } ctx, cancel := context.WithCancel(ctx) return &ClientService{ - Adapter: boxService.NewAdapter(C.TypeUSBIPClient, tag), - ctx: ctx, - cancel: cancel, - logger: logger, - dialer: outboundDialer, - serverAddr: options.ServerOptions.Build(), - matches: options.Devices, - ops: systemUSBIPOps, - allWorkers: make(map[string]*clientBusIDWorker), - ports: make(map[int]struct{}), + Adapter: boxService.NewAdapter(C.TypeUSBIPClient, tag), + ctx: ctx, + cancel: cancel, + logger: logger, + dialer: outboundDialer, + serverAddr: options.ServerOptions.Build(), + matches: options.Devices, + ops: systemUSBIPOps, + allWorkers: make(map[string]*clientBusIDWorker), + ports: make(map[int]struct{}), + activeBusIDs: make(map[string]struct{}), }, nil } @@ -317,6 +321,9 @@ func (c *ClientService) applyRemoteExports(entries []DeviceEntry) { if _, ok := desired[busid]; ok { continue } + if c.isBusIDActive(busid) { + continue + } stopWorkers = append(stopWorkers, worker) delete(c.allWorkers, busid) } @@ -482,6 +489,9 @@ func (c *ClientService) fetchDevList(ctx context.Context) ([]DeviceEntry, error) if err != nil { return nil, E.Cause(err, "read OP_REP_DEVLIST header") } + if header.Version != ProtocolVersion { + return nil, E.New("unexpected reply version 0x", hex16(header.Version)) + } if header.Code != OpRepDevList || header.Status != OpStatusOK { return nil, E.New("OP_REP_DEVLIST status=", header.Status, " code=0x", hex16(header.Code)) } @@ -502,8 +512,9 @@ func (c *ClientService) runBusIDLoop(ctx context.Context, busid, description str continue } c.logger.Info("attached ", busid, " → vhci port ", port) - c.trackPort(port, true) + c.setBusIDActive(busid, true) c.watchPort(ctx, port, busid) + c.setBusIDActive(busid, false) c.trackPort(port, false) if err := ctx.Err(); err != nil { return @@ -530,6 +541,9 @@ func (c *ClientService) attemptAttach(ctx context.Context, busid string) (int, e if err != nil { return -1, E.Cause(err, "read OP_REP_IMPORT header") } + if header.Version != ProtocolVersion { + return -1, E.New("unexpected reply version 0x", hex16(header.Version)) + } if header.Code != OpRepImport { return -1, E.New("unexpected reply code 0x", hex16(header.Code)) } @@ -555,7 +569,11 @@ func (c *ClientService) attemptAttach(ctx context.Context, busid string) (int, e if err != nil { return -1, err } + if !c.reservePort(port) { + return -1, E.New("vhci port ", port, " already reserved") + } if err := c.ops.vhciAttach(port, file.Fd(), info.DevID(), info.Speed); err != nil { + c.trackPort(port, false) return -1, E.Cause(err, "vhci attach") } return port, nil @@ -564,6 +582,9 @@ func (c *ClientService) attemptAttach(ctx context.Context, busid string) (int, e func (c *ClientService) watchPort(ctx context.Context, port int, busid string) { ticker := time.NewTicker(2 * time.Second) defer ticker.Stop() + seenUsed := false + settleDeadline := time.NewTimer(10 * time.Second) + defer settleDeadline.Stop() for { select { case <-ctx.Done(): @@ -571,13 +592,26 @@ func (c *ClientService) watchPort(ctx context.Context, port int, busid string) { c.logger.Warn("detach port ", port, " (", busid, "): ", err) } return + case <-settleDeadline.C: + if !seenUsed { + c.logger.Warn("vhci port ", port, " never reached used state; reattaching ", busid) + return + } case <-ticker.C: used, err := c.ops.vhciPortUsed(port) if err != nil { c.logger.Debug("poll port ", port, ": ", err) continue } - if !used { + if used { + if !seenUsed { + c.logger.Debug("vhci port ", port, " entered used state for ", busid) + } + seenUsed = true + continue + } + if seenUsed { + c.logger.Debug("vhci port ", port, " left used state for ", busid) return } } @@ -587,13 +621,53 @@ func (c *ClientService) watchPort(ctx context.Context, port int, busid string) { func (c *ClientService) trackPort(port int, add bool) { c.portsMu.Lock() defer c.portsMu.Unlock() + if c.ports == nil { + c.ports = make(map[int]struct{}) + } if add { + c.logger.Debug("reserve vhci port ", port) c.ports[port] = struct{}{} } else { + c.logger.Debug("release vhci port ", port) delete(c.ports, port) } } +func (c *ClientService) reservePort(port int) bool { + c.portsMu.Lock() + defer c.portsMu.Unlock() + if c.ports == nil { + c.ports = make(map[int]struct{}) + } + if _, exists := c.ports[port]; exists { + c.logger.Debug("vhci port ", port, " already reserved locally") + return false + } + c.logger.Debug("reserve vhci port ", port) + c.ports[port] = struct{}{} + return true +} + +func (c *ClientService) setBusIDActive(busid string, active bool) { + c.activeMu.Lock() + defer c.activeMu.Unlock() + if c.activeBusIDs == nil { + c.activeBusIDs = make(map[string]struct{}) + } + if active { + c.activeBusIDs[busid] = struct{}{} + } else { + delete(c.activeBusIDs, busid) + } +} + +func (c *ClientService) isBusIDActive(busid string) bool { + c.activeMu.Lock() + defer c.activeMu.Unlock() + _, exists := c.activeBusIDs[busid] + return exists +} + func isBusIDOnlyMatch(m option.USBIPDeviceMatch) bool { return m.BusID != "" && m.VendorID == 0 && m.ProductID == 0 && m.Serial == "" } diff --git a/service/usbip/linux_interop_test.go b/service/usbip/linux_interop_test.go index ebeae0286..200424c5f 100644 --- a/service/usbip/linux_interop_test.go +++ b/service/usbip/linux_interop_test.go @@ -32,6 +32,7 @@ const ( testVendorID uint16 = 0x1d6b testACMProductID uint16 = 0x0104 testHIDProductID uint16 = 0x0105 + testUDCCount = 2 ) var testHIDReportDescriptor = []byte{ @@ -57,6 +58,7 @@ type testUSBIPTools struct { type testVirtualFunction struct { name string + instance string nodePattern string configure func(functionPath string) error } @@ -92,6 +94,11 @@ type readResult struct { err error } +var ( + testUDCMu sync.Mutex + testAllocatedUDC = make(map[string]struct{}) +) + func requireUSBIPTools(t *testing.T) testUSBIPTools { t.Helper() requireRoot(t) @@ -107,6 +114,152 @@ func requireUSBIPTools(t *testing.T) testUSBIPTools { } } +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() + require.NoError(t, 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() + require.NoErrorf(t, err, "modprobe dummy_hcd num=%d\n%s", minCount, string(output)) + + require.Eventually(t, func() bool { + return len(currentUDCNames()) >= minCount + }, 5*time.Second, 100*time.Millisecond) + + return currentUDCNames() +} + +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 resetUSBIPInteropState(t *testing.T) { + t.Helper() + requireRoot(t) + + records, err := readVHCIStatus() + if err == nil { + for _, record := range records { + if record.state == 6 { + _ = vhciDetach(record.port) + } + } + require.Eventually(t, func() bool { + records, err = readVHCIStatus() + if err != nil { + return false + } + for _, record := range records { + if record.state == 6 { + return false + } + } + return true + }, 10*time.Second, 100*time.Millisecond) + } + + devices, err := listUSBDevices() + if err != nil { + return + } + for _, device := range devices { + if !strings.HasPrefix(device.Serial, "codex-usbip-") { + continue + } + driver, err := currentDriver(device.BusID) + if err != nil || driver != "usbip-host" { + continue + } + _ = hostUnbind(device.BusID) + _ = hostMatchBusID(device.BusID, false) + _ = 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 + }, 10*time.Second, 100*time.Millisecond) + require.Eventually(t, func() bool { + return len(importedNodeSnapshot("/dev/ttyACM*")) == 0 && len(importedNodeSnapshot("/dev/hidraw*")) == 0 + }, 10*time.Second, 100*time.Millisecond) + + testUDCMu.Lock() + testAllocatedUDC = make(map[string]struct{}) + testUDCMu.Unlock() +} + func loopbackListenAddr() *badoption.Addr { addr := badoption.Addr(netip.MustParseAddr("127.0.0.1")) return &addr @@ -305,6 +458,35 @@ func waitForNewImportedNode(t *testing.T, pattern string, before map[string]stru 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 } @@ -399,6 +581,16 @@ func requireRead(t *testing.T, results <-chan readResult, expected []byte) { } } +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() @@ -434,14 +626,21 @@ func newTestVirtualGadget(t *testing.T, productID uint16, productName string, fu requireKernelModule(t, "configfs") requireKernelModule(t, "libcomposite") - requireKernelModule(t, "dummy_hcd") - udcs, err := os.ReadDir("/sys/class/udc") - require.NoError(t, err) - require.NotEmpty(t, udcs) + 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 functions { + for _, function := range resolvedFunctions { if function.nodePattern == "" { continue } @@ -451,9 +650,9 @@ func newTestVirtualGadget(t *testing.T, productID uint16, productName string, fu 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: functions, - nodes: make(map[string]string, len(functions)), - udcName: udcs[0].Name(), + functions: resolvedFunctions, + nodes: make(map[string]string, len(resolvedFunctions)), + udcName: reserveTestUDC(t), } require.NoError(t, os.MkdirAll(filepath.Join(gadget.path, "strings/0x409"), 0o755)) @@ -465,13 +664,13 @@ func newTestVirtualGadget(t *testing.T, productID uint16, productName string, fu 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 functions { - functionPath := filepath.Join(gadget.path, "functions", function.name) - require.NoError(t, os.MkdirAll(functionPath, 0o755)) + 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.name))) + 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)) @@ -508,6 +707,8 @@ func newTestVirtualGadget(t *testing.T, productID uint16, productName string, fu func (g *testVirtualGadget) Close() { g.closeOnce.Do(func() { + defer releaseTestUDC(g.udcName) + if g.busid != "" { if driver, err := currentDriver(g.busid); err == nil && driver == "usbip-host" { _ = hostUnbind(g.busid) @@ -518,15 +719,40 @@ func (g *testVirtualGadget) Close() { _ = writeSysfsLine(filepath.Join(g.path, "UDC"), "") for _, function := range g.functions { - _ = os.Remove(filepath.Join(g.path, "configs/c.1", function.name)) + _ = os.Remove(filepath.Join(g.path, "configs/c.1", function.instance)) } for _, function := range g.functions { - _ = os.RemoveAll(filepath.Join(g.path, "functions", function.name)) + _ = 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) + + deadline := time.Now().Add(5 * time.Second) + for time.Now().Before(deadline) { + if _, err := os.Stat(g.path); err == nil { + time.Sleep(100 * time.Millisecond) + continue + } + if g.busid != "" { + if _, err := os.Stat(sysBusDevicePath(g.busid)); err == nil { + time.Sleep(100 * time.Millisecond) + continue + } + } + remainingNode := false + for _, path := range g.nodes { + if _, err := os.Stat(path); err == nil { + remainingNode = true + break + } + } + if !remainingNode { + return + } + time.Sleep(100 * time.Millisecond) + } }) } @@ -593,37 +819,63 @@ func (g *testACMGadget) exerciseImportedIO(t *testing.T, importedTTY string) { func (g *testHIDGadget) exerciseImportedIO(t *testing.T, importedHID string) { t.Helper() - gadgetHID := openBinaryDevice(t, g.hidPath) - imported := openBinaryDevice(t, importedHID) - defer gadgetHID.Close() - defer imported.Close() - gadgetToHost := []byte{1, 2, 3, 4, 5, 6, 7, 8} hostToGadget := []byte{8, 7, 6, 5, 4, 3, 2, 1} - hostRead := readExactlyAsync(imported, len(gadgetToHost)) - _, err := gadgetHID.Write(gadgetToHost) - require.NoError(t, err) - requireRead(t, hostRead, gadgetToHost) + require.Eventually(t, func() bool { + gadgetHID, err := os.OpenFile(g.hidPath, os.O_RDWR, 0) + if err != nil { + return false + } + defer gadgetHID.Close() - gadgetRead := readExactlyAsync(gadgetHID, len(hostToGadget)) - _, err = imported.Write(hostToGadget) - require.NoError(t, err) - requireRead(t, gadgetRead, hostToGadget) + 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() + + 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) require.NoError(t, ensureVHCI()) @@ -651,6 +903,7 @@ func TestUSBIPInteropOurServerWithOfficialClientACM(t *testing.T) { func TestUSBIPInteropOurServerWithOfficialClientHID(t *testing.T) { requireRoot(t) + resetUSBIPInteropState(t) tools := requireUSBIPTools(t) require.NoError(t, ensureVHCI()) @@ -676,6 +929,7 @@ func TestUSBIPInteropOurServerWithOfficialClientHID(t *testing.T) { func TestUSBIPInteropOurClientWithOfficialServerACM(t *testing.T) { requireRoot(t) + resetUSBIPInteropState(t) tools := requireUSBIPTools(t) require.NoError(t, ensureVHCI()) @@ -704,6 +958,7 @@ func TestUSBIPInteropOurClientWithOfficialServerACM(t *testing.T) { func TestUSBIPInteropOurClientWithOfficialServerHID(t *testing.T) { requireRoot(t) + resetUSBIPInteropState(t) tools := requireUSBIPTools(t) require.NoError(t, ensureVHCI()) @@ -732,6 +987,7 @@ func TestUSBIPInteropOurClientWithOfficialServerHID(t *testing.T) { func TestUSBIPOfficialServerHasStaticDiscoveryOnly(t *testing.T) { requireRoot(t) + resetUSBIPInteropState(t) tools := requireUSBIPTools(t) require.NoError(t, ensureVHCI()) @@ -766,9 +1022,11 @@ func TestUSBIPOfficialServerHasStaticDiscoveryOnly(t *testing.T) { func TestUSBIPControlHotplugACMReattach(t *testing.T) { requireRoot(t) + resetUSBIPInteropState(t) require.NoError(t, ensureVHCI()) + ensureTestUDCs(t, testUDCCount) - _, address := startRealUSBIPServer(t, []option.USBIPDeviceMatch{{ + server, address := startRealUSBIPServer(t, []option.USBIPDeviceMatch{{ VendorID: option.USBIPHexUint16(testVendorID), ProductID: option.USBIPHexUint16(testACMProductID), }}) @@ -785,6 +1043,9 @@ func TestUSBIPControlHotplugACMReattach(t *testing.T) { 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) @@ -794,7 +1055,9 @@ func TestUSBIPControlHotplugACMReattach(t *testing.T) { func TestUSBIPControlImportAllACMAndHID(t *testing.T) { requireRoot(t) + resetUSBIPInteropState(t) require.NoError(t, ensureVHCI()) + ensureTestUDCs(t, testUDCCount) _, address := startRealUSBIPServer(t, []option.USBIPDeviceMatch{ {VendorID: option.USBIPHexUint16(testVendorID), ProductID: option.USBIPHexUint16(testACMProductID)}, @@ -811,6 +1074,8 @@ func TestUSBIPControlImportAllACMAndHID(t *testing.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) diff --git a/service/usbip/linux_test.go b/service/usbip/linux_test.go index 01fdcee6f..deaf7c9d8 100644 --- a/service/usbip/linux_test.go +++ b/service/usbip/linux_test.go @@ -4,10 +4,12 @@ package usbip import ( "context" + "encoding/binary" "errors" "fmt" "net" "os" + "os/exec" "path/filepath" "slices" "sync" @@ -18,7 +20,6 @@ import ( "github.com/sagernet/sing-box/log" "github.com/sagernet/sing-box/option" M "github.com/sagernet/sing/common/metadata" - "github.com/sagernet/sing/common/shell" "github.com/stretchr/testify/require" ) @@ -199,6 +200,21 @@ func newTestUSBIPOps(t *testing.T) usbipOps { } func newTestLogger() log.ContextLogger { + if os.Getenv("CODEX_USBIP_TEST_LOG") != "" { + factory := log.NewDefaultFactory( + context.Background(), + log.Formatter{ + BaseTime: time.Now(), + DisableColors: true, + }, + os.Stderr, + "", + nil, + false, + ) + factory.SetLevel(log.LevelTrace) + return factory.NewLogger("usbip") + } return log.NewNOPFactory().NewLogger("usbip") } @@ -261,12 +277,23 @@ func requireRoot(t *testing.T) { func requireKernelModule(t *testing.T, module string) { t.Helper() + if _, err := os.Stat(filepath.Join("/sys/module", module)); err == nil { + return + } modprobePath, err := findModprobePath() require.NoError(t, err) - output, err := shell.Exec(modprobePath, module).Read() - require.NoErrorf(t, err, "modprobe %s: %s", module, output) + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + command := exec.CommandContext(ctx, modprobePath, module) + command.Env = os.Environ() + output, err := command.CombinedOutput() + if ctx.Err() != nil { + t.Fatalf("modprobe %s timed out: %s", module, string(output)) + } + require.NoErrorf(t, err, "modprobe %s: %s", module, string(output)) } func writeSysfsLine(path string, content string) error { @@ -364,6 +391,30 @@ func TestBuildTargetsDedupesFixedBusID(t *testing.T) { }, client.buildTargets()) } +func TestClientApplyRemoteExportsKeepsActiveBusIDWorker(t *testing.T) { + t.Parallel() + + canceled := false + client := &ClientService{ + ctx: context.Background(), + logger: newTestLogger(), + allWorkers: map[string]*clientBusIDWorker{"1-1": {cancel: func() { canceled = true }}}, + activeBusIDs: map[string]struct{}{"1-1": {}}, + ops: newTestUSBIPOps(t), + } + + client.applyRemoteExports(nil) + + require.False(t, canceled) + require.Contains(t, client.allWorkers, "1-1") + + client.setBusIDActive("1-1", false) + client.applyRemoteExports(nil) + + require.True(t, canceled) + require.NotContains(t, client.allWorkers, "1-1") +} + func TestAssignMatchedBusIDs(t *testing.T) { t.Parallel() @@ -502,6 +553,59 @@ func TestServerReconcileExportsBindsMatchesAndSkipsHub(t *testing.T) { }, server.snapshotExports()) } +func TestServerReconcileExportsSkipsVHCIDevices(t *testing.T) { + t.Parallel() + + physical := newTestDevice("1-1", 0x1d6b, 0x0002, "physical", SpeedHigh) + imported := newTestDevice("3-1", 0x1d6b, 0x0002, "imported", SpeedHigh) + imported.Path = "/sys/devices/platform/vhci_hcd.0/usb3/3-1" + + store := newTestDeviceStore(physical, imported) + ops := newTestUSBIPOps(t) + var bound []string + ops.listUSBDevices = store.listUSBDevices + ops.currentDriver = func(busid string) (string, error) { + return "usb", nil + } + ops.unbindFromDriver = func(busid, driver string) error { + bound = append(bound, "unbind "+busid+" "+driver) + return nil + } + ops.hostMatchBusID = func(busid string, add bool) error { + bound = append(bound, "match "+busid) + return nil + } + ops.hostBind = func(busid string) error { + bound = append(bound, "bind "+busid) + return nil + } + + server := &ServerService{ + ctx: context.Background(), + logger: newTestLogger(), + matches: []option.USBIPDeviceMatch{{VendorID: 0x1d6b, ProductID: 0x0002}}, + exports: make(map[string]serverExport), + controlSubs: make(map[uint64]*serverControlConn), + ops: ops, + } + + changed, err := server.reconcileExports() + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, []string{ + "unbind 1-1 usb", + "match 1-1", + "bind 1-1", + }, bound) + require.Equal(t, map[string]serverExport{ + "1-1": { + busid: "1-1", + managed: true, + originalDriver: "usb", + }, + }, server.snapshotExports()) +} + func TestServerReconcileExportsReleasesRemovedExports(t *testing.T) { t.Parallel() @@ -547,6 +651,63 @@ func TestServerReconcileExportsReleasesRemovedExports(t *testing.T) { }, actions) } +func TestServerReleaseExportLeavesCooptedSocketUntouched(t *testing.T) { + t.Parallel() + + ops := newTestUSBIPOps(t) + var calls []string + ops.writeUsbipSockfd = func(busid string, fd int) error { + calls = append(calls, fmt.Sprintf("%s=%d", busid, fd)) + return nil + } + + server := &ServerService{ + logger: newTestLogger(), + exports: map[string]serverExport{"1-1": {busid: "1-1"}}, + ops: ops, + } + + err := server.releaseExport(serverExport{busid: "1-1"}, true) + require.NoError(t, err) + require.Empty(t, calls) + require.Empty(t, server.snapshotExports()) +} + +func TestServerReleaseExportRetainsTrackingOnFailure(t *testing.T) { + t.Parallel() + + expectedErr := errors.New("host unbind failed") + export := serverExport{ + busid: "1-1", + managed: true, + originalDriver: "usbhid", + } + + ops := newTestUSBIPOps(t) + ops.writeUsbipSockfd = func(string, int) error { + return nil + } + ops.hostUnbind = func(string) error { + return expectedErr + } + ops.hostMatchBusID = func(string, bool) error { + return nil + } + ops.bindToDriver = func(string, string) error { + return nil + } + + server := &ServerService{ + logger: newTestLogger(), + exports: map[string]serverExport{"1-1": export}, + ops: ops, + } + + err := server.releaseExport(export, true) + require.ErrorIs(t, err, expectedErr) + require.Equal(t, map[string]serverExport{"1-1": export}, server.snapshotExports()) +} + func TestServerBuildDevListEntriesFiltersUnavailableAndRefreshFailures(t *testing.T) { t.Parallel() @@ -682,6 +843,210 @@ func TestClientAttemptAttachUsesImportReplyAndVHCIAttach(t *testing.T) { require.Positive(t, store.lastSockfd("1-1")) } +func TestClientFetchDevListRejectsUnexpectedReplyVersion(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + defer listener.Close() + + serverErr := make(chan error, 1) + go func() { + conn, acceptErr := listener.Accept() + if acceptErr != nil { + serverErr <- acceptErr + return + } + defer conn.Close() + + header, readErr := ReadOpHeader(conn) + if readErr != nil { + serverErr <- readErr + return + } + if header.Code != OpReqDevList { + serverErr <- fmt.Errorf("unexpected request code 0x%s", hex16(header.Code)) + return + } + if writeErr := binary.Write(conn, binary.BigEndian, OpHeader{ + Version: ProtocolVersion + 1, + Code: OpRepDevList, + Status: OpStatusOK, + }); writeErr != nil { + serverErr <- writeErr + return + } + if writeErr := binary.Write(conn, binary.BigEndian, uint32(0)); writeErr != nil { + serverErr <- writeErr + return + } + serverErr <- nil + }() + + client := &ClientService{ + ctx: ctx, + cancel: cancel, + logger: newTestLogger(), + dialer: testDialer{}, + serverAddr: M.SocksaddrFromNet(listener.Addr()), + ops: newTestUSBIPOps(t), + } + + entries, err := client.fetchDevList(ctx) + require.Nil(t, entries) + require.ErrorContains(t, err, "unexpected reply version") + require.NoError(t, <-serverErr) +} + +func TestClientFetchDevListReturnsOnContextCancelWhileServerStalls(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + defer listener.Close() + + requestReady := make(chan struct{}) + serverErr := make(chan error, 1) + go func() { + conn, acceptErr := listener.Accept() + if acceptErr != nil { + serverErr <- acceptErr + return + } + defer conn.Close() + + header, readErr := ReadOpHeader(conn) + if readErr != nil { + serverErr <- readErr + return + } + if header.Code != OpReqDevList { + serverErr <- fmt.Errorf("unexpected request code 0x%s", hex16(header.Code)) + return + } + close(requestReady) + + var buf [1]byte + _, readErr = conn.Read(buf[:]) + if readErr == nil { + serverErr <- errors.New("expected client close after cancellation") + return + } + serverErr <- nil + }() + + client := &ClientService{ + ctx: ctx, + cancel: cancel, + logger: newTestLogger(), + dialer: testDialer{}, + serverAddr: M.SocksaddrFromNet(listener.Addr()), + ops: newTestUSBIPOps(t), + } + + fetchErr := make(chan error, 1) + go func() { + _, fetchErrValue := client.fetchDevList(ctx) + fetchErr <- fetchErrValue + }() + + select { + case <-requestReady: + case <-time.After(3 * time.Second): + t.Fatal("fetchDevList did not reach stalled read path") + } + cancel() + + select { + case err = <-fetchErr: + require.Error(t, err) + case <-time.After(3 * time.Second): + t.Fatal("fetchDevList did not exit after cancellation") + } + require.NoError(t, <-serverErr) +} + +func TestClientAttemptAttachRejectsUnexpectedReplyVersion(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + defer listener.Close() + + device := newTestDevice("1-1", 0x1d6b, 0x0002, "serial-1", SpeedHigh) + info := device.toProtocol() + + serverErr := make(chan error, 1) + go func() { + conn, acceptErr := listener.Accept() + if acceptErr != nil { + serverErr <- acceptErr + return + } + defer conn.Close() + + header, readErr := ReadOpHeader(conn) + if readErr != nil { + serverErr <- readErr + return + } + if header.Code != OpReqImport { + serverErr <- fmt.Errorf("unexpected request code 0x%s", hex16(header.Code)) + return + } + busid, readErr := ReadOpReqImportBody(conn) + if readErr != nil { + serverErr <- readErr + return + } + if busid != "1-1" { + serverErr <- fmt.Errorf("unexpected busid %s", busid) + return + } + if writeErr := binary.Write(conn, binary.BigEndian, OpHeader{ + Version: ProtocolVersion + 1, + Code: OpRepImport, + Status: OpStatusOK, + }); writeErr != nil { + serverErr <- writeErr + return + } + if writeErr := binary.Write(conn, binary.BigEndian, &info); writeErr != nil { + serverErr <- writeErr + return + } + serverErr <- nil + }() + + ops := newTestUSBIPOps(t) + ops.vhciPickFreePort = func(uint32) (int, error) { + return -1, errors.New("unexpected vhci attach path") + } + + client := &ClientService{ + ctx: ctx, + cancel: cancel, + logger: newTestLogger(), + dialer: testDialer{}, + serverAddr: M.SocksaddrFromNet(listener.Addr()), + ops: ops, + } + + port, err := client.attemptAttach(ctx, "1-1") + require.Equal(t, -1, port) + require.ErrorContains(t, err, "unexpected reply version") + require.NoError(t, <-serverErr) +} + func TestClientRunControlSessionSyncsAssignmentsOnChanged(t *testing.T) { t.Parallel() diff --git a/service/usbip/server_linux.go b/service/usbip/server_linux.go index ca1d89b91..e990a3a3a 100644 --- a/service/usbip/server_linux.go +++ b/service/usbip/server_linux.go @@ -4,10 +4,13 @@ package usbip import ( "context" + "errors" "io" "net" "os" + "path/filepath" "slices" + "strings" "sync" "time" @@ -20,6 +23,7 @@ import ( "github.com/sagernet/sing/common" E "github.com/sagernet/sing/common/exceptions" N "github.com/sagernet/sing/common/network" + "golang.org/x/sys/unix" ) type serverExport struct { @@ -51,6 +55,8 @@ type ServerService struct { controlSeq uint64 controlNextID uint64 controlSubs map[uint64]*serverControlConn + + reconcileMu sync.Mutex } func NewServerService(ctx context.Context, logger log.ContextLogger, tag string, options option.USBIPServerServiceOptions) (adapter.Service, error) { @@ -89,7 +95,7 @@ func (s *ServerService) Start(stage adapter.StartStage) error { if err := s.ops.ensureHostDriver(); err != nil { return err } - if _, err := s.reconcileExports(); err != nil { + if err := s.reconcileAndBroadcast(false); err != nil { s.rollbackExports() return err } @@ -103,6 +109,7 @@ func (s *ServerService) Start(stage adapter.StartStage) error { s.mu.Unlock() go s.acceptLoop(tcpListener) go s.ueventLoop() + go s.reconcileLoop() return nil } @@ -131,6 +138,10 @@ func (s *ServerService) reconcileExports() (bool, error) { if !Matches(m, devices[i].key()) { continue } + if isVHCIImportedDevice(devices[i].Path) { + s.logger.Debug("skip vhci-imported device ", devices[i].BusID, " matched by ", describeMatch(m)) + continue + } if devices[i].DeviceClass == 0x09 { s.logger.Warn("skip hub device ", devices[i].BusID, " matched by ", describeMatch(m)) continue @@ -202,38 +213,36 @@ func (s *ServerService) bindOne(d *sysfsDevice) error { } func (s *ServerService) releaseExport(export serverExport, restore bool) error { - s.deleteExport(export.busid) - - var releaseErr error - if err := s.ops.writeUsbipSockfd(export.busid, -1); err != nil && !os.IsNotExist(err) { - releaseErr = err - } if !export.managed { + s.deleteExport(export.busid) s.logger.Info("stopped tracking ", export.busid, " on usbip-host") - return releaseErr + return nil } - if err := s.ops.hostUnbind(export.busid); err != nil && !os.IsNotExist(err) && releaseErr == nil { - releaseErr = err + if err := s.ops.writeUsbipSockfd(export.busid, -1); err != nil && !os.IsNotExist(err) { + return err } - if err := s.ops.hostMatchBusID(export.busid, false); err != nil && releaseErr == nil { - releaseErr = err + if err := s.ops.hostUnbind(export.busid); err != nil && !os.IsNotExist(err) && !(isMissingUSBDeviceError(err) && !restore) { + return err + } + if err := s.ops.hostMatchBusID(export.busid, false); err != nil { + return err } if !restore { + s.deleteExport(export.busid) s.logger.Info("removed export state for disappeared device ", export.busid) - return releaseErr + return nil } if export.originalDriver == "" { + s.deleteExport(export.busid) s.logger.Info("released ", export.busid, " from usbip-host") - return releaseErr + return nil } if err := s.ops.bindToDriver(export.busid, export.originalDriver); err != nil { - if releaseErr == nil { - releaseErr = err - } - return releaseErr + return err } + s.deleteExport(export.busid) s.logger.Info("restored ", export.busid, " to ", export.originalDriver) - return releaseErr + return nil } func (s *ServerService) rollbackExports() { @@ -247,6 +256,20 @@ func (s *ServerService) rollbackExports() { } } +func (s *ServerService) reconcileAndBroadcast(notify bool) error { + s.reconcileMu.Lock() + defer s.reconcileMu.Unlock() + + changed, err := s.reconcileExports() + if err != nil { + return err + } + if notify && changed { + s.broadcastChanged() + } + return nil +} + func (s *ServerService) currentExports() []string { s.mu.Lock() defer s.mu.Unlock() @@ -525,18 +548,30 @@ func (s *ServerService) ueventLoop() { } break } - changed, reconcileErr := s.reconcileExports() - if reconcileErr != nil { - s.logger.Warn("reconcile exports: ", reconcileErr) - continue - } - if changed { - s.broadcastChanged() + if err := s.reconcileAndBroadcast(true); err != nil { + s.logger.Warn("reconcile exports: ", err) } } } } +func (s *ServerService) reconcileLoop() { + ticker := time.NewTicker(time.Second) + defer ticker.Stop() + + for { + select { + case <-s.ctx.Done(): + return + case <-ticker.C: + } + + if err := s.reconcileAndBroadcast(true); err != nil { + s.logger.Warn("reconcile exports: ", err) + } + } +} + func (s *ServerService) registerControlConn(conn net.Conn) (*serverControlConn, uint64) { s.controlMu.Lock() defer s.controlMu.Unlock() @@ -602,6 +637,21 @@ func sysBusDevicePath(busid string) string { return sysBusUSBDevices + "/" + busid } +func isVHCIImportedDevice(path string) bool { + if strings.Contains(path, "vhci_hcd") { + return true + } + realPath, err := filepath.EvalSymlinks(path) + if err != nil { + return false + } + return strings.Contains(realPath, "vhci_hcd") +} + +func isMissingUSBDeviceError(err error) bool { + return errors.Is(err, unix.ENOENT) || errors.Is(err, unix.ENODEV) +} + func describeMatch(m option.USBIPDeviceMatch) string { var parts []string if m.BusID != "" {