usbip: add linux socketpair relay fallback
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user