diff --git a/common/net/bufconn_unsafe.go b/common/net/bufconn_unsafe.go index 349321df..fc988432 100644 --- a/common/net/bufconn_unsafe.go +++ b/common/net/bufconn_unsafe.go @@ -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) diff --git a/common/net/bufconn_unsafe_test.go b/common/net/bufconn_unsafe_test.go new file mode 100644 index 00000000..c905c4e1 --- /dev/null +++ b/common/net/bufconn_unsafe_test.go @@ -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) + }) +} diff --git a/component/sniffer/dispatcher.go b/component/sniffer/dispatcher.go index c7710055..06e3eb5a 100644 --- a/component/sniffer/dispatcher.go +++ b/component/sniffer/dispatcher.go @@ -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 } diff --git a/component/sniffer/dispatcher_test.go b/component/sniffer/dispatcher_test.go index 399c2960..8d734062 100644 --- a/component/sniffer/dispatcher_test.go +++ b/component/sniffer/dispatcher_test.go @@ -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") })