Fix USB/IP Darwin protocol handling

This commit is contained in:
世界
2026-04-24 04:53:54 +08:00
parent 1cf178440f
commit 3cfe767ae4
4 changed files with 290 additions and 33 deletions
+57 -11
View File
@@ -485,7 +485,7 @@ func (c *ClientService) runBusIDLoop(ctx context.Context, busid, description str
}
c.logger.Info("attached ", busid, " through IOUSBHostControllerInterface")
c.setBusIDActive(busid, true)
controller.Wait()
waitDarwinController(ctx, controller)
c.setBusIDActive(busid, false)
if err := ctx.Err(); err != nil {
controller.Close()
@@ -502,6 +502,15 @@ func (c *ClientService) runBusIDLoop(ctx context.Context, busid, description str
}
}
func waitDarwinController(ctx context.Context, controller *darwinVirtualController) {
select {
case <-controller.done:
case <-ctx.Done():
controller.Close()
controller.Wait()
}
}
func (c *ClientService) attemptAttach(ctx context.Context, busid string) (*darwinVirtualController, error) {
conn, err := c.dialer.DialContext(ctx, N.NetworkTCP, c.serverAddr)
if err != nil {
@@ -592,6 +601,19 @@ type darwinControlState struct {
setup [8]byte
}
type darwinPendingSubmit struct {
direction uint32
reply chan SubmitResponse
}
type darwinEndpointStateMachine interface {
Close()
respond(message darwinCIMessage, status int) error
processDoorbell(doorbell uint32) error
currentTransfer() darwinCITransfer
complete(transfer darwinCITransfer, status int, length int) error
}
type darwinVirtualController struct {
ctx context.Context
cancel context.CancelFunc
@@ -608,14 +630,14 @@ type darwinVirtualController struct {
writeMu sync.Mutex
pendingMu sync.Mutex
pending map[uint32]chan SubmitResponse
pending map[uint32]darwinPendingSubmit
stateMu sync.Mutex
powered bool
connected bool
nextAddress uint8
devices map[uint8]*darwinUSBHostDeviceSM
endpoints map[darwinEndpointKey]*darwinUSBHostEndpointSM
endpoints map[darwinEndpointKey]darwinEndpointStateMachine
controlStates map[uint8]darwinControlState
}
@@ -630,10 +652,10 @@ func newDarwinVirtualController(ctx context.Context, logger log.ContextLogger, c
startTime: time.Now(),
events: make(chan darwinControllerEvent, 64),
done: make(chan struct{}),
pending: make(map[uint32]chan SubmitResponse),
pending: make(map[uint32]darwinPendingSubmit),
nextAddress: 1,
devices: make(map[uint8]*darwinUSBHostDeviceSM),
endpoints: make(map[darwinEndpointKey]*darwinUSBHostEndpointSM),
endpoints: make(map[darwinEndpointKey]darwinEndpointStateMachine),
controlStates: make(map[uint8]darwinControlState),
}
}
@@ -703,7 +725,11 @@ func (c *darwinVirtualController) readLoop() {
}
switch header.Command {
case RetSubmit:
response, err := ReadSubmitResponseBody(c.conn, header)
payloadDirection, ok := c.pendingSubmitDirection(header.SeqNum)
if !ok {
payloadDirection = header.Direction
}
response, err := ReadSubmitResponseBody(c.conn, header, payloadDirection)
if err != nil {
c.logger.Debug("read RET_SUBMIT: ", err)
c.failPending()
@@ -862,6 +888,9 @@ func (c *darwinVirtualController) handleDoorbell(doorbell uint32) {
return
}
status, length := c.handleTransfer(key, transfer.message)
if transfer.message.noResponse() {
return
}
if err := endpoint.complete(transfer, darwinUSBIPStatusToCIStatus(status), length); err != nil {
c.logger.Debug("complete transfer: ", err)
c.Close()
@@ -914,6 +943,7 @@ func (c *darwinVirtualController) handleControlDataTransfer(key darwinEndpointKe
Endpoint: 0,
},
TransferBufferLength: int32(length),
NumberOfPackets: nonIsoPacketCount,
Setup: state.setup,
Buffer: buffer,
})
@@ -944,7 +974,8 @@ func (c *darwinVirtualController) handleControlStatusTransfer(key darwinEndpoint
Direction: USBIPDirOut,
Endpoint: 0,
},
Setup: state.setup,
NumberOfPackets: nonIsoPacketCount,
Setup: state.setup,
})
if err != nil {
return -int32(unix.EIO), 0
@@ -969,6 +1000,7 @@ func (c *darwinVirtualController) handleNormalTransfer(key darwinEndpointKey, me
Endpoint: uint32(key.endpoint & 0x0f),
},
TransferBufferLength: int32(length),
NumberOfPackets: nonIsoPacketCount,
Buffer: buffer,
})
if err != nil {
@@ -1017,9 +1049,12 @@ func (c *darwinVirtualController) handleIsoTransfer(key darwinEndpointKey, messa
func (c *darwinVirtualController) sendSubmit(command SubmitCommand) (SubmitResponse, error) {
seq := c.seq.Add(1)
command.Header.SeqNum = seq
if command.NumberOfPackets == 0 && len(command.IsoPackets) == 0 {
command.NumberOfPackets = nonIsoPacketCount
}
reply := make(chan SubmitResponse, 1)
c.pendingMu.Lock()
c.pending[seq] = reply
c.pending[seq] = darwinPendingSubmit{direction: command.Header.Direction, reply: reply}
c.pendingMu.Unlock()
defer func() {
c.pendingMu.Lock()
@@ -1043,10 +1078,21 @@ func (c *darwinVirtualController) sendSubmit(command SubmitCommand) (SubmitRespo
}
}
func (c *darwinVirtualController) pendingSubmitDirection(seq uint32) (uint32, bool) {
c.pendingMu.Lock()
defer c.pendingMu.Unlock()
pending, ok := c.pending[seq]
if !ok {
return 0, false
}
return pending.direction, true
}
func (c *darwinVirtualController) deliverSubmit(response SubmitResponse) {
c.pendingMu.Lock()
reply := c.pending[response.Header.SeqNum]
pending := c.pending[response.Header.SeqNum]
c.pendingMu.Unlock()
reply := pending.reply
if reply == nil {
return
}
@@ -1059,9 +1105,9 @@ func (c *darwinVirtualController) deliverSubmit(response SubmitResponse) {
func (c *darwinVirtualController) failPending() {
c.pendingMu.Lock()
defer c.pendingMu.Unlock()
for seq, reply := range c.pending {
for seq, pending := range c.pending {
delete(c.pending, seq)
close(reply)
close(pending.reply)
}
}
+143 -3
View File
@@ -12,6 +12,7 @@ import (
"sync"
"testing"
"time"
"unsafe"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/log"
@@ -114,6 +115,144 @@ func requireDarwinUserHCI(t *testing.T) {
hostController.Close()
}
func TestDarwinVirtualControllerReadsCompliantSubmitResponsePayload(t *testing.T) {
t.Parallel()
clientConn, serverConn := net.Pipe()
defer serverConn.Close()
controller := newDarwinVirtualController(context.Background(), newTestLogger(), clientConn, DeviceInfoTruncated{
BusNum: 1,
DevNum: 1,
})
go controller.readLoop()
t.Cleanup(controller.Close)
responseCh := make(chan SubmitResponse, 1)
errCh := make(chan error, 1)
go func() {
response, err := controller.sendSubmit(SubmitCommand{
Header: DataHeader{
Command: CmdSubmit,
DevID: 0x00010001,
Direction: USBIPDirIn,
Endpoint: 1,
},
TransferBufferLength: 3,
})
if err != nil {
errCh <- err
return
}
responseCh <- response
}()
header, err := ReadDataHeader(serverConn)
require.NoError(t, err)
command, err := ReadSubmitCommandBody(serverConn, header)
require.NoError(t, err)
require.Equal(t, USBIPDirIn, command.Header.Direction)
require.Equal(t, int32(nonIsoPacketCount), command.NumberOfPackets)
require.NoError(t, WriteSubmitResponse(serverConn, SubmitResponse{
Header: DataHeader{
Command: RetSubmit,
SeqNum: header.SeqNum,
Direction: USBIPDirIn,
},
Status: 0,
ActualLength: 3,
Buffer: []byte{1, 2, 3},
}))
select {
case err := <-errCh:
require.NoError(t, err)
case response := <-responseCh:
require.Equal(t, DataHeader{Command: RetSubmit, SeqNum: header.SeqNum}, response.Header)
require.Equal(t, int32(nonIsoPacketCount), response.NumberOfPackets)
require.Equal(t, []byte{1, 2, 3}, response.Buffer)
case <-time.After(5 * time.Second):
t.Fatal("timed out waiting for submit response")
}
}
func TestWaitDarwinControllerClosesOnContextCancel(t *testing.T) {
t.Parallel()
clientConn, serverConn := net.Pipe()
defer serverConn.Close()
controller := newDarwinVirtualController(context.Background(), newTestLogger(), clientConn, DeviceInfoTruncated{})
go controller.readLoop()
ctx, cancel := context.WithCancel(context.Background())
cancel()
done := make(chan struct{})
go func() {
waitDarwinController(ctx, controller)
close(done)
}()
select {
case <-done:
case <-time.After(5 * time.Second):
t.Fatal("timed out waiting for controller cancellation")
}
select {
case <-controller.done:
default:
t.Fatal("controller read loop still active after cancellation")
}
}
type fakeDarwinEndpointStateMachine struct {
transfer darwinCITransfer
currentRead bool
completeCalled int
}
func (f *fakeDarwinEndpointStateMachine) Close() {}
func (f *fakeDarwinEndpointStateMachine) respond(darwinCIMessage, int) error {
return nil
}
func (f *fakeDarwinEndpointStateMachine) processDoorbell(uint32) error {
return nil
}
func (f *fakeDarwinEndpointStateMachine) currentTransfer() darwinCITransfer {
if f.currentRead {
return darwinCITransfer{}
}
f.currentRead = true
return f.transfer
}
func (f *fakeDarwinEndpointStateMachine) complete(darwinCITransfer, int, int) error {
f.completeCalled++
return nil
}
func TestDarwinHandleDoorbellSkipsNoResponseCompletion(t *testing.T) {
t.Parallel()
controller := newDarwinVirtualController(context.Background(), newTestLogger(), nil, DeviceInfoTruncated{})
message := darwinCIMessage{
control: (1 << 15) | (1 << 14) | 0x3c,
data0: (uint32(2) << 8) | 1,
}
endpoint := &fakeDarwinEndpointStateMachine{
transfer: darwinCITransfer{
ptr: unsafe.Pointer(&message),
message: message,
},
}
controller.endpoints[darwinEndpointKey{device: 1, endpoint: 2}] = endpoint
controller.handleDoorbell((uint32(2) << 8) | 1)
require.Zero(t, endpoint.completeCalled)
}
func startDarwinFakeUSBIPServer(t *testing.T) *darwinFakeUSBIPServer {
t.Helper()
@@ -300,7 +439,8 @@ func (s *darwinFakeUSBIPServer) submitResponse(command SubmitCommand) SubmitResp
Direction: command.Header.Direction,
Endpoint: command.Header.Endpoint,
},
Status: 0,
Status: 0,
NumberOfPackets: nonIsoPacketCount,
}
if command.Header.Endpoint != 0 {
return response
@@ -459,7 +599,7 @@ func TestDarwinUSBIPServerSelectedDeviceConfiguresDevice(t *testing.T) {
dataHeader, err := ReadDataHeader(conn)
require.NoError(t, err)
require.Equal(t, RetSubmit, dataHeader.Command)
response, err := ReadSubmitResponseBody(conn, dataHeader)
response, err := ReadSubmitResponseBody(conn, dataHeader, USBIPDirIn)
require.NoError(t, err)
require.Equal(t, int32(0), response.Status)
require.GreaterOrEqual(t, len(response.Buffer), 18)
@@ -517,7 +657,7 @@ func darwinSubmitControl(t *testing.T, conn net.Conn, command SubmitCommand) Sub
require.NoError(t, err)
require.Equal(t, RetSubmit, dataHeader.Command)
require.Equal(t, command.Header.SeqNum, dataHeader.SeqNum)
response, err := ReadSubmitResponseBody(conn, dataHeader)
response, err := ReadSubmitResponseBody(conn, dataHeader, command.Header.Direction)
require.NoError(t, err)
return response
}
+33 -11
View File
@@ -21,6 +21,7 @@ const (
isoPacketDescriptorWireSize = 16
maxUSBIPTransferBufferLength = 16 << 20
maxUSBIPIsoPackets = 4096
nonIsoPacketCount = -1
)
type DataHeader struct {
@@ -109,7 +110,7 @@ func ReadSubmitCommandBody(r io.Reader, header DataHeader) (SubmitCommand, error
return command, nil
}
func ReadSubmitResponseBody(r io.Reader, header DataHeader) (SubmitResponse, error) {
func ReadSubmitResponseBody(r io.Reader, header DataHeader, payloadDirection uint32) (SubmitResponse, error) {
var raw [28]byte
if _, err := io.ReadFull(r, raw[:]); err != nil {
return SubmitResponse{}, err
@@ -127,7 +128,7 @@ func ReadSubmitResponseBody(r io.Reader, header DataHeader) (SubmitResponse, err
if bufferLength < 0 {
bufferLength = 0
}
buffer, isoPackets, err := readUSBIPPayload(r, header.Direction, bufferLength, response.NumberOfPackets, false)
buffer, isoPackets, err := readUSBIPPayload(r, payloadDirection, bufferLength, response.NumberOfPackets, false)
if err != nil {
return SubmitResponse{}, err
}
@@ -162,7 +163,8 @@ func WriteSubmitCommand(w io.Writer, command SubmitCommand) error {
if err := validateUSBIPBufferLength(command.TransferBufferLength); err != nil {
return err
}
if err := validateUSBIPIsoPacketCount(command.NumberOfPackets); err != nil {
packetCount := normalizeUSBIPIsoPacketCount(command.NumberOfPackets, command.IsoPackets)
if err := validateUSBIPIsoPacketCount(packetCount); err != nil {
return err
}
if err := writeDataHeader(w, command.Header); err != nil {
@@ -172,7 +174,7 @@ func WriteSubmitCommand(w io.Writer, command SubmitCommand) error {
binary.BigEndian.PutUint32(raw[0:4], uint32(command.TransferFlags))
binary.BigEndian.PutUint32(raw[4:8], uint32(command.TransferBufferLength))
binary.BigEndian.PutUint32(raw[8:12], uint32(command.StartFrame))
binary.BigEndian.PutUint32(raw[12:16], uint32(command.NumberOfPackets))
binary.BigEndian.PutUint32(raw[12:16], uint32(packetCount))
binary.BigEndian.PutUint32(raw[16:20], uint32(command.Interval))
copy(raw[20:28], command.Setup[:])
if _, err := w.Write(raw[:]); err != nil {
@@ -188,23 +190,26 @@ func WriteSubmitResponse(w io.Writer, response SubmitResponse) error {
if err := validateUSBIPBufferLength(response.ActualLength); err != nil {
return err
}
if err := validateUSBIPIsoPacketCount(response.NumberOfPackets); err != nil {
packetCount := normalizeUSBIPIsoPacketCount(response.NumberOfPackets, response.IsoPackets)
if err := validateUSBIPIsoPacketCount(packetCount); err != nil {
return err
}
if err := writeDataHeader(w, response.Header); err != nil {
payloadDirection := response.Header.Direction
header := responseDataHeader(response.Header)
if err := writeDataHeader(w, header); err != nil {
return err
}
var raw [28]byte
binary.BigEndian.PutUint32(raw[0:4], uint32(response.Status))
binary.BigEndian.PutUint32(raw[4:8], uint32(response.ActualLength))
binary.BigEndian.PutUint32(raw[8:12], uint32(response.StartFrame))
binary.BigEndian.PutUint32(raw[12:16], uint32(response.NumberOfPackets))
binary.BigEndian.PutUint32(raw[12:16], uint32(packetCount))
binary.BigEndian.PutUint32(raw[16:20], uint32(response.ErrorCount))
copy(raw[20:28], response.Setup[:])
if _, err := w.Write(raw[:]); err != nil {
return err
}
return writeUSBIPPayload(w, response.Header.Direction, response.Buffer, response.IsoPackets, false)
return writeUSBIPPayload(w, payloadDirection, response.Buffer, response.IsoPackets, false)
}
func WriteUnlinkCommand(w io.Writer, command UnlinkCommand) error {
@@ -218,7 +223,7 @@ func WriteUnlinkCommand(w io.Writer, command UnlinkCommand) error {
}
func WriteUnlinkResponse(w io.Writer, response UnlinkResponse) error {
if err := writeDataHeader(w, response.Header); err != nil {
if err := writeDataHeader(w, responseDataHeader(response.Header)); err != nil {
return err
}
var raw [unlinkBodySize]byte
@@ -276,8 +281,25 @@ func shouldCarryUSBIPBuffer(direction uint32, command bool) bool {
return direction == USBIPDirIn
}
func normalizeUSBIPIsoPacketCount(count int32, packets []IsoPacketDescriptor) int32 {
if count == 0 && len(packets) == 0 {
return nonIsoPacketCount
}
if count == 0 && len(packets) > 0 {
return int32(len(packets))
}
return count
}
func responseDataHeader(header DataHeader) DataHeader {
return DataHeader{
Command: header.Command,
SeqNum: header.SeqNum,
}
}
func readUSBIPIsoPackets(r io.Reader, count int32) ([]IsoPacketDescriptor, error) {
if count == 0 {
if count <= 0 {
return nil, nil
}
packets := make([]IsoPacketDescriptor, int(count))
@@ -321,7 +343,7 @@ func validateUSBIPBufferLength(length int32) error {
}
func validateUSBIPIsoPacketCount(count int32) error {
if count < 0 {
if count < nonIsoPacketCount {
return E.New("USB/IP iso packet count is negative: ", count)
}
if count > maxUSBIPIsoPackets {
+57 -8
View File
@@ -25,7 +25,7 @@ func TestUSBIPSubmitCommandRoundTripOut(t *testing.T) {
TransferFlags: 0x400,
TransferBufferLength: 3,
StartFrame: 11,
NumberOfPackets: 0,
NumberOfPackets: nonIsoPacketCount,
Interval: 4,
Setup: [8]byte{0, 1, 2, 3, 4, 5, 6, 7},
Buffer: []byte{1, 2, 3},
@@ -55,6 +55,7 @@ func TestUSBIPSubmitCommandRoundTripInOmitsCommandPayload(t *testing.T) {
Endpoint: 1,
},
TransferBufferLength: 4,
NumberOfPackets: nonIsoPacketCount,
Buffer: []byte{9, 8, 7, 6},
}
@@ -99,10 +100,11 @@ func TestUSBIPSubmitResponseRoundTripInWithIsoPackets(t *testing.T) {
header, err := ReadDataHeader(&buffer)
require.NoError(t, err)
require.Equal(t, expected.Header, header)
require.Equal(t, DataHeader{Command: RetSubmit, SeqNum: expected.Header.SeqNum}, header)
actual, err := ReadSubmitResponseBody(&buffer, header)
actual, err := ReadSubmitResponseBody(&buffer, header, expected.Header.Direction)
require.NoError(t, err)
expected.Header = DataHeader{Command: RetSubmit, SeqNum: expected.Header.SeqNum}
require.Equal(t, expected, actual)
}
@@ -117,9 +119,10 @@ func TestUSBIPSubmitResponseRoundTripOutOmitsResponsePayload(t *testing.T) {
Direction: USBIPDirOut,
Endpoint: 2,
},
Status: 0,
ActualLength: 3,
Buffer: []byte{1, 2, 3},
Status: 0,
ActualLength: 3,
NumberOfPackets: nonIsoPacketCount,
Buffer: []byte{1, 2, 3},
}
var buffer bytes.Buffer
@@ -128,13 +131,53 @@ func TestUSBIPSubmitResponseRoundTripOutOmitsResponsePayload(t *testing.T) {
header, err := ReadDataHeader(&buffer)
require.NoError(t, err)
actual, err := ReadSubmitResponseBody(&buffer, header)
actual, err := ReadSubmitResponseBody(&buffer, header, expected.Header.Direction)
require.NoError(t, err)
require.Equal(t, expected.Header, actual.Header)
require.Equal(t, DataHeader{Command: RetSubmit, SeqNum: expected.Header.SeqNum}, actual.Header)
require.Equal(t, expected.ActualLength, actual.ActualLength)
require.Equal(t, expected.NumberOfPackets, actual.NumberOfPackets)
require.Empty(t, actual.Buffer)
}
func TestUSBIPSubmitResponseReadsPayloadFromOriginalDirection(t *testing.T) {
t.Parallel()
expected := SubmitResponse{
Header: DataHeader{
Command: RetSubmit,
SeqNum: 42,
Direction: USBIPDirIn,
},
Status: 0,
ActualLength: 3,
NumberOfPackets: nonIsoPacketCount,
Buffer: []byte{4, 5, 6},
}
var buffer bytes.Buffer
require.NoError(t, WriteSubmitResponse(&buffer, expected))
header, err := ReadDataHeader(&buffer)
require.NoError(t, err)
require.Equal(t, DataHeader{Command: RetSubmit, SeqNum: expected.Header.SeqNum}, header)
actual, err := ReadSubmitResponseBody(&buffer, header, USBIPDirIn)
require.NoError(t, err)
expected.Header = header
require.Equal(t, expected, actual)
}
func TestUSBIPSubmitCommandAcceptsNonISOPacketSentinel(t *testing.T) {
t.Parallel()
var raw [28]byte
binary.BigEndian.PutUint32(raw[12:16], 0xffffffff)
command, err := ReadSubmitCommandBody(bytes.NewReader(raw[:]), DataHeader{Command: CmdSubmit, Direction: USBIPDirIn})
require.NoError(t, err)
require.Equal(t, int32(nonIsoPacketCount), command.NumberOfPackets)
require.Empty(t, command.IsoPackets)
}
func TestUSBIPUnlinkRoundTrip(t *testing.T) {
t.Parallel()
@@ -175,6 +218,7 @@ func TestUSBIPUnlinkRoundTrip(t *testing.T) {
require.NoError(t, err)
actualResponse, err := ReadUnlinkResponseBody(&buffer, header)
require.NoError(t, err)
response.Header = DataHeader{Command: RetUnlink, SeqNum: command.Header.SeqNum}
require.Equal(t, response, actualResponse)
}
@@ -282,4 +326,9 @@ func TestUSBIPRejectsInvalidDataPlaneLengths(t *testing.T) {
binary.BigEndian.PutUint32(raw[12:16], uint32(maxUSBIPIsoPackets+1))
_, err = ReadSubmitCommandBody(bytes.NewReader(raw[:]), DataHeader{Command: CmdSubmit, Direction: USBIPDirIn})
require.ErrorContains(t, err, "too large")
raw = [28]byte{}
binary.BigEndian.PutUint32(raw[12:16], 0xfffffffe)
_, err = ReadSubmitCommandBody(bytes.NewReader(raw[:]), DataHeader{Command: CmdSubmit, Direction: USBIPDirIn})
require.ErrorContains(t, err, "negative")
}