diff --git a/service/usbip/client_darwin.go b/service/usbip/client_darwin.go index 41c40c2fc..fe8cf2eef 100644 --- a/service/usbip/client_darwin.go +++ b/service/usbip/client_darwin.go @@ -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) } } diff --git a/service/usbip/darwin_integration_test.go b/service/usbip/darwin_integration_test.go index ddcb474f2..6768b5575 100644 --- a/service/usbip/darwin_integration_test.go +++ b/service/usbip/darwin_integration_test.go @@ -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 } diff --git a/service/usbip/data_protocol.go b/service/usbip/data_protocol.go index ad007cc0c..cc6b86459 100644 --- a/service/usbip/data_protocol.go +++ b/service/usbip/data_protocol.go @@ -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 { diff --git a/service/usbip/data_protocol_test.go b/service/usbip/data_protocol_test.go index e82057eb8..084700189 100644 --- a/service/usbip/data_protocol_test.go +++ b/service/usbip/data_protocol_test.go @@ -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") }