From 1cf178440f1213747819230873d494dcb961b2e3 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Fri, 24 Apr 2026 04:47:10 +0800 Subject: [PATCH] usbip: add linux socketpair relay fallback --- service/usbip/client_linux.go | 34 +- service/usbip/handoff_linux.go | 104 ++++++ service/usbip/linux_test.go | 585 +++++++++++++++++++++++++++++++++ service/usbip/server_linux.go | 45 +-- 4 files changed, 737 insertions(+), 31 deletions(-) create mode 100644 service/usbip/handoff_linux.go diff --git a/service/usbip/client_linux.go b/service/usbip/client_linux.go index 6398aa004..57e192c6d 100644 --- a/service/usbip/client_linux.go +++ b/service/usbip/client_linux.go @@ -29,8 +29,10 @@ const ( controlSessionIdleHint = "control session lost" ) -var errImmediateReconnect = errors.New("usbip control reconnect") -var errControlUnsupported = errors.New("usbip control unsupported") +var ( + errImmediateReconnect = errors.New("usbip control reconnect") + errControlUnsupported = errors.New("usbip control unsupported") +) type clientTarget struct { fixedBusID string @@ -542,7 +544,12 @@ func (c *ClientService) attemptAttach(ctx context.Context, busid string) (int, e if err != nil { return -1, E.Cause(err, "dial ", c.serverAddr) } - defer conn.Close() + relayStarted := false + defer func() { + if !relayStarted { + _ = conn.Close() + } + }() stopCloseOnCancel := closeConnOnContextDone(ctx, conn) defer stopCloseOnCancel() if err := WriteOpReqImport(conn, busid); err != nil { @@ -565,15 +572,16 @@ func (c *ClientService) attemptAttach(ctx context.Context, busid string) (int, e if err != nil { return -1, E.Cause(err, "read OP_REP_IMPORT body") } - tcp, ok := conn.(*net.TCPConn) - if !ok { - return -1, E.New("dialed conn is not *net.TCPConn (type=", conn, ")") - } - file, err := tcp.File() + handoff, err := newUSBIPConnHandoff(conn) if err != nil { - return -1, E.Cause(err, "dup socket fd") + return -1, E.Cause(err, "prepare handoff") } - defer file.Close() + defer func() { + if !relayStarted { + _ = handoff.Close() + } + }() + c.logger.Debug("usbip client handoff ", busid, ": ", handoff.mode()) c.attachMu.Lock() defer c.attachMu.Unlock() port, err := c.ops.vhciPickFreePort(info.Speed) @@ -583,10 +591,14 @@ func (c *ClientService) attemptAttach(ctx context.Context, busid string) (int, e 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 { + if err := c.ops.vhciAttach(port, handoff.kernelFD(), info.DevID(), info.Speed); err != nil { c.trackPort(port, false) return -1, E.Cause(err, "vhci attach") } + if err := handoff.closeKernelFD(); err != nil { + c.logger.Debug("close kernel fd ", busid, ": ", err) + } + relayStarted = handoff.startRelay(ctx, c.logger, "client", busid) return port, nil } diff --git a/service/usbip/handoff_linux.go b/service/usbip/handoff_linux.go new file mode 100644 index 000000000..775eb1efa --- /dev/null +++ b/service/usbip/handoff_linux.go @@ -0,0 +1,104 @@ +//go:build linux + +package usbip + +import ( + "context" + "net" + "os" + + "github.com/sagernet/sing-box/log" + "github.com/sagernet/sing/common" + sBufio "github.com/sagernet/sing/common/bufio" + E "github.com/sagernet/sing/common/exceptions" + N "github.com/sagernet/sing/common/network" + + "golang.org/x/sys/unix" +) + +type usbipConnHandoff struct { + conn net.Conn + file *os.File + relayConn net.Conn +} + +func newUSBIPConnHandoff(conn net.Conn) (*usbipConnHandoff, error) { + if tcpConn, _ := N.UnwrapReader(conn).(*net.TCPConn); tcpConn != nil { + file, err := tcpConn.File() + if err != nil { + return nil, E.Cause(err, "dup TCP socket fd") + } + return &usbipConnHandoff{ + conn: conn, + file: file, + }, nil + } + + fds, err := unix.Socketpair(unix.AF_UNIX, unix.SOCK_STREAM|unix.SOCK_CLOEXEC, 0) + if err != nil { + return nil, E.Cause(err, "create USB/IP relay socketpair") + } + kernelFile := os.NewFile(uintptr(fds[0]), "usbip-kernel") + relayFile := os.NewFile(uintptr(fds[1]), "usbip-relay") + relayConn, err := net.FileConn(relayFile) + _ = relayFile.Close() + if err != nil { + _ = kernelFile.Close() + return nil, E.Cause(err, "wrap USB/IP relay socket") + } + return &usbipConnHandoff{ + conn: conn, + file: kernelFile, + relayConn: relayConn, + }, nil +} + +func (h *usbipConnHandoff) kernelFD() uintptr { + return h.file.Fd() +} + +func (h *usbipConnHandoff) relay() bool { + return h.relayConn != nil +} + +func (h *usbipConnHandoff) mode() string { + if h.relay() { + return "relay" + } + return "direct" +} + +func (h *usbipConnHandoff) closeKernelFD() error { + if h.file == nil { + return nil + } + err := h.file.Close() + h.file = nil + return err +} + +func (h *usbipConnHandoff) Close() error { + return E.Errors( + h.closeKernelFD(), + common.Close(h.relayConn), + ) +} + +func (h *usbipConnHandoff) startRelay(ctx context.Context, logger log.ContextLogger, side string, busid string) bool { + if !h.relay() { + return false + } + relayConn := h.relayConn + h.relayConn = nil + go func() { + err := sBufio.CopyConn(ctx, h.conn, relayConn) + if err == nil { + logger.Debug("usbip ", side, " relay ", busid, " closed") + } else if ctx.Err() == nil && !E.IsClosedOrCanceled(err) { + logger.Warn("usbip ", side, " relay ", busid, ": ", err) + } else { + logger.Debug("usbip ", side, " relay ", busid, ": ", err) + } + }() + return true +} diff --git a/service/usbip/linux_test.go b/service/usbip/linux_test.go index 39935ea94..967a3f17d 100644 --- a/service/usbip/linux_test.go +++ b/service/usbip/linux_test.go @@ -7,6 +7,7 @@ import ( "encoding/binary" "errors" "fmt" + "io" "net" "os" "os/exec" @@ -22,6 +23,7 @@ import ( M "github.com/sagernet/sing/common/metadata" "github.com/stretchr/testify/require" + "golang.org/x/sys/unix" ) type testDialer struct{} @@ -47,6 +49,25 @@ func (d failingDialer) ListenPacket(context.Context, M.Socksaddr) (net.PacketCon return nil, errors.New("unused") } +type opaqueConn struct { + net.Conn +} + +type wrappingDialer struct{} + +func (wrappingDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) { + var dialer net.Dialer + conn, err := dialer.DialContext(ctx, network, destination.String()) + if err != nil { + return nil, err + } + return opaqueConn{Conn: conn}, nil +} + +func (wrappingDialer) ListenPacket(context.Context, M.Socksaddr) (net.PacketConn, error) { + return nil, errors.New("unused") +} + type testDeviceStore struct { mu sync.Mutex devices map[string]sysfsDevice @@ -274,6 +295,71 @@ func startDispatchServer(t *testing.T, server *ServerService) (M.Socksaddr, func } } +func duplicateConnFromFD(t *testing.T, fd uintptr, name string) net.Conn { + t.Helper() + + conn, err := duplicateNetConnFromFD(fd, name) + require.NoError(t, err) + return conn +} + +func duplicateNetConnFromFD(fd uintptr, name string) (net.Conn, error) { + dupFD, err := unix.Dup(int(fd)) + if err != nil { + return nil, err + } + file := os.NewFile(uintptr(dupFD), name) + conn, err := net.FileConn(file) + closeErr := file.Close() + if err != nil { + return nil, err + } + if closeErr != nil { + return nil, closeErr + } + return conn, nil +} + +func duplicateHandoffKernelConn(t *testing.T, handoff *usbipConnHandoff) net.Conn { + t.Helper() + + conn := duplicateConnFromFD(t, handoff.kernelFD(), "usbip-test-kernel") + require.NoError(t, handoff.closeKernelFD()) + return conn +} + +func requireConnRead(t *testing.T, conn net.Conn, expected []byte) { + t.Helper() + + buffer := make([]byte, len(expected)) + _, err := io.ReadFull(conn, buffer) + require.NoError(t, err) + require.Equal(t, expected, buffer) +} + +func requireConnEOF(t *testing.T, conn net.Conn) { + t.Helper() + + buffer := make([]byte, 1) + n, err := conn.Read(buffer) + require.Zero(t, n) + require.ErrorIs(t, err, io.EOF) +} + +func setConnDeadline(t *testing.T, conn net.Conn) { + t.Helper() + + require.NoError(t, conn.SetDeadline(time.Now().Add(3*time.Second))) +} + +func requireStreamSocketFD(t *testing.T, fd uintptr) { + t.Helper() + + socketType, err := unix.GetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_TYPE) + require.NoError(t, err) + require.Equal(t, unix.SOCK_STREAM, socketType) +} + type testUSBGadget struct { path string serial string @@ -529,6 +615,64 @@ func TestLinuxHelpers(t *testing.T) { require.False(t, isUSBUEvent([]byte("ACTION=add\x00SUBSYSTEM=net\x00"))) } +func TestUSBIPConnHandoffDirectTCP(t *testing.T) { + t.Parallel() + + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + defer listener.Close() + + accepted := make(chan net.Conn, 1) + go func() { + conn, _ := listener.Accept() + accepted <- conn + }() + + conn, err := net.Dial("tcp", listener.Addr().String()) + require.NoError(t, err) + defer conn.Close() + acceptedConn := <-accepted + defer acceptedConn.Close() + + handoff, err := newUSBIPConnHandoff(conn) + require.NoError(t, err) + defer handoff.Close() + + require.False(t, handoff.relay()) + require.Equal(t, "direct", handoff.mode()) + requireStreamSocketFD(t, handoff.kernelFD()) +} + +func TestUSBIPConnHandoffRelaySocketpairCopies(t *testing.T) { + t.Parallel() + + left, right := net.Pipe() + defer right.Close() + handoff, err := newUSBIPConnHandoff(opaqueConn{Conn: left}) + require.NoError(t, err) + defer handoff.Close() + require.True(t, handoff.relay()) + require.Equal(t, "relay", handoff.mode()) + requireStreamSocketFD(t, handoff.kernelFD()) + + kernelConn := duplicateHandoffKernelConn(t, handoff) + defer kernelConn.Close() + setConnDeadline(t, right) + setConnDeadline(t, kernelConn) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + require.True(t, handoff.startRelay(ctx, newTestLogger(), "test", "relay")) + + _, err = right.Write([]byte("ping")) + require.NoError(t, err) + requireConnRead(t, kernelConn, []byte("ping")) + + _, err = kernelConn.Write([]byte("pong")) + require.NoError(t, err) + requireConnRead(t, right, []byte("pong")) +} + func TestServerStartRequiresHostDriver(t *testing.T) { t.Parallel() @@ -806,6 +950,229 @@ func TestServerBuildDevListEntriesFiltersUnavailableAndRefreshFailures(t *testin require.Equal(t, "ok", entries[0].Info.SerialString()) } +func TestServerHandleImportWithOpaqueConnRelay(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + device := newTestDevice("1-1", 0x1d6b, 0x0002, "serial-1", SpeedHigh) + store := newTestDeviceStore(device) + store.setStatus("1-1", usbipStatusAvailable) + + kernelConnCh := make(chan net.Conn, 1) + kernelErrCh := make(chan error, 1) + ops := newTestUSBIPOps(t) + ops.readUsbipStatus = store.readUsbipStatus + ops.readSysfsDevice = store.readSysfsDevice + ops.writeUsbipSockfd = func(busid string, fd int) error { + if fd < 0 { + return nil + } + if busid != "1-1" { + kernelErrCh <- fmt.Errorf("unexpected busid %s", busid) + return nil + } + socketType, err := unix.GetsockoptInt(fd, unix.SOL_SOCKET, unix.SO_TYPE) + if err != nil { + kernelErrCh <- err + return nil + } + if socketType != unix.SOCK_STREAM { + kernelErrCh <- fmt.Errorf("unexpected socket type %d", socketType) + return nil + } + kernelConn, err := duplicateNetConnFromFD(uintptr(fd), "usbip-server-test-kernel") + if err != nil { + kernelErrCh <- err + return nil + } + kernelConnCh <- kernelConn + return nil + } + + server := &ServerService{ + ctx: ctx, + cancel: cancel, + logger: newTestLogger(), + exports: map[string]serverExport{"1-1": {busid: "1-1"}}, + controlSubs: make(map[uint64]*serverControlConn), + ops: ops, + } + + serverConn, clientConn := net.Pipe() + defer clientConn.Close() + go server.dispatchConn(opaqueConn{Conn: serverConn}) + + setConnDeadline(t, clientConn) + require.NoError(t, WriteOpReqImport(clientConn, "1-1")) + 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) + + var kernelConn net.Conn + select { + case kernelConn = <-kernelConnCh: + case err = <-kernelErrCh: + require.NoError(t, err) + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for server relay kernel conn") + } + defer kernelConn.Close() + setConnDeadline(t, kernelConn) + + _, err = clientConn.Write([]byte("server-in")) + require.NoError(t, err) + requireConnRead(t, kernelConn, []byte("server-in")) + + _, err = kernelConn.Write([]byte("server-out")) + require.NoError(t, err) + requireConnRead(t, clientConn, []byte("server-out")) +} + +func TestServerHandleImportRelayClosesHandoffOnSockfdFailure(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + device := newTestDevice("1-1", 0x1d6b, 0x0002, "serial-1", SpeedHigh) + store := newTestDeviceStore(device) + store.setStatus("1-1", usbipStatusAvailable) + + expectedErr := errors.New("sockfd handoff failed") + kernelConnCh := make(chan net.Conn, 1) + kernelErrCh := make(chan error, 1) + ops := newTestUSBIPOps(t) + ops.readUsbipStatus = store.readUsbipStatus + ops.readSysfsDevice = store.readSysfsDevice + ops.writeUsbipSockfd = func(busid string, fd int) error { + if fd < 0 { + return nil + } + if busid != "1-1" { + kernelErrCh <- fmt.Errorf("unexpected busid %s", busid) + return expectedErr + } + kernelConn, err := duplicateNetConnFromFD(uintptr(fd), "usbip-server-sockfd-failure-kernel") + if err != nil { + kernelErrCh <- err + } else { + kernelConnCh <- kernelConn + } + return expectedErr + } + + server := &ServerService{ + ctx: ctx, + cancel: cancel, + logger: newTestLogger(), + exports: map[string]serverExport{"1-1": {busid: "1-1"}}, + controlSubs: make(map[uint64]*serverControlConn), + ops: ops, + } + + serverConn, clientConn := net.Pipe() + defer clientConn.Close() + go server.dispatchConn(opaqueConn{Conn: serverConn}) + + setConnDeadline(t, clientConn) + require.NoError(t, WriteOpReqImport(clientConn, "1-1")) + header, err := ReadOpHeader(clientConn) + require.NoError(t, err) + require.Equal(t, OpRepImport, header.Code) + require.Equal(t, OpStatusError, header.Status) + + var kernelConn net.Conn + select { + case kernelConn = <-kernelConnCh: + case err = <-kernelErrCh: + require.NoError(t, err) + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for failed server relay kernel conn") + } + defer kernelConn.Close() + setConnDeadline(t, kernelConn) + requireConnEOF(t, kernelConn) +} + +func TestServerHandleImportRelayClosesHandoffOnReplyFailure(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + device := newTestDevice("1-1", 0x1d6b, 0x0002, "serial-1", SpeedHigh) + store := newTestDeviceStore(device) + store.setStatus("1-1", usbipStatusAvailable) + + kernelConnCh := make(chan net.Conn, 1) + kernelErrCh := make(chan error, 1) + rollbackCh := make(chan string, 1) + allowReply := make(chan struct{}) + ops := newTestUSBIPOps(t) + ops.readUsbipStatus = store.readUsbipStatus + ops.readSysfsDevice = store.readSysfsDevice + ops.writeUsbipSockfd = func(busid string, fd int) error { + if fd < 0 { + rollbackCh <- busid + return nil + } + if busid != "1-1" { + kernelErrCh <- fmt.Errorf("unexpected busid %s", busid) + <-allowReply + return nil + } + kernelConn, err := duplicateNetConnFromFD(uintptr(fd), "usbip-server-reply-failure-kernel") + if err != nil { + kernelErrCh <- err + } else { + kernelConnCh <- kernelConn + } + <-allowReply + return nil + } + + server := &ServerService{ + ctx: ctx, + cancel: cancel, + logger: newTestLogger(), + exports: map[string]serverExport{"1-1": {busid: "1-1"}}, + controlSubs: make(map[uint64]*serverControlConn), + ops: ops, + } + + serverConn, clientConn := net.Pipe() + go server.dispatchConn(opaqueConn{Conn: serverConn}) + + setConnDeadline(t, clientConn) + require.NoError(t, WriteOpReqImport(clientConn, "1-1")) + + var kernelConn net.Conn + select { + case kernelConn = <-kernelConnCh: + case err := <-kernelErrCh: + require.NoError(t, err) + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for reply-failure relay kernel conn") + } + defer kernelConn.Close() + require.NoError(t, clientConn.Close()) + close(allowReply) + + select { + case busid := <-rollbackCh: + require.Equal(t, "1-1", busid) + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for import rollback") + } + setConnDeadline(t, kernelConn) + requireConnEOF(t, kernelConn) +} + func TestServerDispatchConnHandlesControlPingAndChanged(t *testing.T) { t.Parallel() @@ -912,6 +1279,224 @@ func TestClientAttemptAttachUsesImportReplyAndVHCIAttach(t *testing.T) { require.Positive(t, store.lastSockfd("1-1")) } +func TestClientAttemptAttachWithOpaqueConnRelay(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) + serverConnCh := make(chan net.Conn, 1) + serverErrCh := make(chan error, 1) + serverDone := make(chan struct{}) + go func() { + conn, acceptErr := listener.Accept() + if acceptErr != nil { + serverErrCh <- acceptErr + return + } + header, readErr := ReadOpHeader(conn) + if readErr != nil { + _ = conn.Close() + serverErrCh <- readErr + return + } + if header.Code != OpReqImport { + _ = conn.Close() + serverErrCh <- fmt.Errorf("unexpected request code 0x%s", hex16(header.Code)) + return + } + busid, readErr := ReadOpReqImportBody(conn) + if readErr != nil { + _ = conn.Close() + serverErrCh <- readErr + return + } + if busid != "1-1" { + _ = conn.Close() + serverErrCh <- fmt.Errorf("unexpected busid %s", busid) + return + } + info := device.toProtocol() + if writeErr := WriteOpRepImport(conn, OpStatusOK, &info); writeErr != nil { + _ = conn.Close() + serverErrCh <- writeErr + return + } + serverConnCh <- conn + <-serverDone + _ = conn.Close() + serverErrCh <- nil + }() + defer close(serverDone) + + kernelConnCh := make(chan net.Conn, 1) + ops := newTestUSBIPOps(t) + ops.vhciPickFreePort = func(speed uint32) (int, error) { + require.Equal(t, SpeedHigh, speed) + return 4, nil + } + ops.vhciAttach = func(port int, fd uintptr, devid uint32, speed uint32) error { + require.Equal(t, 4, port) + requireStreamSocketFD(t, fd) + info := device.toProtocol() + require.Equal(t, info.DevID(), devid) + require.Equal(t, SpeedHigh, speed) + kernelConnCh <- duplicateConnFromFD(t, fd, "usbip-client-test-kernel") + return nil + } + + client := &ClientService{ + ctx: ctx, + cancel: cancel, + logger: newTestLogger(), + dialer: wrappingDialer{}, + serverAddr: M.SocksaddrFromNet(listener.Addr()), + ops: ops, + } + + port, err := client.attemptAttach(ctx, "1-1") + require.NoError(t, err) + require.Equal(t, 4, port) + + var serverConn net.Conn + select { + case serverConn = <-serverConnCh: + case err = <-serverErrCh: + require.NoError(t, err) + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for server conn") + } + var kernelConn net.Conn + select { + case kernelConn = <-kernelConnCh: + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for client relay kernel conn") + } + defer kernelConn.Close() + setConnDeadline(t, serverConn) + setConnDeadline(t, kernelConn) + + _, err = serverConn.Write([]byte("client-in")) + require.NoError(t, err) + requireConnRead(t, kernelConn, []byte("client-in")) + + _, err = kernelConn.Write([]byte("client-out")) + require.NoError(t, err) + requireConnRead(t, serverConn, []byte("client-out")) +} + +func TestClientAttemptAttachRelayClosesHandoffOnVHCIAttachFailure(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) + serverErrCh := make(chan error, 1) + go func() { + conn, acceptErr := listener.Accept() + if acceptErr != nil { + serverErrCh <- acceptErr + return + } + defer conn.Close() + header, readErr := ReadOpHeader(conn) + if readErr != nil { + serverErrCh <- readErr + return + } + if header.Code != OpReqImport { + serverErrCh <- fmt.Errorf("unexpected request code 0x%s", hex16(header.Code)) + return + } + busid, readErr := ReadOpReqImportBody(conn) + if readErr != nil { + serverErrCh <- readErr + return + } + if busid != "1-1" { + serverErrCh <- fmt.Errorf("unexpected busid %s", busid) + return + } + info := device.toProtocol() + if writeErr := WriteOpRepImport(conn, OpStatusOK, &info); writeErr != nil { + serverErrCh <- writeErr + return + } + buffer := make([]byte, 1) + n, readErr := conn.Read(buffer) + if n != 0 { + serverErrCh <- fmt.Errorf("unexpected server read bytes after attach failure: %d", n) + return + } + if !errors.Is(readErr, io.EOF) { + serverErrCh <- readErr + return + } + serverErrCh <- nil + }() + + expectedErr := errors.New("vhci attach failed") + kernelConnCh := make(chan net.Conn, 1) + ops := newTestUSBIPOps(t) + ops.vhciPickFreePort = func(speed uint32) (int, error) { + require.Equal(t, SpeedHigh, speed) + return 4, nil + } + ops.vhciAttach = func(port int, fd uintptr, devid uint32, speed uint32) error { + require.Equal(t, 4, port) + requireStreamSocketFD(t, fd) + info := device.toProtocol() + require.Equal(t, info.DevID(), devid) + require.Equal(t, SpeedHigh, speed) + kernelConnCh <- duplicateConnFromFD(t, fd, "usbip-client-vhci-failure-kernel") + return expectedErr + } + + client := &ClientService{ + ctx: ctx, + cancel: cancel, + logger: newTestLogger(), + dialer: wrappingDialer{}, + serverAddr: M.SocksaddrFromNet(listener.Addr()), + ops: ops, + } + + port, err := client.attemptAttach(ctx, "1-1") + require.Equal(t, -1, port) + require.ErrorIs(t, err, expectedErr) + + var kernelConn net.Conn + select { + case kernelConn = <-kernelConnCh: + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for failed client relay kernel conn") + } + defer kernelConn.Close() + setConnDeadline(t, kernelConn) + requireConnEOF(t, kernelConn) + + select { + case err = <-serverErrCh: + require.NoError(t, err) + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for server side close") + } + client.portsMu.Lock() + _, reserved := client.ports[4] + client.portsMu.Unlock() + require.False(t, reserved) +} + func TestClientFetchDevListRejectsUnexpectedReplyVersion(t *testing.T) { t.Parallel() diff --git a/service/usbip/server_linux.go b/service/usbip/server_linux.go index e990a3a3a..469e3c00b 100644 --- a/service/usbip/server_linux.go +++ b/service/usbip/server_linux.go @@ -23,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" ) @@ -345,12 +346,17 @@ func (s *ServerService) dispatchConn(conn net.Conn) { } func (s *ServerService) handleStandardConn(conn net.Conn, header OpHeader) { - defer conn.Close() + closeConn := true + defer func() { + if closeConn { + _ = conn.Close() + } + }() switch header.Code { case OpReqDevList: s.handleDevList(conn) case OpReqImport: - s.handleImport(conn) + closeConn = !s.handleImport(conn) default: s.logger.Debug("unknown opcode 0x", hex16(header.Code)) } @@ -456,54 +462,53 @@ func (s *ServerService) buildDevListEntries() []DeviceEntry { return entries } -func (s *ServerService) handleImport(conn net.Conn) { +func (s *ServerService) handleImport(conn net.Conn) bool { busid, err := ReadOpReqImportBody(conn) if err != nil { s.logger.Debug("read import body: ", err) - return + return false } if !s.isExported(busid) { s.logger.Info("import rejected (unknown busid): ", busid) _ = WriteOpRepImport(conn, OpStatusError, nil) - return + return false } status, err := s.ops.readUsbipStatus(busid) if err != nil || status != usbipStatusAvailable { s.logger.Info("import rejected (busid ", busid, " status=", status, " err=", err, ")") _ = WriteOpRepImport(conn, OpStatusError, nil) - return + return false } dev, err := s.ops.readSysfsDevice(busid, sysBusDevicePath(busid)) if err != nil { s.logger.Warn("refresh ", busid, ": ", err) _ = WriteOpRepImport(conn, OpStatusError, nil) - return + return false } - tcp, ok := conn.(*net.TCPConn) - if !ok { - s.logger.Warn("import requires *net.TCPConn, got ", conn) - _ = WriteOpRepImport(conn, OpStatusError, nil) - return - } - file, err := tcp.File() + handoff, err := newUSBIPConnHandoff(conn) if err != nil { - s.logger.Warn("dup socket fd: ", err) + s.logger.Warn("prepare handoff ", busid, ": ", err) _ = WriteOpRepImport(conn, OpStatusError, nil) - return + return false } - defer file.Close() - if err := s.ops.writeUsbipSockfd(busid, int(file.Fd())); err != nil { + defer handoff.Close() + s.logger.Debug("usbip server handoff ", busid, ": ", handoff.mode()) + if err := s.ops.writeUsbipSockfd(busid, int(handoff.kernelFD())); err != nil { s.logger.Warn("hand off ", busid, " to kernel: ", err) _ = WriteOpRepImport(conn, OpStatusError, nil) - return + return false + } + if err := handoff.closeKernelFD(); err != nil { + s.logger.Debug("close kernel fd ", busid, ": ", err) } info := dev.toProtocol() if err := WriteOpRepImport(conn, OpStatusOK, &info); err != nil { s.logger.Warn("reply import ", busid, ": ", err) _ = s.ops.writeUsbipSockfd(busid, -1) - return + return false } s.logger.Info("attached ", busid, " to remote ", conn.RemoteAddr()) + return handoff.startRelay(s.ctx, s.logger, "server", busid) } func (s *ServerService) isExported(busid string) bool {