From ebd569670d43d3eab79afd02663455569e7c1b16 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Mon, 6 Jul 2026 21:05:21 +0800 Subject: [PATCH] Drain pending packets in Linux batch read --- tun_linux.go | 108 ++++++++++++++++++++++++++++++++++++++++++++++----- 1 file changed, 98 insertions(+), 10 deletions(-) diff --git a/tun_linux.go b/tun_linux.go index 9cd6300..8761615 100644 --- a/tun_linux.go +++ b/tun_linux.go @@ -39,6 +39,9 @@ type NativeTun struct { writeAccess sync.Mutex vnetHdr bool writeBuffer []byte + readRawConn syscall.RawConn + pendingBuffer []byte + pendingLength int vnetHdrWriteBuf []byte gsoToWrite []int tcpGROTable *tcpGROTable @@ -192,6 +195,11 @@ func (t *NativeTun) enableGSO() error { } t.vnetHdr = true t.writeBuffer = make([]byte, virtioNetHdrLen+int(gsoMaxSize)) + t.pendingBuffer = make([]byte, virtioNetHdrLen+int(gsoMaxSize)) + t.readRawConn, err = t.tunFile.SyscallConn() + if err != nil { + return E.Cause(err, "enable offload: get raw conn") + } t.tcpGROTable = newTCPGROTable() t.udpGROTable = newUDPGROTable() err = setUDPOffload(t.tunFd) @@ -364,16 +372,24 @@ func (t *NativeTun) Read(p []byte) (n int, err error) { // each buffer. It mutates sizes to reflect the size of each element of bufs, // and returns the number of packets read. func handleVirtioRead(in []byte, bufs [][]byte, sizes []int, offset int) (int, error) { + payload, options, err := parseVirtioRead(in) + if err != nil { + return 0, err + } + return GSOSplit(payload, options, bufs, sizes, offset) +} + +func parseVirtioRead(in []byte) ([]byte, GSOOptions, error) { var hdr virtioNetHdr err := hdr.decode(in) if err != nil { - return 0, err + return nil, GSOOptions{}, err } in = in[virtioNetHdrLen:] options, err := hdr.toGSOOptions() if err != nil { - return 0, err + return nil, GSOOptions{}, err } // Don't trust HdrLen from the kernel as it can be equal to the length @@ -384,18 +400,26 @@ func handleVirtioRead(in []byte, bufs [][]byte, sizes []int, offset int) (int, e options.HdrLen = options.CsumStart + 8 } else if options.GSOType != GSONone { if len(in) <= int(options.CsumStart+12) { - return 0, errors.New("packet is too short") + return nil, GSOOptions{}, errors.New("packet is too short") } tcpHLen := uint16(in[options.CsumStart+12] >> 4 * 4) if tcpHLen < 20 || tcpHLen > 60 { // A TCP header must be between 20 and 60 bytes in length. - return 0, fmt.Errorf("tcp header len is invalid: %d", tcpHLen) + return nil, GSOOptions{}, fmt.Errorf("tcp header len is invalid: %d", tcpHLen) } options.HdrLen = options.CsumStart + tcpHLen } - return GSOSplit(in, options, bufs, sizes, offset) + return in, options, nil +} + +func gsoSegmentCount(payload []byte, options GSOOptions) int { + dataLength := len(payload) - int(options.HdrLen) + if options.GSOType == GSONone || options.GSOSize == 0 || dataLength < int(options.GSOSize) { + return 1 + } + return (dataLength + int(options.GSOSize) - 1) / int(options.GSOSize) } func (t *NativeTun) Write(p []byte) (n int, err error) { @@ -430,14 +454,78 @@ func (t *NativeTun) BatchSize() int { return idealBatchSize } -func (t *NativeTun) BatchRead(buffers [][]byte, offset int, readN []int) (n int, err error) { +func (t *NativeTun) BatchRead(buffers [][]byte, offset int, readN []int) (int, error) { t.readAccess.Lock() defer t.readAccess.Unlock() - n, err = t.tunFile.Read(t.writeBuffer) - if err != nil { - return + var used int + if t.pendingLength > 0 { + pendingLength := t.pendingLength + t.pendingLength = 0 + count, err := handleVirtioRead(t.pendingBuffer[:pendingLength], buffers, readN, offset) + if err != nil { + return count, err + } + used = count } - return handleVirtioRead(t.writeBuffer[:n], buffers, readN, offset) + for used < len(buffers) { + var ( + readLength int + err error + ) + if used == 0 { + readLength, err = t.tunFile.Read(t.writeBuffer) + if err != nil { + return 0, err + } + } else { + readLength, err = t.readNonblocking(t.writeBuffer) + if err != nil || readLength == 0 { + break + } + } + payload, options, parseErr := parseVirtioRead(t.writeBuffer[:readLength]) + if parseErr != nil { + if used > 0 { + break + } + return 0, parseErr + } + if used > 0 && gsoSegmentCount(payload, options) > len(buffers)-used { + t.writeBuffer, t.pendingBuffer = t.pendingBuffer, t.writeBuffer + t.pendingLength = readLength + break + } + count, splitErr := GSOSplit(payload, options, buffers[used:], readN[used:], offset) + if splitErr != nil { + if used > 0 { + break + } + return count, splitErr + } + used += count + } + return used, nil +} + +func (t *NativeTun) readNonblocking(buffer []byte) (int, error) { + var ( + readLength int + readErr error + ) + controlErr := t.readRawConn.Read(func(fd uintptr) bool { + readLength, readErr = syscall.Read(int(fd), buffer) + return true + }) + if controlErr != nil { + return 0, controlErr + } + if readErr != nil { + if readErr == syscall.EAGAIN { + return 0, nil + } + return 0, readErr + } + return readLength, nil } func (t *NativeTun) BatchWrite(buffers [][]byte, offset int) (int, error) {