diff --git a/service/usbip/client_control.go b/service/usbip/client_control.go index 16c2135d4..777cf347d 100644 --- a/service/usbip/client_control.go +++ b/service/usbip/client_control.go @@ -167,6 +167,21 @@ func (c *ClientService) requestImportLease(ctx context.Context, busid string) (c }, nil } +func (c *ClientService) runStandardSession() error { + return c.runStandardSessionWithInterval(clientReconnectDelay) +} + +func (c *ClientService) runStandardSessionWithInterval(interval time.Duration) error { + for { + if err := c.syncRemoteState(); err != nil { + return E.Cause(err, "devlist sync") + } + if !sleepCtx(c.ctx, interval) { + return nil + } + } +} + func (c *ClientService) applyControlSnapshot(snapshot controlDeviceSnapshot) { devices := deviceInfoV2Map(snapshot.Devices) values := sortedDeviceInfoV2Values(devices) diff --git a/service/usbip/client_darwin.go b/service/usbip/client_darwin.go index 3590a40e9..84e6292c5 100644 --- a/service/usbip/client_darwin.go +++ b/service/usbip/client_darwin.go @@ -290,14 +290,6 @@ func (c *ClientService) runControlSession() error { } } -func (c *ClientService) runStandardSession() error { - if err := c.syncRemoteState(); err != nil { - return E.Cause(err, "initial devlist sync") - } - <-c.ctx.Done() - return nil -} - func (c *ClientService) controlPingLoop(session *clientControlSession, done <-chan struct{}) { ticker := time.NewTicker(controlPingInterval) defer ticker.Stop() diff --git a/service/usbip/client_linux.go b/service/usbip/client_linux.go index 0b3a02025..4256dcf3c 100644 --- a/service/usbip/client_linux.go +++ b/service/usbip/client_linux.go @@ -316,14 +316,6 @@ func (c *ClientService) runControlSession() error { } } -func (c *ClientService) runStandardSession() error { - if err := c.syncRemoteState(); err != nil { - return E.Cause(err, "initial devlist sync") - } - <-c.ctx.Done() - return nil -} - func (c *ClientService) controlPingLoop(session *clientControlSession, done <-chan struct{}) { ticker := time.NewTicker(controlPingInterval) defer ticker.Stop() diff --git a/service/usbip/client_standard_test.go b/service/usbip/client_standard_test.go new file mode 100644 index 000000000..477ed334c --- /dev/null +++ b/service/usbip/client_standard_test.go @@ -0,0 +1,135 @@ +//go:build linux || (darwin && cgo) + +package usbip + +import ( + "context" + "errors" + "fmt" + "net" + "testing" + "time" + + "github.com/sagernet/sing-box/log" + "github.com/sagernet/sing-box/option" + M "github.com/sagernet/sing/common/metadata" + + "github.com/stretchr/testify/require" +) + +type standardTestDialer struct{} + +func (standardTestDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) { + var dialer net.Dialer + return dialer.DialContext(ctx, network, destination.String()) +} + +func (standardTestDialer) ListenPacket(context.Context, M.Socksaddr) (net.PacketConn, error) { + return nil, errors.New("unused") +} + +func TestClientStandardSessionPollsDevList(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) + t.Cleanup(func() { _ = listener.Close() }) + + entry := standardTestDeviceEntry("1-1") + serverErr := make(chan error, 1) + go func() { + serverErr <- serveStandardDevLists(listener, [][]DeviceEntry{ + nil, + {entry}, + }) + }() + + target := clientTarget{fixedBusID: "1-1"} + worker := &clientAssignedWorker{target: target, updates: make(chan string, 1)} + client := &ClientService{ + ctx: ctx, + logger: log.NewNOPFactory().NewLogger("usbip"), + dialer: standardTestDialer{}, + serverAddr: standardTestSocksaddr(listener.Addr()), + matches: []option.USBIPDeviceMatch{{BusID: "1-1"}}, + targets: []clientTarget{target}, + assigned: []string{""}, + assignedWorkers: []*clientAssignedWorker{worker}, + allWorkers: make(map[string]*clientBusIDWorker), + allDesired: make(map[string]struct{}), + activeBusIDs: make(map[string]struct{}), + } + + sessionErr := make(chan error, 1) + go func() { + sessionErr <- client.runStandardSessionWithInterval(10 * time.Millisecond) + }() + + select { + case update := <-worker.updates: + require.Equal(t, "1-1", update) + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for standard devlist refresh") + } + cancel() + + select { + case err := <-sessionErr: + require.NoError(t, err) + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for standard session shutdown") + } + require.NoError(t, <-serverErr) +} + +func serveStandardDevLists(listener net.Listener, responses [][]DeviceEntry) error { + for _, entries := range responses { + conn, err := listener.Accept() + if err != nil { + return err + } + if err := handleStandardDevListConn(conn, entries); err != nil { + return err + } + } + return nil +} + +func handleStandardDevListConn(conn net.Conn, entries []DeviceEntry) error { + defer conn.Close() + header, err := ReadOpHeader(conn) + if err != nil { + return err + } + if header.Version != ProtocolVersion || header.Code != OpReqDevList || header.Status != OpStatusOK { + return fmt.Errorf("unexpected devlist request: version=0x%s code=0x%s status=%d", hex16(header.Version), hex16(header.Code), header.Status) + } + return WriteOpRepDevList(conn, entries) +} + +func standardTestDeviceEntry(busid string) DeviceEntry { + var info DeviceInfoTruncated + copy(info.BusID[:], busid) + info.BusNum = 1 + info.DevNum = 1 + info.Speed = SpeedHigh + info.IDVendor = 0x1d6b + info.IDProduct = 0x0002 + info.BConfigurationValue = 1 + info.BNumConfigurations = 1 + info.BNumInterfaces = 1 + return DeviceEntry{ + Info: info, + Interfaces: []DeviceInterface{{ + BInterfaceClass: 0xff, + }}, + } +} + +func standardTestSocksaddr(address net.Addr) M.Socksaddr { + tcpAddr := address.(*net.TCPAddr) + return M.ParseSocksaddrHostPort(tcpAddr.IP.String(), uint16(tcpAddr.Port)) +} diff --git a/service/usbip/server_darwin.go b/service/usbip/server_darwin.go index a070f6dd3..4dcce5168 100644 --- a/service/usbip/server_darwin.go +++ b/service/usbip/server_darwin.go @@ -230,6 +230,19 @@ func (s *ServerService) currentExports() []serverExport { return out } +func (s *ServerService) allExports() []serverExport { + s.mu.Lock() + defer s.mu.Unlock() + out := make([]serverExport, 0, len(s.exports)) + for _, export := range s.exports { + out = append(out, export) + } + slices.SortFunc(out, func(a, b serverExport) int { + return stringsCompare(a.busid, b.busid) + }) + return out +} + func (s *ServerService) snapshotExports() map[string]serverExport { s.mu.Lock() defer s.mu.Unlock() @@ -436,10 +449,12 @@ func (s *ServerService) handleImportBusID(conn net.Conn, busid string, extended _ = writeReply(conn, OpStatusError, nil) return } + s.broadcastChanged() releaseClaim := true defer func() { if releaseClaim { s.releaseClaim(busid) + s.broadcastChanged() } }() info := export.entry.Info @@ -577,13 +592,19 @@ func (s *ServerService) refreshControlState() { } func (s *ServerService) buildDeviceStateV2() []DeviceInfoV2 { - exports := s.currentExports() + exports := s.allExports() if len(exports) == 0 { return nil } devices := make([]DeviceInfoV2, 0, len(exports)) for _, export := range exports { - devices = append(devices, deviceInfoV2FromEntry(export.entry, "darwin-iokit", darwinStableID(export.registryID), deviceStateAvailable, 0, "available")) + state := deviceStateAvailable + reason := "available" + if export.busy { + state = deviceStateBusy + reason = "busy" + } + devices = append(devices, deviceInfoV2FromEntry(export.entry, "darwin-iokit", darwinStableID(export.registryID), state, 0, reason)) } return devices } diff --git a/service/usbip/server_darwin_test.go b/service/usbip/server_darwin_test.go index a03c96236..d0cf04d1d 100644 --- a/service/usbip/server_darwin_test.go +++ b/service/usbip/server_darwin_test.go @@ -4,8 +4,11 @@ package usbip import ( "context" + "net" "testing" + "time" + "github.com/sagernet/sing-box/log" "github.com/stretchr/testify/require" ) @@ -46,3 +49,81 @@ func TestDarwinServerReconcileAndBroadcastSkipsAfterCancel(t *testing.T) { require.NoError(t, server.reconcileAndBroadcast(true)) } + +func TestDarwinServerBuildDeviceStateIncludesBusyExports(t *testing.T) { + t.Parallel() + + available := standardTestDeviceEntry("available") + busy := standardTestDeviceEntry("busy") + server := &ServerService{ + exports: map[string]serverExport{ + "available": { + busid: "available", + registryID: 1, + entry: available, + }, + "busy": { + busid: "busy", + registryID: 2, + entry: busy, + busy: true, + }, + }, + } + + devices := deviceInfoV2Map(server.buildDeviceStateV2()) + require.Equal(t, deviceStateAvailable, devices["available"].State) + require.Equal(t, deviceStateBusy, devices["busy"].State) +} + +func TestDarwinServerImportBroadcastsBusyState(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + const busid = "1-1" + entry := standardTestDeviceEntry(busid) + server := &ServerService{ + ctx: ctx, + logger: log.NewNOPFactory().NewLogger("usbip"), + exports: map[string]serverExport{busid: {busid: busid, registryID: 1, entry: entry}}, + controlSubs: make(map[uint64]*serverControlConn), + controlState: make(map[string]DeviceInfoV2), + } + + serverConn, clientConn := net.Pipe() + done := make(chan struct{}) + go func() { + defer close(done) + server.handleImportBusID(serverConn, busid, false) + }() + + header, err := ReadOpHeader(clientConn) + require.NoError(t, err) + require.Equal(t, OpRepImport, header.Code) + require.Equal(t, OpStatusOK, header.Status) + _, err = ReadOpRepImportBody(clientConn) + require.NoError(t, err) + + require.Eventually(t, func() bool { + return darwinServerControlState(server, busid) == deviceStateBusy + }, time.Second, 10*time.Millisecond) + + require.NoError(t, clientConn.Close()) + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timed out waiting for Darwin import session shutdown") + } + + require.Eventually(t, func() bool { + return darwinServerControlState(server, busid) == deviceStateAvailable + }, time.Second, 10*time.Millisecond) +} + +func darwinServerControlState(server *ServerService, busid string) string { + server.controlMu.Lock() + defer server.controlMu.Unlock() + return server.controlState[busid].State +}