diff --git a/service/usbip/handoff_linux.go b/service/usbip/handoff_linux.go index c15ebb62a..5fd574e3f 100644 --- a/service/usbip/handoff_linux.go +++ b/service/usbip/handoff_linux.go @@ -86,10 +86,9 @@ func (h *usbipConnHandoff) Close() error { func (h *usbipConnHandoff) startRelay(ctx context.Context, logger log.ContextLogger, side string, busid string) bool { if !h.relay() { - go func() { - <-ctx.Done() - _ = h.conn.Close() - }() + if err := h.conn.Close(); err != nil && !E.IsClosedOrCanceled(err) { + logger.Debug("close usbip ", side, " userspace socket ", busid, ": ", err) + } return true } relayConn := h.relayConn diff --git a/service/usbip/linux_test.go b/service/usbip/linux_test.go index 7ba1aaebf..3e1296cb9 100644 --- a/service/usbip/linux_test.go +++ b/service/usbip/linux_test.go @@ -708,6 +708,10 @@ func TestUSBIPConnHandoffDirectTCP(t *testing.T) { require.False(t, handoff.relay()) require.Equal(t, "direct", handoff.mode()) requireStreamSocketFD(t, handoff.kernelFD()) + require.True(t, handoff.startRelay(context.Background(), newTestLogger(), "test", "direct")) + + _, err = conn.Write([]byte("closed")) + require.Error(t, err) } func TestUSBIPConnHandoffRelaySocketpairCopies(t *testing.T) { @@ -994,6 +998,122 @@ func TestServerReleaseExportRetainsTrackingOnFailure(t *testing.T) { require.Equal(t, map[string]serverExport{"1-1": export}, server.snapshotExports()) } +func TestServerCloseSerializesRollbackWithActiveReconcile(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(context.Background()) + device := newTestDevice("1-1", 0x1d6b, 0x0002, "regular", SpeedHigh) + listEntered := make(chan struct{}) + releaseList := make(chan struct{}) + reconcileDone := make(chan error, 1) + closeDone := make(chan error, 1) + + var actionsMu sync.Mutex + var actions []string + record := func(action string) { + actionsMu.Lock() + defer actionsMu.Unlock() + actions = append(actions, action) + } + + ops := newTestUSBIPOps(t) + ops.listUSBDevices = func() ([]sysfsDevice, error) { + close(listEntered) + <-releaseList + return []sysfsDevice{device}, nil + } + ops.currentDriver = func(string) (string, error) { + return "", nil + } + ops.hostMatchBusID = func(busid string, add bool) error { + if add { + record("match add " + busid) + } else { + record("match del " + busid) + } + return nil + } + ops.hostBind = func(busid string) error { + record("hostbind " + busid) + return nil + } + ops.readSysfsDevice = func(string, string) (sysfsDevice, error) { + return device, nil + } + ops.readUsbipStatus = func(string) (int, error) { + return usbipStatusAvailable, nil + } + ops.hostUnbind = func(busid string) error { + record("hostunbind " + busid) + return nil + } + + server := &ServerService{ + ctx: ctx, + cancel: cancel, + logger: newTestLogger(), + matches: []option.USBIPDeviceMatch{{BusID: "1-1"}}, + exports: make(map[string]serverExport), + controlSubs: make(map[uint64]*serverControlConn), + controlState: make(map[string]DeviceInfoV2), + ops: ops, + } + + go func() { + reconcileDone <- server.reconcileAndBroadcast(true) + }() + select { + case <-listEntered: + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for active reconcile") + } + go func() { + closeDone <- server.Close() + }() + close(releaseList) + + select { + case err := <-reconcileDone: + require.NoError(t, err) + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for reconcile") + } + select { + case err := <-closeDone: + require.NoError(t, err) + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for close") + } + + actionsMu.Lock() + defer actionsMu.Unlock() + require.Equal(t, []string{ + "match add 1-1", + "hostbind 1-1", + "hostunbind 1-1", + "match del 1-1", + }, actions) + require.Empty(t, server.snapshotExports()) +} + +func TestServerReconcileAndBroadcastSkipsAfterCancel(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + ops := newTestUSBIPOps(t) + server := &ServerService{ + ctx: ctx, + logger: newTestLogger(), + exports: make(map[string]serverExport), + controlSubs: make(map[uint64]*serverControlConn), + controlState: make(map[string]DeviceInfoV2), + ops: ops, + } + + require.NoError(t, server.reconcileAndBroadcast(true)) +} + func TestServerBuildDevListEntriesFiltersUnavailableAndRefreshFailures(t *testing.T) { t.Parallel() diff --git a/service/usbip/server_darwin.go b/service/usbip/server_darwin.go index b7d09aab7..a070f6dd3 100644 --- a/service/usbip/server_darwin.go +++ b/service/usbip/server_darwin.go @@ -119,6 +119,8 @@ func (s *ServerService) Close() error { } s.closeControlSubscribers() err := common.Close(common.PtrOrNil(s.listener)) + s.reconcileMu.Lock() + defer s.reconcileMu.Unlock() s.rollbackExports() return err } @@ -189,6 +191,9 @@ func (s *ServerService) reconcileExports() (bool, error) { func (s *ServerService) reconcileAndBroadcast(notify bool) error { s.reconcileMu.Lock() defer s.reconcileMu.Unlock() + if s.ctx != nil && s.ctx.Err() != nil { + return nil + } changed, err := s.reconcileExports() if err != nil { return err diff --git a/service/usbip/server_darwin_test.go b/service/usbip/server_darwin_test.go index 77204a643..a03c96236 100644 --- a/service/usbip/server_darwin_test.go +++ b/service/usbip/server_darwin_test.go @@ -3,6 +3,7 @@ package usbip import ( + "context" "testing" "github.com/stretchr/testify/require" @@ -29,3 +30,19 @@ func TestDarwinServerPendingSubmitUnlinkState(t *testing.T) { _, active = session.markSubmitUnlinked(8) require.False(t, active) } + +func TestDarwinServerReconcileAndBroadcastSkipsAfterCancel(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + server := &ServerService{ + ctx: ctx, + logger: newTestLogger(), + exports: make(map[string]serverExport), + controlSubs: make(map[uint64]*serverControlConn), + controlState: make(map[string]DeviceInfoV2), + } + + require.NoError(t, server.reconcileAndBroadcast(true)) +} diff --git a/service/usbip/server_linux.go b/service/usbip/server_linux.go index 1badf1731..e5bffd496 100644 --- a/service/usbip/server_linux.go +++ b/service/usbip/server_linux.go @@ -133,6 +133,8 @@ func (s *ServerService) Close() error { } s.closeControlSubscribers() err := common.Close(common.PtrOrNil(s.listener)) + s.reconcileMu.Lock() + defer s.reconcileMu.Unlock() s.rollbackExports() return err } @@ -306,6 +308,9 @@ func (s *ServerService) reconcileAndBroadcast(notify bool) error { s.reconcileMu.Lock() defer s.reconcileMu.Unlock() + if s.ctx != nil && s.ctx.Err() != nil { + return nil + } if _, err := s.reconcileExports(); err != nil { return err }