Fix unbounded allocations when reading untrusted binary data

This commit is contained in:
世界
2026-08-10 13:26:42 +08:00
parent 45ca32dcb9
commit f260e771ea
11 changed files with 151 additions and 90 deletions
+7
View File
@@ -7,6 +7,7 @@ import (
"io"
"time"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/observable"
"github.com/sagernet/sing/common/varbin"
)
@@ -103,6 +104,9 @@ func (s *SavedBinary) UnmarshalBinary(data []byte) error {
if err != nil {
return err
}
if contentLength > uint64(reader.Len()) {
return E.New("invalid content length: ", contentLength)
}
s.Content = make([]byte, contentLength)
_, err = io.ReadFull(reader, s.Content)
if err != nil {
@@ -118,6 +122,9 @@ func (s *SavedBinary) UnmarshalBinary(data []byte) error {
if err != nil {
return err
}
if etagLength > uint64(reader.Len()) {
return E.New("invalid etag length: ", etagLength)
}
etagBytes := make([]byte, etagLength)
_, err = io.ReadFull(reader, etagBytes)
if err != nil {
+5 -1
View File
@@ -211,7 +211,11 @@ func TestGeositeWriteReadCompat(t *testing.T) {
for code, expectedItems := range tc.input {
items, err := reader.Read(code)
require.NoError(t, err)
require.Equal(t, expectedItems, items, "items mismatch for code: %s", code)
if len(expectedItems) == 0 {
require.Empty(t, items, "items mismatch for code: %s", code)
} else {
require.Equal(t, expectedItems, items, "items mismatch for code: %s", code)
}
}
})
}
+21 -12
View File
@@ -4,6 +4,7 @@ import (
"bufio"
"encoding/binary"
"io"
"math"
"os"
"sync"
"sync/atomic"
@@ -62,10 +63,9 @@ func (r *Reader) readMetadata() error {
if err != nil {
return err
}
keys := make([]string, entryLength)
domainIndex := make(map[string]int)
domainLength := make(map[string]int)
for i := 0; i < int(entryLength); i++ {
for range entryLength {
var (
code string
codeIndex uint64
@@ -75,7 +75,6 @@ func (r *Reader) readMetadata() error {
if err != nil {
return err
}
keys[i] = code
codeIndex, err = binary.ReadUvarint(reader)
if err != nil {
return err
@@ -84,6 +83,9 @@ func (r *Reader) readMetadata() error {
if err != nil {
return err
}
if codeIndex > math.MaxInt32 || codeLength > math.MaxInt32 {
return E.New("invalid metadata entry: ", code)
}
domainIndex[code] = int(codeIndex)
domainLength[code] = int(codeLength)
}
@@ -107,17 +109,22 @@ func (r *Reader) Read(code string) ([]Item, error) {
return nil, err
}
r.bufferedReader.Reset(r.reader)
itemList := make([]Item, r.domainLength[code])
for i := range itemList {
typeByte, err := r.bufferedReader.ReadByte()
length := r.domainLength[code]
var itemList []Item
for range length {
var (
typeByte byte
value string
)
typeByte, err = r.bufferedReader.ReadByte()
if err != nil {
return nil, err
}
itemList[i].Type = ItemType(typeByte)
itemList[i].Value, err = readString(r.bufferedReader)
value, err = readString(r.bufferedReader)
if err != nil {
return nil, err
}
itemList = append(itemList, Item{Type: ItemType(typeByte), Value: value})
}
return itemList, nil
}
@@ -144,12 +151,14 @@ func readString(reader io.ByteReader) (string, error) {
if err != nil {
return "", err
}
bytes := make([]byte, length)
for i := range bytes {
bytes[i], err = reader.ReadByte()
var result []byte
for range length {
var value byte
value, err = reader.ReadByte()
if err != nil {
return "", err
}
result = append(result, value)
}
return string(bytes), nil
return string(result), nil
}
+25 -47
View File
@@ -77,13 +77,14 @@ func Read(reader io.Reader, recover bool) (ruleSetCompat option.PlainRuleSetComp
return
}
ruleSetCompat.Version = version
ruleSetCompat.Options.Rules = make([]option.HeadlessRule, length)
for i := range length {
ruleSetCompat.Options.Rules[i], err = readRule(bReader, recover)
var rule option.HeadlessRule
rule, err = readRule(bReader, recover, 0)
if err != nil {
err = E.Cause(err, "read rule[", i, "]")
return
}
ruleSetCompat.Options.Rules = append(ruleSetCompat.Options.Rules, rule)
}
return
}
@@ -119,7 +120,13 @@ func Write(writer io.Writer, ruleSet option.PlainRuleSet, generateVersion uint8)
return compressWriter.Close()
}
func readRule(reader varbin.Reader, recover bool) (rule option.HeadlessRule, err error) {
const maxLogicalRuleDepth = 100
func readRule(reader varbin.Reader, recover bool, depth int) (rule option.HeadlessRule, err error) {
if depth > maxLogicalRuleDepth {
err = E.New("logical rule nested too deep")
return
}
var ruleType uint8
err = binary.Read(reader, binary.BigEndian, &ruleType)
if err != nil {
@@ -131,7 +138,7 @@ func readRule(reader varbin.Reader, recover bool) (rule option.HeadlessRule, err
rule.DefaultOptions, err = readDefaultRule(reader, recover)
case 1:
rule.Type = C.RuleTypeLogical
rule.LogicalOptions, err = readLogicalRule(reader, recover)
rule.LogicalOptions, err = readLogicalRule(reader, recover, depth)
default:
err = E.New("unknown rule type: ", ruleType)
}
@@ -160,7 +167,7 @@ func readDefaultRule(reader varbin.Reader, recover bool) (rule option.DefaultHea
switch itemType {
case ruleItemQueryType:
var rawQueryType []uint16
rawQueryType, err = readRuleItemUint16(reader)
rawQueryType, err = varbin.ReadSlice[uint16](reader, binary.BigEndian)
if err != nil {
return
}
@@ -200,11 +207,11 @@ func readDefaultRule(reader varbin.Reader, recover bool) (rule option.DefaultHea
rule.IPCIDR = common.Map(rule.IPSet.Prefixes(), netip.Prefix.String)
}
case ruleItemSourcePort:
rule.SourcePort, err = readRuleItemUint16(reader)
rule.SourcePort, err = varbin.ReadSlice[uint16](reader, binary.BigEndian)
case ruleItemSourcePortRange:
rule.SourcePortRange, err = readRuleItemString(reader)
case ruleItemPort:
rule.Port, err = readRuleItemUint16(reader)
rule.Port, err = varbin.ReadSlice[uint16](reader, binary.BigEndian)
case ruleItemPortRange:
rule.PortRange, err = readRuleItemString(reader)
case ruleItemProcessName:
@@ -230,7 +237,7 @@ func readDefaultRule(reader varbin.Reader, recover bool) (rule option.DefaultHea
rule.AdGuardDomain = matcher.Dump()
}
case ruleItemNetworkType:
rule.NetworkType, err = readRuleItemUint8[option.InterfaceType](reader)
rule.NetworkType, err = varbin.ReadSlice[option.InterfaceType](reader, binary.BigEndian)
case ruleItemNetworkIsExpensive:
rule.NetworkIsExpensive = true
case ruleItemNetworkIsConstrained:
@@ -242,7 +249,7 @@ func readDefaultRule(reader varbin.Reader, recover bool) (rule option.DefaultHea
if err != nil {
return
}
for i := uint64(0); i < size; i++ {
for range size {
var key uint8
err = binary.Read(reader, binary.BigEndian, &key)
if err != nil {
@@ -510,18 +517,14 @@ func readRuleItemString(reader varbin.Reader) ([]string, error) {
if err != nil {
return nil, err
}
result := make([]string, length)
for i := range result {
strLen, err := binary.ReadUvarint(reader)
var result []string
for range length {
var value []byte
value, err = varbin.ReadSlice[byte](reader, binary.BigEndian)
if err != nil {
return nil, err
}
buf := make([]byte, strLen)
_, err = io.ReadFull(reader, buf)
if err != nil {
return nil, err
}
result[i] = string(buf)
result = append(result, string(value))
}
return result, nil
}
@@ -548,19 +551,6 @@ func writeRuleItemString(writer varbin.Writer, itemType uint8, value []string) e
return nil
}
func readRuleItemUint8[E ~uint8](reader varbin.Reader) ([]E, error) {
length, err := binary.ReadUvarint(reader)
if err != nil {
return nil, err
}
result := make([]E, length)
_, err = io.ReadFull(reader, *(*[]byte)(unsafe.Pointer(&result)))
if err != nil {
return nil, err
}
return result, nil
}
func writeRuleItemUint8[E ~uint8](writer varbin.Writer, itemType uint8, value []E) error {
err := writer.WriteByte(itemType)
if err != nil {
@@ -574,19 +564,6 @@ func writeRuleItemUint8[E ~uint8](writer varbin.Writer, itemType uint8, value []
return err
}
func readRuleItemUint16(reader varbin.Reader) ([]uint16, error) {
length, err := binary.ReadUvarint(reader)
if err != nil {
return nil, err
}
result := make([]uint16, length)
err = binary.Read(reader, binary.BigEndian, result)
if err != nil {
return nil, err
}
return result, nil
}
func writeRuleItemUint16(writer varbin.Writer, itemType uint8, value []uint16) error {
err := writer.WriteByte(itemType)
if err != nil {
@@ -625,7 +602,7 @@ func writeRuleItemCIDR(writer varbin.Writer, itemType uint8, value []string) err
return writeIPSet(writer, ipSet)
}
func readLogicalRule(reader varbin.Reader, recovery bool) (logicalRule option.LogicalHeadlessRule, err error) {
func readLogicalRule(reader varbin.Reader, recovery bool, depth int) (logicalRule option.LogicalHeadlessRule, err error) {
mode, err := reader.ReadByte()
if err != nil {
return
@@ -643,13 +620,14 @@ func readLogicalRule(reader varbin.Reader, recovery bool) (logicalRule option.Lo
if err != nil {
return
}
logicalRule.Rules = make([]option.HeadlessRule, length)
for i := range length {
logicalRule.Rules[i], err = readRule(reader, recovery)
var rule option.HeadlessRule
rule, err = readRule(reader, recovery, depth+1)
if err != nil {
err = E.Cause(err, "read logical rule [", i, "]")
return
}
logicalRule.Rules = append(logicalRule.Rules, rule)
}
err = binary.Read(reader, binary.BigEndian, &logicalRule.Invert)
if err != nil {
+2 -2
View File
@@ -251,7 +251,7 @@ func TestUint8SliceCompat(t *testing.T) {
requireUint8SliceEqual(t, tc.input, readBack)
// Old write -> new read
readBack2, err := readRuleItemUint8[uint8](bufio.NewReader(bytes.NewReader(oldBuf.Bytes())))
readBack2, err := varbin.ReadSlice[uint8](bufio.NewReader(bytes.NewReader(oldBuf.Bytes())), binary.BigEndian)
require.NoError(t, err)
requireUint8SliceEqual(t, tc.input, readBack2)
})
@@ -300,7 +300,7 @@ func TestUint16SliceCompat(t *testing.T) {
requireUint16SliceEqual(t, tc.input, readBack)
// Old write -> new read
readBack2, err := readRuleItemUint16(bufio.NewReader(bytes.NewReader(oldBuf.Bytes())))
readBack2, err := varbin.ReadSlice[uint16](bufio.NewReader(bytes.NewReader(oldBuf.Bytes())), binary.BigEndian)
require.NoError(t, err)
requireUint16SliceEqual(t, tc.input, readBack2)
})
+7 -3
View File
@@ -4,6 +4,7 @@ import (
"encoding/binary"
"io"
"net/netip"
"os"
M "github.com/sagernet/sing/common/metadata"
"github.com/sagernet/sing/common/varbin"
@@ -14,8 +15,11 @@ func readPrefix(reader varbin.Reader) (netip.Prefix, error) {
if err != nil {
return netip.Prefix{}, err
}
addrSlice := make([]byte, addrLen)
_, err = io.ReadFull(reader, addrSlice)
if addrLen != 4 && addrLen != 16 {
return netip.Prefix{}, os.ErrInvalid
}
var addrBytes [16]byte
_, err = io.ReadFull(reader, addrBytes[:addrLen])
if err != nil {
return netip.Prefix{}, err
}
@@ -23,7 +27,7 @@ func readPrefix(reader varbin.Reader) (netip.Prefix, error) {
if err != nil {
return netip.Prefix{}, err
}
return netip.PrefixFrom(M.AddrFromIP(addrSlice), int(prefixBits)), nil
return netip.PrefixFrom(M.AddrFromIP(addrBytes[:addrLen]), int(prefixBits)), nil
}
func writePrefix(writer varbin.Writer, prefix netip.Prefix) error {
+22 -18
View File
@@ -36,34 +36,38 @@ func readIPSet(reader varbin.Reader) (*netipx.IPSet, error) {
if err != nil {
return nil, err
}
mySet := &myIPSet{
rr: make([]myIPRange, length),
}
for i := range mySet.rr {
fromLen, err := binary.ReadUvarint(reader)
mySet := &myIPSet{}
for range length {
var from, to netip.Addr
from, err = readIPSetAddr(reader)
if err != nil {
return nil, err
}
fromBytes := make([]byte, fromLen)
_, err = io.ReadFull(reader, fromBytes)
to, err = readIPSetAddr(reader)
if err != nil {
return nil, err
}
toLen, err := binary.ReadUvarint(reader)
if err != nil {
return nil, err
}
toBytes := make([]byte, toLen)
_, err = io.ReadFull(reader, toBytes)
if err != nil {
return nil, err
}
mySet.rr[i].from = M.AddrFromIP(fromBytes)
mySet.rr[i].to = M.AddrFromIP(toBytes)
mySet.rr = append(mySet.rr, myIPRange{from: from, to: to})
}
return (*netipx.IPSet)(unsafe.Pointer(mySet)), nil
}
func readIPSetAddr(reader varbin.Reader) (netip.Addr, error) {
addrLen, err := binary.ReadUvarint(reader)
if err != nil {
return netip.Addr{}, err
}
if addrLen != 4 && addrLen != 16 {
return netip.Addr{}, os.ErrInvalid
}
var addrBytes [16]byte
_, err = io.ReadFull(reader, addrBytes[:addrLen])
if err != nil {
return netip.Addr{}, err
}
return M.AddrFromIP(addrBytes[:addrLen]), nil
}
func writeIPSet(writer varbin.Writer, set *netipx.IPSet) error {
err := writer.WriteByte(1)
if err != nil {
+53
View File
@@ -0,0 +1,53 @@
package srs
import (
"bytes"
"compress/zlib"
"encoding/binary"
"testing"
C "github.com/sagernet/sing-box/constant"
"github.com/stretchr/testify/require"
)
func craftRuleSet(body []byte) []byte {
var buffer bytes.Buffer
buffer.Write(MagicBytes[:])
buffer.WriteByte(C.RuleSetVersionCurrent)
compressWriter := zlib.NewWriter(&buffer)
compressWriter.Write(body)
compressWriter.Close()
return buffer.Bytes()
}
func uvarint(value uint64) []byte {
return binary.AppendUvarint(nil, value)
}
func TestReadMalformed(t *testing.T) {
t.Parallel()
defaultRule := []byte{0x00}
cases := []struct {
name string
body []byte
}{
{"huge_rule_count", uvarint(1 << 40)},
{"huge_string_count", bytes.Join([][]byte{uvarint(1), defaultRule, {ruleItemDomainKeyword}, uvarint(1 << 40)}, nil)},
{"huge_string_length", bytes.Join([][]byte{uvarint(1), defaultRule, {ruleItemDomainKeyword}, uvarint(1), uvarint(1 << 40)}, nil)},
{"huge_uint16_count", bytes.Join([][]byte{uvarint(1), defaultRule, {ruleItemQueryType}, uvarint(1 << 40)}, nil)},
{"huge_ip_range_count", bytes.Join([][]byte{uvarint(1), defaultRule, {ruleItemIPCIDR, 0x01}, binary.BigEndian.AppendUint64(nil, 1<<40)}, nil)},
{"bad_ip_range_address_length", bytes.Join([][]byte{uvarint(1), defaultRule, {ruleItemIPCIDR, 0x01}, binary.BigEndian.AppendUint64(nil, 1), uvarint(99)}, nil)},
{"huge_prefix_address_length", bytes.Join([][]byte{uvarint(1), defaultRule, {ruleItemDefaultInterfaceAddress}, uvarint(1), uvarint(1 << 40)}, nil)},
{"empty_domain_matcher", bytes.Join([][]byte{uvarint(1), defaultRule, {ruleItemDomain, 0x00}, uvarint(0), uvarint(0), uvarint(0)}, nil)},
{"huge_domain_matcher_bitmap", bytes.Join([][]byte{uvarint(1), defaultRule, {ruleItemDomain, 0x00}, uvarint(1 << 40)}, nil)},
{"deep_logical_nesting", bytes.Join([][]byte{uvarint(1), bytes.Repeat([]byte{0x01, 0x00, 0x01}, 10000)}, nil)},
}
for _, testCase := range cases {
t.Run(testCase.name, func(t *testing.T) {
t.Parallel()
_, err := Read(bytes.NewReader(craftRuleSet(testCase.body)), false)
require.Error(t, err)
})
}
}
+6 -4
View File
@@ -257,14 +257,16 @@ func readString(reader io.ByteReader) (string, error) {
if err != nil {
return "", err
}
buf := make([]byte, length)
for i := range buf {
buf[i], err = reader.ReadByte()
var result []byte
for range length {
var value byte
value, err = reader.ReadByte()
if err != nil {
return "", err
}
result = append(result, value)
}
return string(buf), nil
return string(result), nil
}
func writeString(buffer *bytes.Buffer, value string) {
+1 -1
View File
@@ -34,7 +34,7 @@ require (
github.com/sagernet/gomobile v0.1.12
github.com/sagernet/gvisor v0.0.0-20250811.0-sing-box-mod.1
github.com/sagernet/quic-go v0.59.0-sing-box-mod.4
github.com/sagernet/sing v0.8.12
github.com/sagernet/sing v0.8.13
github.com/sagernet/sing-mux v0.3.5
github.com/sagernet/sing-quic v0.6.4-0.20260803041914-d83826c306d7
github.com/sagernet/sing-shadowsocks v0.2.8
+2 -2
View File
@@ -236,8 +236,8 @@ github.com/sagernet/nftables v0.3.0-mod.2 h1:ck2KMU02OxL1eDFgGaWYglMDpoOZ7OHzxje
github.com/sagernet/nftables v0.3.0-mod.2/go.mod h1:8kslHG4VvYNihcco+i6uxIX7qbT8A56T0y5q7U44ZaQ=
github.com/sagernet/quic-go v0.59.0-sing-box-mod.4 h1:6qvrUW79S+CrPwWz6cMePXohgjHoKxLo3c+MDhNwc3o=
github.com/sagernet/quic-go v0.59.0-sing-box-mod.4/go.mod h1:OqILvS182CyOol5zNNo6bguvOGgXzV459+chpRaUC+4=
github.com/sagernet/sing v0.8.12 h1:v77YM0gB0uCQ74f1RmiTDtkV8eNT6Qtniobxh012a5I=
github.com/sagernet/sing v0.8.12/go.mod h1:olXxWQNqRW/l2Q6JI3b2Qmz8iQnIFlOeeH8bx6JhgUA=
github.com/sagernet/sing v0.8.13 h1:yVoXnx9nPxfjlwD4Tp+Wd9zuW2tfiSVrcRDBZNbKRCw=
github.com/sagernet/sing v0.8.13/go.mod h1:olXxWQNqRW/l2Q6JI3b2Qmz8iQnIFlOeeH8bx6JhgUA=
github.com/sagernet/sing-mux v0.3.5 h1:RHnhVEc+SFqkrK4xMygYjDwwLhzp2Bj3lztSukONfhI=
github.com/sagernet/sing-mux v0.3.5/go.mod h1:QvlKMyNBNrQoyX4x+gq028uPbLM2XeRpWtDsWBJbFSk=
github.com/sagernet/sing-quic v0.6.4-0.20260803041914-d83826c306d7 h1:D46kmyvKMNVFvL3KXdq3T4vy8r29xCMmSqNuQHJBybE=