usbip: add linux socketpair relay fallback

This commit is contained in:
世界
2026-04-24 04:47:10 +08:00
parent 5450351607
commit 1cf178440f
4 changed files with 737 additions and 31 deletions
+23 -11
View File
@@ -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
}
+104
View File
@@ -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
}
+585
View File
@@ -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()
+25 -20
View File
@@ -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 {