From 26bd37a0128c2fdd1cdde1eca5657ada6e1aed3a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Fri, 24 Apr 2026 23:15:11 +0800 Subject: [PATCH] Fix USB/IP matched retention and snapshot ordering --- service/usbip/client_control.go | 54 ++++++++++++++++++++++++++ service/usbip/client_darwin.go | 19 +++++---- service/usbip/client_linux.go | 19 +++++---- service/usbip/client_standard_test.go | 55 +++++++++++++++++++++++++++ service/usbip/linux_test.go | 30 +++++++++++++++ service/usbip/server_darwin.go | 9 +++-- service/usbip/server_darwin_test.go | 28 ++++++++++++++ service/usbip/server_linux.go | 9 +++-- 8 files changed, 199 insertions(+), 24 deletions(-) diff --git a/service/usbip/client_control.go b/service/usbip/client_control.go index 777cf347d..df57b14c8 100644 --- a/service/usbip/client_control.go +++ b/service/usbip/client_control.go @@ -245,3 +245,57 @@ func (c *ClientService) resetControlDeviceStateFromEntries(entries []DeviceEntry c.remoteDevicesV2 = devices c.remoteMu.Unlock() } + +func (c *ClientService) matchedKeysForAssignmentLocked(entries []DeviceEntry, knownKeys map[string]DeviceKey) map[string]DeviceKey { + if len(c.matchedKnownKeys) == 0 && len(entries) == 0 && len(knownKeys) == 0 { + return nil + } + assignmentKeys := make(map[string]DeviceKey, len(c.matchedKnownKeys)+len(entries)+len(knownKeys)) + for busid, key := range c.matchedKnownKeys { + assignmentKeys[busid] = key + } + for i := range entries { + key := entryDeviceKey(entries[i]) + if key.BusID == "" { + continue + } + assignmentKeys[key.BusID] = key + } + for busid, key := range knownKeys { + if busid == "" { + continue + } + assignmentKeys[busid] = key + } + return assignmentKeys +} + +func (c *ClientService) retainMatchedKnownKeysLocked(assignmentKeys map[string]DeviceKey, entries []DeviceEntry, assigned []string) { + if len(assignmentKeys) == 0 { + c.matchedKnownKeys = nil + return + } + retained := make(map[string]DeviceKey, len(entries)+len(assigned)) + for i := range entries { + busid := entries[i].Info.BusIDString() + if busid == "" { + continue + } + if key, ok := assignmentKeys[busid]; ok { + retained[busid] = key + } + } + for _, busid := range assigned { + if busid == "" { + continue + } + if key, ok := assignmentKeys[busid]; ok { + retained[busid] = key + } + } + if len(retained) == 0 { + c.matchedKnownKeys = nil + return + } + c.matchedKnownKeys = retained +} diff --git a/service/usbip/client_darwin.go b/service/usbip/client_darwin.go index 4bf976b37..1dc8bb7c6 100644 --- a/service/usbip/client_darwin.go +++ b/service/usbip/client_darwin.go @@ -57,12 +57,13 @@ type ClientService struct { serverAddr M.Socksaddr matches []option.USBIPDeviceMatch - stateMu sync.Mutex - targets []clientTarget - assigned []string - assignedWorkers []*clientAssignedWorker - allWorkers map[string]*clientBusIDWorker - allDesired map[string]struct{} + stateMu sync.Mutex + targets []clientTarget + assigned []string + assignedWorkers []*clientAssignedWorker + allWorkers map[string]*clientBusIDWorker + allDesired map[string]struct{} + matchedKnownKeys map[string]DeviceKey wg sync.WaitGroup @@ -394,11 +395,13 @@ func (c *ClientService) applyMatchedExportsWithRetained(entries []DeviceEntry, k c.stateMu.Unlock() return } - activeCurrent := c.activeCurrentAssignmentsLocked(c.assigned, knownKeys) - nextAssigned := assignMatchedBusIDsWithRetained(c.targets, c.assigned, entries, knownKeys, activeCurrent) + assignmentKeys := c.matchedKeysForAssignmentLocked(entries, knownKeys) + activeCurrent := c.activeCurrentAssignmentsLocked(c.assigned, assignmentKeys) + nextAssigned := assignMatchedBusIDsWithRetained(c.targets, c.assigned, entries, assignmentKeys, activeCurrent) workers := append([]*clientAssignedWorker(nil), c.assignedWorkers...) previous := append([]string(nil), c.assigned...) c.assigned = nextAssigned + c.retainMatchedKnownKeysLocked(assignmentKeys, entries, nextAssigned) c.stateMu.Unlock() for i, worker := range workers { diff --git a/service/usbip/client_linux.go b/service/usbip/client_linux.go index 4256dcf3c..2853f2d66 100644 --- a/service/usbip/client_linux.go +++ b/service/usbip/client_linux.go @@ -68,12 +68,13 @@ type ClientService struct { matches []option.USBIPDeviceMatch // empty = import all remote exports ops usbipOps - stateMu sync.Mutex - targets []clientTarget - assigned []string - assignedWorkers []*clientAssignedWorker - allWorkers map[string]*clientBusIDWorker - allDesired map[string]struct{} + stateMu sync.Mutex + targets []clientTarget + assigned []string + assignedWorkers []*clientAssignedWorker + allWorkers map[string]*clientBusIDWorker + allDesired map[string]struct{} + matchedKnownKeys map[string]DeviceKey attachMu sync.Mutex // serializes vhci port pick + attach wg sync.WaitGroup @@ -425,11 +426,13 @@ func (c *ClientService) applyMatchedExportsWithRetained(entries []DeviceEntry, k c.stateMu.Unlock() return } - activeCurrent := c.activeCurrentAssignmentsLocked(c.assigned, knownKeys) - nextAssigned := assignMatchedBusIDsWithRetained(c.targets, c.assigned, entries, knownKeys, activeCurrent) + assignmentKeys := c.matchedKeysForAssignmentLocked(entries, knownKeys) + activeCurrent := c.activeCurrentAssignmentsLocked(c.assigned, assignmentKeys) + nextAssigned := assignMatchedBusIDsWithRetained(c.targets, c.assigned, entries, assignmentKeys, activeCurrent) workers := append([]*clientAssignedWorker(nil), c.assignedWorkers...) previous := append([]string(nil), c.assigned...) c.assigned = nextAssigned + c.retainMatchedKnownKeysLocked(assignmentKeys, entries, nextAssigned) c.stateMu.Unlock() for i, worker := range workers { diff --git a/service/usbip/client_standard_test.go b/service/usbip/client_standard_test.go index 477ed334c..7600f8afa 100644 --- a/service/usbip/client_standard_test.go +++ b/service/usbip/client_standard_test.go @@ -85,6 +85,61 @@ func TestClientStandardSessionPollsDevList(t *testing.T) { require.NoError(t, <-serverErr) } +func TestClientStandardMatchedAssignmentSurvivesHiddenActiveDevice(t *testing.T) { + t.Parallel() + + entry := standardTestDeviceEntry("1-1") + tests := []struct { + name string + match option.USBIPDeviceMatch + target clientTarget + }{ + { + name: "fixed busid", + match: option.USBIPDeviceMatch{BusID: "1-1"}, + target: clientTarget{fixedBusID: "1-1"}, + }, + { + name: "device key", + match: option.USBIPDeviceMatch{VendorID: 0x1d6b, ProductID: 0x0002}, + target: clientTarget{match: option.USBIPDeviceMatch{VendorID: 0x1d6b, ProductID: 0x0002}}, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + worker := &clientAssignedWorker{target: test.target, updates: make(chan string, 2)} + client := &ClientService{ + matches: []option.USBIPDeviceMatch{test.match}, + targets: []clientTarget{test.target}, + assigned: []string{"1-1"}, + assignedWorkers: []*clientAssignedWorker{worker}, + activeBusIDs: map[string]struct{}{"1-1": {}}, + } + + client.applyMatchedExports([]DeviceEntry{entry}) + require.Equal(t, []string{"1-1"}, client.assigned) + select { + case update := <-worker.updates: + t.Fatalf("unexpected assignment update %q", update) + default: + } + + client.applyMatchedExports(nil) + require.Equal(t, []string{"1-1"}, client.assigned) + select { + case update := <-worker.updates: + t.Fatalf("unexpected assignment update %q", update) + default: + } + + client.setBusIDActive("1-1", false) + client.applyMatchedExports(nil) + require.Equal(t, []string{""}, client.assigned) + require.Equal(t, "", <-worker.updates) + }) + } +} + func serveStandardDevLists(listener net.Listener, responses [][]DeviceEntry) error { for _, entries := range responses { conn, err := listener.Accept() diff --git a/service/usbip/linux_test.go b/service/usbip/linux_test.go index 520673df7..fa8a90852 100644 --- a/service/usbip/linux_test.go +++ b/service/usbip/linux_test.go @@ -1487,6 +1487,36 @@ func TestServerDispatchConnHandlesControlPingAndChanged(t *testing.T) { require.Equal(t, uint64(1), delta.Sequence) } +func TestServerRegisterControlConnQueuesSnapshotBeforeBroadcast(t *testing.T) { + t.Parallel() + + server := &ServerService{ + logger: newTestLogger(), + exports: make(map[string]serverExport), + controlSubs: make(map[uint64]*serverControlConn), + controlState: make(map[string]DeviceInfoV2), + } + serverConn, clientConn := net.Pipe() + defer serverConn.Close() + defer clientConn.Close() + + sub, seq := server.registerControlConn(serverConn, controlCapabilities) + require.Zero(t, seq) + require.Contains(t, server.controlSubs, sub.id) + + require.True(t, server.broadcastControlState(map[string]DeviceInfoV2{ + "1-1": {BusID: "1-1", State: deviceStateAvailable}, + }, true)) + + first := <-sub.send + require.Equal(t, controlFrameDeviceSnapshot, first.Frame.Type) + require.Zero(t, first.Frame.Sequence) + + second := <-sub.send + require.Equal(t, controlFrameDeviceDelta, second.Frame.Type) + require.Equal(t, uint64(1), second.Frame.Sequence) +} + func TestServerReconcileBroadcastsStatusOnlyDeviceDelta(t *testing.T) { t.Parallel() diff --git a/service/usbip/server_darwin.go b/service/usbip/server_darwin.go index 0e2e02015..5a0f961da 100644 --- a/service/usbip/server_darwin.go +++ b/service/usbip/server_darwin.go @@ -361,9 +361,6 @@ func (s *ServerService) handleControlConn(conn net.Conn) { s.logger.Debug("write control ack: ", err) return } - if supportsControlExtensions(capabilities) { - s.enqueueControlSnapshot(sub, seq) - } readDone := make(chan struct{}) go s.readControlConn(sub, readDone) for { @@ -491,14 +488,18 @@ func (s *ServerService) registerControlConn(conn net.Conn, capabilities uint32) s.controlMu.Lock() defer s.controlMu.Unlock() s.controlNextID++ + sequence := s.controlSeq sub := &serverControlConn{ id: s.controlNextID, capabilities: capabilities, conn: conn, send: make(chan controlOutboundMessage, 16), } + if supportsControlExtensions(capabilities) { + s.enqueueControlSnapshot(sub, sequence) + } s.controlSubs[sub.id] = sub - return sub, s.controlSeq + return sub, sequence } func (s *ServerService) unregisterControlConn(id uint64) { diff --git a/service/usbip/server_darwin_test.go b/service/usbip/server_darwin_test.go index ca67354a4..4b2a6811d 100644 --- a/service/usbip/server_darwin_test.go +++ b/service/usbip/server_darwin_test.go @@ -167,6 +167,34 @@ func TestDarwinServerBuildDeviceStateIncludesBusyExports(t *testing.T) { require.Equal(t, deviceStateBusy, devices["busy"].State) } +func TestDarwinServerRegisterControlConnQueuesSnapshotBeforeBroadcast(t *testing.T) { + t.Parallel() + + server := &ServerService{ + logger: newTestLogger(), + exports: make(map[string]serverExport), + controlSubs: make(map[uint64]*serverControlConn), + controlState: make(map[string]DeviceInfoV2), + } + serverConn, clientConn := net.Pipe() + defer serverConn.Close() + defer clientConn.Close() + + sub, seq := server.registerControlConn(serverConn, controlCapabilities) + require.Zero(t, seq) + require.Contains(t, server.controlSubs, sub.id) + + server.broadcastChanged() + + first := <-sub.send + require.Equal(t, controlFrameDeviceSnapshot, first.Frame.Type) + require.Zero(t, first.Frame.Sequence) + + second := <-sub.send + require.Equal(t, controlFrameDeviceDelta, second.Frame.Type) + require.Equal(t, uint64(1), second.Frame.Sequence) +} + func TestDarwinServerImportBroadcastsBusyState(t *testing.T) { t.Parallel() diff --git a/service/usbip/server_linux.go b/service/usbip/server_linux.go index c6346a615..6134a5ed4 100644 --- a/service/usbip/server_linux.go +++ b/service/usbip/server_linux.go @@ -471,9 +471,6 @@ func (s *ServerService) handleControlConn(conn net.Conn) { s.logger.Debug("write control ack: ", err) return } - if supportsControlExtensions(capabilities) { - s.enqueueControlSnapshot(sub, seq) - } readDone := make(chan struct{}) go s.readControlConn(sub, readDone) @@ -691,14 +688,18 @@ func (s *ServerService) registerControlConn(conn net.Conn, capabilities uint32) s.controlMu.Lock() defer s.controlMu.Unlock() s.controlNextID++ + sequence := s.controlSeq sub := &serverControlConn{ id: s.controlNextID, capabilities: capabilities, conn: conn, send: make(chan controlOutboundMessage, 16), } + if supportsControlExtensions(capabilities) { + s.enqueueControlSnapshot(sub, sequence) + } s.controlSubs[sub.id] = sub - return sub, s.controlSeq + return sub, sequence } func (s *ServerService) unregisterControlConn(id uint64) {