chore: allow the sniffer to read more data if possible

This commit is contained in:
wwqgtxx committed 2026-07-28 15:10:05 +08:00
1 parent cb8e9900f6
commit 2bade2940a
4 files changed
+213 -4

No files matched your search

+26 -1
View File
@@ -2,11 +2,12 @@ package net
import (
"io"
"math/bits"
"unsafe"
)
// bufioReader copy from stdlib bufio/bufio.go
// This structure has remained unchanged from go1.5 to go1.21.
// This structure has remained unchanged from go1.5 to go1.26.
type bufioReader struct {
buf []byte
rd io.Reader // reader provided by the client
@@ -16,8 +17,32 @@ type bufioReader struct {
lastRuneSize int // size of last rune read for UnreadRune; -1 means invalid
}
// Grow increases the read buffer to at least size while preserving buffered data.
// The capacity grows geometrically to avoid repeated allocations for small increments.
func (c *BufferedConn) Grow(size int) {
b := (*bufioReader)(unsafe.Pointer(c.r))
if size <= len(b.buf) {
return
}
newSize := uint(1) << bits.Len(uint(size-1))
if newSize > ^uint(0)>>1 {
newSize = uint(size)
}
newBuf := make([]byte, int(newSize))
buffered := copy(newBuf, b.buf[b.r:b.w])
b.buf = newBuf
b.r = 0
b.w = buffered
}
func (c *BufferedConn) AppendData(buf []byte) (ok bool) {
b := (*bufioReader)(unsafe.Pointer(c.r))
needed := b.w - b.r + len(buf)
if needed > len(b.buf) {
c.Grow(needed)
}
pos := len(b.buf) - b.w - len(buf)
if pos >= -b.r { // len(b.buf)-(b.w - b.r) >= len(buf)
if pos < 0 { // len(b.buf)-b.w < len(buf)
+149
View File
@@ -0,0 +1,149 @@
package net
import (
"bufio"
"bytes"
"io"
stdnet "net"
"reflect"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestBufioReaderLayout(t *testing.T) {
standard := reflect.TypeOf(bufio.Reader{})
mirror := reflect.TypeOf(bufioReader{})
require.Equal(t, standard.Size(), mirror.Size())
require.Equal(t, standard.Align(), mirror.Align())
require.Equal(t, standard.NumField(), mirror.NumField())
for i := 0; i < standard.NumField(); i++ {
standardField := standard.Field(i)
mirrorField := mirror.Field(i)
assert.Equal(t, standardField.Name, mirrorField.Name)
assert.Equal(t, standardField.Type, mirrorField.Type)
assert.Equal(t, standardField.Offset, mirrorField.Offset)
assert.Equal(t, standardField.Anonymous, mirrorField.Anonymous)
}
}
type testReaderConn struct {
stdnet.Conn
*bytes.Reader
}
func (c *testReaderConn) Read(p []byte) (int, error) {
return c.Reader.Read(p)
}
func TestBufferedConnGrow(t *testing.T) {
data := make([]byte, 32*1024)
for i := range data {
data[i] = byte(i)
}
conn := NewBufferedConn(&testReaderConn{Reader: bytes.NewReader(data)})
reader := conn.Reader()
first, err := conn.Peek(1)
require.NoError(t, err)
require.Equal(t, data[:1], first)
require.Less(t, conn.Buffered(), len(data))
conn.Grow(len(data))
peeked, err := conn.Peek(len(data))
require.NoError(t, err)
assert.Equal(t, data, peeked)
assert.Same(t, reader, conn.Reader())
assert.Equal(t, 32*1024, conn.Reader().Size())
_, err = conn.Peek(32*1024 + 1)
assert.ErrorIs(t, err, bufio.ErrBufferFull)
assert.Equal(t, 32*1024, conn.Reader().Size())
read := make([]byte, len(data))
_, err = io.ReadFull(conn, read)
require.NoError(t, err)
assert.Equal(t, data, read)
conn = NewBufferedConn(&testReaderConn{Reader: bytes.NewReader(data)})
initialSize := conn.Reader().Size()
peeked, err = conn.Peek(32*1024 + 1)
assert.ErrorIs(t, err, bufio.ErrBufferFull)
assert.Equal(t, data[:initialSize], peeked)
assert.Equal(t, initialSize, conn.Reader().Size())
}
func TestBufferedConnAppendData(t *testing.T) {
t.Run("prepends to unread underlying data", func(t *testing.T) {
conn := NewBufferedConn(&testReaderConn{Reader: bytes.NewReader([]byte("underlying"))})
require.True(t, conn.AppendData([]byte("appended")))
data, err := io.ReadAll(conn)
require.NoError(t, err)
assert.Equal(t, []byte("appendedunderlying"), data)
})
t.Run("slides buffered data to make room", func(t *testing.T) {
conn := NewBufferedConn(&testReaderConn{Reader: bytes.NewReader(bytes.Repeat([]byte("a"), 4096))})
_, err := conn.Peek(4096)
require.NoError(t, err)
_, err = conn.Discard(2048)
require.NoError(t, err)
require.True(t, conn.AppendData(bytes.Repeat([]byte("b"), 2048)))
data, err := io.ReadAll(conn)
require.NoError(t, err)
assert.Equal(t, append(bytes.Repeat([]byte("a"), 2048), bytes.Repeat([]byte("b"), 2048)...), data)
})
t.Run("grows to make room", func(t *testing.T) {
original := bytes.Repeat([]byte("a"), 4096)
appended := bytes.Repeat([]byte("b"), 2048)
conn := NewBufferedConn(&testReaderConn{Reader: bytes.NewReader(original)})
_, err := conn.Peek(4096)
require.NoError(t, err)
_, err = conn.Discard(1024)
require.NoError(t, err)
require.True(t, conn.AppendData(appended))
assert.Equal(t, 8192, conn.Reader().Size())
data, err := io.ReadAll(conn)
require.NoError(t, err)
assert.Equal(t, append(original[1024:], appended...), data)
})
t.Run("grows beyond 32 KiB", func(t *testing.T) {
original := bytes.Repeat([]byte("a"), 32*1024)
appended := bytes.Repeat([]byte("b"), 1025)
conn := NewBufferedConn(&testReaderConn{Reader: bytes.NewReader(original)})
conn.Grow(len(original))
_, err := conn.Peek(32 * 1024)
require.NoError(t, err)
_, err = conn.Discard(1024)
require.NoError(t, err)
require.True(t, conn.AppendData(appended))
assert.Greater(t, conn.Reader().Size(), 32*1024)
data, err := io.ReadAll(conn)
require.NoError(t, err)
assert.Equal(t, append(original[1024:], appended...), data)
})
t.Run("appends after explicit growth", func(t *testing.T) {
original := bytes.Repeat([]byte("a"), 5000)
appended := bytes.Repeat([]byte("b"), 1000)
conn := NewBufferedConn(&testReaderConn{Reader: bytes.NewReader(original)})
conn.Grow(len(original))
_, err := conn.Peek(len(original))
require.NoError(t, err)
require.True(t, conn.AppendData(appended))
data, err := io.ReadAll(conn)
require.NoError(t, err)
assert.Equal(t, append(original, appended...), data)
})
}
+13
View File
@@ -21,6 +21,11 @@ var (
ErrNoClue = errors.New("not enough information for making a decision")
)
// maxSniffBufferSize bounds the per-connection read-ahead memory used by TCP
// sniffing. 64 KiB covers the HTTP/2 preface and several default-sized frames
// while keeping lengths read from untrusted protocol headers bounded.
const maxSniffBufferSize = 64 * 1024
type Dispatcher struct {
enable bool
sniffers map[sniffer.Sniffer]SnifferConfig
@@ -242,6 +247,14 @@ func (sd *Dispatcher) sniffDomain(conn *N.BufferedConn, metadata *C.Metadata) (s
if need.length <= len(data) || !time.Now().Before(deadline) {
break
}
// Request enough capacity for the next retry. Grow rounds capacity up
// geometrically, while this power-of-two limit keeps automatic allocation
// bounded when a protocol advertises a much larger length.
growTo := need.length
if growTo > maxSniffBufferSize {
growTo = maxSniffBufferSize
}
conn.Grow(growTo)
//log.Debugln("[Sniffer] [%s] [%s] %v, got length: %d, want: %d", metadata.DstIP, s.Protocol(), need, len(data), need.length)
want = need.length
}
+25 -3
View File
@@ -5,6 +5,7 @@ package sniffer
// parsing itself (a complete buffer in, a domain out) belongs in sniff_test.go.
import (
"bytes"
"io"
"net"
"testing"
@@ -95,6 +96,7 @@ func needAtLeast(n int) error {
func TestDispatcherFeedLoop(t *testing.T) {
segment := []byte("0123456789")
threeSegments := [][]byte{segment, segment, segment}
largeSize := 5000
tests := []struct {
name string
@@ -104,7 +106,8 @@ func TestDispatcherFeedLoop(t *testing.T) {
wantErr bool
// size of each buffer handed to the sniffer, so both the number of
// rounds and how much they grew by are pinned down
seen []int
seen []int
bufferSize int
}{
{
// a sniffer that discovers its needs incrementally (HTTP/2: preface,
@@ -122,6 +125,21 @@ func TestDispatcherFeedLoop(t *testing.T) {
host: "example.com",
seen: []int{10, 20, 30},
},
{
// requests beyond bufio.Reader's default capacity should expand the
// buffer before the dispatcher retries Peek
name: "grows the peek buffer on demand",
chunks: [][]byte{segment, bytes.Repeat([]byte("x"), largeSize-len(segment))},
reply: func(data []byte) (string, error) {
if len(data) < largeSize {
return "", needAtLeast(largeSize)
}
return "example.com", nil
},
host: "example.com",
seen: []int{10, largeSize},
bufferSize: 8192,
},
{
// asking for one more byte still advances a whole segment at a time,
// because every retry is handed everything that arrived
@@ -153,8 +171,9 @@ func TestDispatcherFeedLoop(t *testing.T) {
reply: func(data []byte) (string, error) {
return "", needAtLeast(1 << 20)
},
wantErr: true,
seen: []int{10},
wantErr: true,
seen: []int{10},
bufferSize: maxSniffBufferSize,
},
{
name: "gives up when the data never completes",
@@ -186,6 +205,9 @@ func TestDispatcherFeedLoop(t *testing.T) {
assert.Equal(t, test.host, host)
}
assert.Equal(t, test.seen, s.seen)
if test.bufferSize != 0 {
assert.Equal(t, test.bufferSize, conn.Reader().Size())
}
// sniffing must not leave a deadline behind for the relay that follows
assert.True(t, raw.deadline.IsZero(), "read deadline was not cleared")
})