diff --git a/adapter/experimental.go b/adapter/experimental.go index 1bd8d2d9..67191141 100644 --- a/adapter/experimental.go +++ b/adapter/experimental.go @@ -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 { diff --git a/common/geosite/compat_test.go b/common/geosite/compat_test.go index 9c66aea3..6989a5fd 100644 --- a/common/geosite/compat_test.go +++ b/common/geosite/compat_test.go @@ -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) + } } }) } diff --git a/common/geosite/reader.go b/common/geosite/reader.go index ecd63a7e..0d9ca8f3 100644 --- a/common/geosite/reader.go +++ b/common/geosite/reader.go @@ -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 } diff --git a/common/srs/binary.go b/common/srs/binary.go index d2c865e1..95b054b2 100644 --- a/common/srs/binary.go +++ b/common/srs/binary.go @@ -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 { diff --git a/common/srs/compat_test.go b/common/srs/compat_test.go index 46f3c114..6300c5f1 100644 --- a/common/srs/compat_test.go +++ b/common/srs/compat_test.go @@ -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) }) diff --git a/common/srs/ip_cidr.go b/common/srs/ip_cidr.go index 7c81abda..cba22f88 100644 --- a/common/srs/ip_cidr.go +++ b/common/srs/ip_cidr.go @@ -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 { diff --git a/common/srs/ip_set.go b/common/srs/ip_set.go index a10ac08c..01def6bb 100644 --- a/common/srs/ip_set.go +++ b/common/srs/ip_set.go @@ -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 { diff --git a/common/srs/malformed_test.go b/common/srs/malformed_test.go new file mode 100644 index 00000000..eff932fb --- /dev/null +++ b/common/srs/malformed_test.go @@ -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) + }) + } +} diff --git a/experimental/libbox/profile_import.go b/experimental/libbox/profile_import.go index c337d015..c0906f87 100644 --- a/experimental/libbox/profile_import.go +++ b/experimental/libbox/profile_import.go @@ -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) { diff --git a/go.mod b/go.mod index 7fce6138..ad40dff3 100644 --- a/go.mod +++ b/go.mod @@ -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 diff --git a/go.sum b/go.sum index 4b2cdf69..be75c2a1 100644 --- a/go.sum +++ b/go.sum @@ -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=