Fix USB/IP matched retention and snapshot ordering
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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) {
|
||||
|
||||
Reference in New Issue
Block a user