mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-09-30 19:07:58 +03:00
Geodata: Reduce matcher memory on mobile and desktop, skip regexes that cannot match (#6867)
https://github.com/XTLS/Xray-core/pull/6867#issuecomment-5895171934
This commit is contained in:
@@ -82,19 +82,10 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
|
||||
}
|
||||
g.Add(m, uint32(i))
|
||||
case *DomainRule_Geosite:
|
||||
domains, err := loadSiteWithAttrs(v.Geosite.File, v.Geosite.Code, v.Geosite.Attrs)
|
||||
err := loadSiteMatchers(v.Geosite, func(m strmatcher.Matcher) { g.Add(m, uint32(i)) })
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for j, d := range domains {
|
||||
domains[j] = nil // peak mem
|
||||
m, err := parseDomain(d)
|
||||
if err != nil {
|
||||
errors.LogError(context.Background(), "ignore invalid geosite entry in ", v.Geosite.File, ":", v.Geosite.Code, " at index ", j, ", ", err)
|
||||
continue
|
||||
}
|
||||
g.Add(m, uint32(i))
|
||||
}
|
||||
default:
|
||||
panic("unknown domain rule type")
|
||||
}
|
||||
@@ -108,12 +99,12 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
|
||||
return g, nil
|
||||
}
|
||||
|
||||
type CompactDomainMatcherFactory struct {
|
||||
type CompactMphDomainMatcherFactory struct {
|
||||
sync.Mutex
|
||||
shared *utils.WeakCacheMap[string, strmatcher.LinearAnyMatcher]
|
||||
shared *utils.WeakCacheMap[string, strmatcher.MphValueMatcher]
|
||||
}
|
||||
|
||||
func (f *CompactDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (strmatcher.MatcherSet, error) {
|
||||
func (f *CompactMphDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (*strmatcher.MphValueMatcher, error) {
|
||||
key := rule.File + ":" + rule.Code + "@" + rule.Attrs
|
||||
|
||||
f.Lock()
|
||||
@@ -125,33 +116,23 @@ func (f *CompactDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (strmat
|
||||
}
|
||||
errors.LogDebug(context.Background(), "geodata geosite matcher cache MISS ", key)
|
||||
|
||||
s := strmatcher.NewLinearAnyMatcher()
|
||||
domains, err := loadSiteWithAttrs(rule.File, rule.Code, rule.Attrs)
|
||||
if err != nil {
|
||||
s := strmatcher.NewMphValueMatcher()
|
||||
if err := loadSiteMatchers(rule, func(m strmatcher.Matcher) { s.Add(m, 0) }); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for i, d := range domains {
|
||||
domains[i] = nil // peak mem
|
||||
m, err := parseDomain(d)
|
||||
if err != nil {
|
||||
errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err)
|
||||
continue
|
||||
}
|
||||
s.Add(m)
|
||||
if err := s.Build(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
f.shared.Store(key, s)
|
||||
return s, err
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// BuildMatcher implements DomainMatcherFactory.
|
||||
func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
|
||||
func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
|
||||
if len(rules) == 0 {
|
||||
return nil, errors.New("empty domain rule list")
|
||||
}
|
||||
compact := &CompactDomainMatcher{
|
||||
matchers: make([]strmatcher.MatcherSet, 0, len(rules)),
|
||||
values: make([]uint32, 0, len(rules)),
|
||||
}
|
||||
compact := new(CompactMphDomainMatcher)
|
||||
for i, r := range rules {
|
||||
switch v := r.Value.(type) {
|
||||
case *DomainRule_Custom:
|
||||
@@ -168,8 +149,7 @@ func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainM
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
compact.matchers = append(compact.matchers, m)
|
||||
compact.values = append(compact.values, uint32(i))
|
||||
compact.combiner.Add(m, uint32(i))
|
||||
default:
|
||||
panic("unknown domain rule type")
|
||||
}
|
||||
@@ -177,37 +157,40 @@ func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainM
|
||||
return compact, nil
|
||||
}
|
||||
|
||||
type CompactDomainMatcher struct {
|
||||
type CompactMphDomainMatcher struct {
|
||||
custom strmatcher.ValueMatcher
|
||||
matchers []strmatcher.MatcherSet
|
||||
values []uint32
|
||||
combiner strmatcher.MphValueMatcherCombiner
|
||||
}
|
||||
|
||||
// Match implements DomainMatcher.
|
||||
func (c *CompactDomainMatcher) Match(input string) []uint32 {
|
||||
var result []uint32
|
||||
func (c *CompactMphDomainMatcher) Match(input string) []uint32 {
|
||||
result := c.combiner.Match(input)
|
||||
if c.custom != nil {
|
||||
result = append(result, c.custom.Match(input)...)
|
||||
}
|
||||
for i, m := range c.matchers {
|
||||
if m.MatchAny(input) {
|
||||
result = append(result, c.values[i])
|
||||
}
|
||||
result = append(c.custom.Match(input), result...)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// MatchAny implements DomainMatcher.
|
||||
func (c *CompactDomainMatcher) MatchAny(input string) bool {
|
||||
func (c *CompactMphDomainMatcher) MatchAny(input string) bool {
|
||||
if c.custom != nil && c.custom.MatchAny(input) {
|
||||
return true
|
||||
}
|
||||
for _, m := range c.matchers {
|
||||
if m.MatchAny(input) {
|
||||
return true
|
||||
return c.combiner.MatchAny(input)
|
||||
}
|
||||
|
||||
// loadSiteMatchers calls add with a matcher for every domain of the geosite rule and logs the invalid ones.
|
||||
func loadSiteMatchers(rule *GeoSiteRule, add func(strmatcher.Matcher)) error {
|
||||
i := 0
|
||||
return loadSite(rule.File, rule.Code, rule.Attrs, func(t Domain_Type, value []byte) {
|
||||
m, err := parseDomain(&Domain{Type: t, Value: string(value)})
|
||||
if err != nil {
|
||||
errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err)
|
||||
} else {
|
||||
add(m)
|
||||
}
|
||||
}
|
||||
return false
|
||||
i++
|
||||
})
|
||||
}
|
||||
|
||||
func parseDomain(d *Domain) (strmatcher.Matcher, error) {
|
||||
@@ -231,7 +214,7 @@ func parseDomain(d *Domain) (strmatcher.Matcher, error) {
|
||||
func newDomainMatcherFactory() DomainMatcherFactory {
|
||||
switch runtime.GOOS {
|
||||
case "ios", "android":
|
||||
return &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
|
||||
return &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||
default:
|
||||
return &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"slices"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/xtls/xray-core/common/geodata/strmatcher"
|
||||
@@ -11,7 +12,7 @@ import (
|
||||
)
|
||||
|
||||
func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
|
||||
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
|
||||
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||
matcher, err := factory.BuildMatcher([]*DomainRule{
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
|
||||
@@ -32,7 +33,7 @@ func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
|
||||
func TestCompactDomainMatcher_PreservesMixedRuleIndices(t *testing.T) {
|
||||
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
||||
|
||||
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
|
||||
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||
matcher, err := factory.BuildMatcher([]*DomainRule{
|
||||
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "163.com"}}},
|
||||
@@ -72,3 +73,76 @@ func TestMphDomainMatcher_MatchReturnsDetachedSlice(t *testing.T) {
|
||||
t.Fatalf("Match() after caller mutation = %v, want %v", gotAgain, []uint32{0, 1})
|
||||
}
|
||||
}
|
||||
|
||||
// DNS sorts every Match result in place, so a matcher must never hand out a
|
||||
// slice it keeps, also when only its keyword or regex part matches.
|
||||
func TestDomainMatcher_MatchResultsCanBeSortedConcurrently(t *testing.T) {
|
||||
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
||||
|
||||
rules := []*DomainRule{
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Substr, Value: "exam"}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Regex, Value: `^ex.*\.org$`}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Substr, Value: "exam"}}},
|
||||
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "only.full.test"}}},
|
||||
}
|
||||
cases := []struct {
|
||||
input string
|
||||
want []uint32
|
||||
}{
|
||||
{"example.com", []uint32{0, 1, 2, 4}},
|
||||
{"www.example.com", []uint32{1, 2, 4}},
|
||||
{"exam.net", []uint32{2, 4}}, // keyword part only
|
||||
{"example.org", []uint32{2, 3, 4}},
|
||||
{"163.com", []uint32{5}},
|
||||
{"www.163.com", []uint32{5}},
|
||||
{"only.full.test", []uint32{6}}, // full part only
|
||||
{"nomatch.test", nil},
|
||||
}
|
||||
factories := map[string]DomainMatcherFactory{
|
||||
"mph": &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()},
|
||||
"compact": &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()},
|
||||
}
|
||||
for name, factory := range factories {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
matcher, err := factory.BuildMatcher(rules)
|
||||
if err != nil {
|
||||
t.Fatalf("BuildMatcher() failed: %v", err)
|
||||
}
|
||||
for _, c := range cases {
|
||||
got := matcher.Match(c.input)
|
||||
if sorted := slices.Sorted(slices.Values(got)); !slices.Equal(sorted, c.want) {
|
||||
t.Fatalf("Match(%q) = %v, want %v", c.input, sorted, c.want)
|
||||
}
|
||||
got = got[:cap(got)]
|
||||
for j := range got {
|
||||
got[j] = ^uint32(0)
|
||||
}
|
||||
if again := slices.Sorted(slices.Values(matcher.Match(c.input))); !slices.Equal(again, c.want) {
|
||||
t.Fatalf("Match(%q) after caller mutation = %v, want %v", c.input, again, c.want)
|
||||
}
|
||||
}
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for range 8 {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for range 500 {
|
||||
for _, c := range cases {
|
||||
got := matcher.Match(c.input)
|
||||
slices.Sort(got)
|
||||
if !slices.Equal(got, c.want) {
|
||||
t.Errorf("Match(%q) = %v, want %v", c.input, got, c.want)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
+212
-61
@@ -5,11 +5,14 @@ import (
|
||||
"bytes"
|
||||
"io"
|
||||
"runtime"
|
||||
"slices"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/platform/filesystem"
|
||||
|
||||
"google.golang.org/protobuf/encoding/protowire"
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
@@ -52,17 +55,56 @@ func loadIP(file, code string) ([]*CIDR, error) {
|
||||
return geoip.Cidr, nil
|
||||
}
|
||||
|
||||
func loadSite(file, code string) ([]*Domain, error) {
|
||||
bs, err := loadFile(file, code)
|
||||
// loadSite calls fn, in file order, with the type and value of every domain of the geosite code
|
||||
// that has all the "@"-separated attrs. It decodes the entry while reading the file instead of
|
||||
// unmarshalling it into a []*Domain, so value is only valid during fn.
|
||||
func loadSite(file, code, attrs string, fn func(Domain_Type, []byte)) error {
|
||||
runtime.GC() // peak mem
|
||||
r, err := filesystem.OpenAsset(file)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return errors.New("failed to open ", file).Base(err)
|
||||
}
|
||||
defer runtime.GC() // peak mem
|
||||
var geosite GeoSite
|
||||
if err := proto.Unmarshal(bs, &geosite); err != nil {
|
||||
return nil, errors.New("error unmarshal Site in ", file, ":", code).Base(err)
|
||||
defer r.Close()
|
||||
br := bufio.NewReaderSize(r, 64*1024)
|
||||
n, err := seek(br, []byte(code))
|
||||
if err != nil {
|
||||
return errors.New("failed to load code ", code, " from ", file).Base(err)
|
||||
}
|
||||
return geosite.Domain, nil
|
||||
loadErr := func(err error) error {
|
||||
if err == io.EOF {
|
||||
err = io.ErrUnexpectedEOF
|
||||
}
|
||||
return errors.New("failed to load code ", code, " from ", file).Base(err)
|
||||
}
|
||||
unmarshalErr := func(err error) error {
|
||||
return errors.New("error unmarshal Site in ", file, ":", code).Base(err)
|
||||
}
|
||||
d := newSiteDecoder(attrs, fn)
|
||||
for n > 0 {
|
||||
w, err := br.Peek(min(n, br.Size()))
|
||||
if err != nil {
|
||||
return loadErr(err)
|
||||
}
|
||||
used, err := d.decode(w, len(w) < n)
|
||||
if err != nil {
|
||||
return unmarshalErr(err)
|
||||
}
|
||||
if used == 0 {
|
||||
break // a field longer than the buffer
|
||||
}
|
||||
br.Discard(used)
|
||||
n -= used
|
||||
}
|
||||
if n > 0 {
|
||||
w := make([]byte, n)
|
||||
if _, err := io.ReadFull(br, w); err != nil {
|
||||
return loadErr(err)
|
||||
}
|
||||
if _, err := d.decode(w, false); err != nil {
|
||||
return unmarshalErr(err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func decodeVarint(br *bufio.Reader) (uint64, error) {
|
||||
@@ -82,68 +124,63 @@ func decodeVarint(br *bufio.Reader) (uint64, error) {
|
||||
}
|
||||
|
||||
func find(r io.Reader, code []byte, readBody bool) ([]byte, error) {
|
||||
br := bufio.NewReaderSize(r, 64*1024)
|
||||
bodyL, err := seek(br, code)
|
||||
if err != nil || !readBody {
|
||||
return nil, err
|
||||
}
|
||||
out := make([]byte, bodyL)
|
||||
if _, err := io.ReadFull(br, out); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// seek advances br to the body of the entry for code and returns the body length.
|
||||
func seek(br *bufio.Reader, code []byte) (int, error) {
|
||||
codeL := len(code)
|
||||
if codeL == 0 {
|
||||
return nil, errors.New("empty code")
|
||||
return 0, errors.New("empty code")
|
||||
}
|
||||
|
||||
br := bufio.NewReaderSize(r, 64*1024)
|
||||
need := 2 + codeL // TODO: if code too long
|
||||
prefixBuf := make([]byte, need)
|
||||
|
||||
for {
|
||||
if _, err := br.ReadByte(); err != nil {
|
||||
return nil, err
|
||||
return 0, err
|
||||
}
|
||||
|
||||
x, err := decodeVarint(br)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return 0, err
|
||||
}
|
||||
bodyL := int(x)
|
||||
if bodyL <= 0 {
|
||||
return nil, errors.New("invalid body length: ", bodyL)
|
||||
return 0, errors.New("invalid body length: ", bodyL)
|
||||
}
|
||||
|
||||
prefixL := bodyL
|
||||
if prefixL > need {
|
||||
prefixL = need
|
||||
}
|
||||
prefix := prefixBuf[:prefixL]
|
||||
if _, err := io.ReadFull(br, prefix); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
match := false
|
||||
if bodyL >= need {
|
||||
if int(prefix[1]) == codeL && bytes.Equal(prefix[2:need], code) {
|
||||
if !readBody {
|
||||
return nil, nil
|
||||
}
|
||||
match = true
|
||||
// Peek no more than the buffer holds: a code longer than the buffer cannot match a single
|
||||
// length byte anyway, so a short peek only skips it, as base find (io.ReadFull) does.
|
||||
prefix, err := br.Peek(min(bodyL, need, br.Size()))
|
||||
if err != nil {
|
||||
if err == io.EOF && len(prefix) > 0 {
|
||||
err = io.ErrUnexpectedEOF // as io.ReadFull
|
||||
}
|
||||
return 0, err
|
||||
}
|
||||
|
||||
remain := bodyL - prefixL
|
||||
if match {
|
||||
out := make([]byte, bodyL)
|
||||
copy(out, prefix)
|
||||
if remain > 0 {
|
||||
if _, err := io.ReadFull(br, out[prefixL:]); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
if bodyL >= need && len(prefix) >= need && int(prefix[1]) == codeL && bytes.Equal(prefix[2:], code) {
|
||||
return bodyL, nil
|
||||
}
|
||||
|
||||
if remain > 0 {
|
||||
if _, err := br.Discard(remain); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if _, err := br.Discard(bodyL); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// AttributeMatcher, HasAttrMatcher, AllAttrsMatcher and NewAllAttrsMatcher are the exported
|
||||
// attribute helpers that have been part of this package's API since #5814. The streaming loader
|
||||
// above filters attributes itself without building a *Domain, so it does not use them, but they
|
||||
// are kept for external callers. Their behaviour is unchanged.
|
||||
|
||||
type AttributeMatcher interface {
|
||||
Match(*Domain) bool
|
||||
}
|
||||
@@ -185,23 +222,137 @@ func NewAllAttrsMatcher(attrs string) AttributeMatcher {
|
||||
return m
|
||||
}
|
||||
|
||||
func loadSiteWithAttrs(file, code, attrs string) ([]*Domain, error) {
|
||||
domains, err := loadSite(file, code)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var errInvalidUTF8 = errors.New("string field contains invalid UTF-8")
|
||||
|
||||
matcher := NewAllAttrsMatcher(attrs)
|
||||
if matcher == nil {
|
||||
return domains, nil
|
||||
}
|
||||
type siteDecoder struct {
|
||||
want []string
|
||||
has []bool
|
||||
fn func(Domain_Type, []byte)
|
||||
}
|
||||
|
||||
filtered := make([]*Domain, 0, len(domains))
|
||||
for _, d := range domains {
|
||||
if matcher.Match(d) {
|
||||
filtered = append(filtered, d)
|
||||
func newSiteDecoder(attrs string, fn func(Domain_Type, []byte)) *siteDecoder {
|
||||
d := &siteDecoder{fn: fn}
|
||||
if attrs != "" {
|
||||
d.want = strings.Split(attrs, "@")
|
||||
d.has = make([]bool, len(d.want))
|
||||
}
|
||||
return d
|
||||
}
|
||||
|
||||
// decode walks the whole fields at the start of b, a part of an encoded GeoSite (see geodat.proto),
|
||||
// calls fn for every domain that has all attrs and returns how many bytes it used. A field cut off
|
||||
// by the end of b is an error unless more is set. It accepts and rejects what proto.Unmarshal does.
|
||||
func (d *siteDecoder) decode(b []byte, more bool) (int, error) {
|
||||
used := 0
|
||||
for used < len(b) {
|
||||
f, n, err := consumeField(b[used:])
|
||||
if err == io.ErrUnexpectedEOF && more {
|
||||
break
|
||||
}
|
||||
if err != nil {
|
||||
return used, err
|
||||
}
|
||||
used += n
|
||||
if f.typ != protowire.BytesType {
|
||||
continue
|
||||
}
|
||||
switch f.num {
|
||||
case 1: // code
|
||||
if !utf8.Valid(f.v) {
|
||||
return used, errInvalidUTF8
|
||||
}
|
||||
case 2: // domain
|
||||
t, value, err := decodeDomain(f.v, d.want, d.has)
|
||||
if err != nil {
|
||||
return used, err
|
||||
}
|
||||
if !slices.Contains(d.has, false) {
|
||||
d.fn(t, value)
|
||||
}
|
||||
}
|
||||
}
|
||||
return used, nil
|
||||
}
|
||||
|
||||
return filtered, nil
|
||||
// decodeDomain decodes an encoded Domain and sets has[i] if one of its attributes has the key want[i].
|
||||
func decodeDomain(b []byte, want []string, has []bool) (t Domain_Type, value []byte, err error) {
|
||||
clear(has)
|
||||
for len(b) > 0 {
|
||||
f, n, err := consumeField(b)
|
||||
if err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
b = b[n:]
|
||||
switch {
|
||||
case f.num == 1 && f.typ == protowire.VarintType: // type
|
||||
t = Domain_Type(f.x)
|
||||
case f.num == 2 && f.typ == protowire.BytesType: // value
|
||||
if !utf8.Valid(f.v) {
|
||||
return 0, nil, errInvalidUTF8
|
||||
}
|
||||
value = f.v
|
||||
case f.num == 3 && f.typ == protowire.BytesType: // attribute
|
||||
key, err := decodeAttributeKey(f.v)
|
||||
if err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
for i, w := range want {
|
||||
if string(key) == w {
|
||||
has[i] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return t, value, nil
|
||||
}
|
||||
|
||||
// decodeAttributeKey returns the key of an encoded Domain.Attribute.
|
||||
func decodeAttributeKey(b []byte) ([]byte, error) {
|
||||
var key []byte
|
||||
for len(b) > 0 {
|
||||
f, n, err := consumeField(b)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
b = b[n:]
|
||||
if f.num == 1 && f.typ == protowire.BytesType {
|
||||
if !utf8.Valid(f.v) {
|
||||
return nil, errInvalidUTF8
|
||||
}
|
||||
key = f.v
|
||||
}
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
type protoField struct {
|
||||
num protowire.Number
|
||||
typ protowire.Type
|
||||
v []byte // payload of a length-delimited field
|
||||
x uint64 // value of a varint field
|
||||
}
|
||||
|
||||
// consumeField parses the first field of an encoded message and returns it with its length.
|
||||
func consumeField(b []byte) (protoField, int, error) {
|
||||
num, typ, n := protowire.ConsumeTag(b)
|
||||
if n < 0 {
|
||||
return protoField{}, 0, protowire.ParseError(n)
|
||||
}
|
||||
if num > protowire.MaxValidNumber {
|
||||
return protoField{}, 0, errors.New("invalid field number ", num)
|
||||
}
|
||||
f := protoField{num: num, typ: typ}
|
||||
var m int
|
||||
switch typ {
|
||||
case protowire.BytesType:
|
||||
f.v, m = protowire.ConsumeBytes(b[n:])
|
||||
case protowire.VarintType:
|
||||
f.x, m = protowire.ConsumeVarint(b[n:])
|
||||
default:
|
||||
m = protowire.ConsumeFieldValue(num, typ, b[n:])
|
||||
}
|
||||
if m < 0 {
|
||||
return protoField{}, 0, protowire.ParseError(m)
|
||||
}
|
||||
return f, n + m, nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,283 @@
|
||||
package geodata
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"google.golang.org/protobuf/encoding/protowire"
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
type siteEntry struct {
|
||||
Type Domain_Type
|
||||
Value string
|
||||
}
|
||||
|
||||
// unmarshalSite is what loadSite used to do: proto.Unmarshal, then keep the domains that have all attrs.
|
||||
func unmarshalSite(b []byte, attrs string) ([]siteEntry, error) {
|
||||
var site GeoSite
|
||||
if err := proto.Unmarshal(b, &site); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var entries []siteEntry
|
||||
for _, d := range site.Domain {
|
||||
ok := true
|
||||
for _, key := range strings.Split(attrs, "@") {
|
||||
ok = ok && (attrs == "" || slices.ContainsFunc(d.Attribute, func(a *Domain_Attribute) bool { return a.Key == key }))
|
||||
}
|
||||
if ok {
|
||||
entries = append(entries, siteEntry{d.Type, d.Value})
|
||||
}
|
||||
}
|
||||
return entries, nil
|
||||
}
|
||||
|
||||
func checkDecodeSite(t *testing.T, name string, b []byte, attrs string) {
|
||||
t.Helper()
|
||||
want, wantErr := unmarshalSite(b, attrs)
|
||||
var got []siteEntry
|
||||
_, err := newSiteDecoder(attrs, func(typ Domain_Type, value []byte) {
|
||||
got = append(got, siteEntry{typ, string(value)})
|
||||
}).decode(b, false)
|
||||
if (err == nil) != (wantErr == nil) {
|
||||
t.Fatalf("%s@%s: error %v, proto.Unmarshal: %v", name, attrs, err, wantErr)
|
||||
}
|
||||
if err == nil && !slices.Equal(got, want) {
|
||||
t.Fatalf("%s@%s: got %v, want %v", name, attrs, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeSiteMatchesUnmarshal(t *testing.T) {
|
||||
bs, err := os.ReadFile(filepath.Join("..", "..", "resources", DefaultGeoSiteDat))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for len(bs) > 0 {
|
||||
num, typ, n := protowire.ConsumeTag(bs)
|
||||
if n < 0 || num != 1 || typ != protowire.BytesType {
|
||||
t.Fatal("unexpected GeoSiteList field")
|
||||
}
|
||||
entry, m := protowire.ConsumeBytes(bs[n:])
|
||||
if m < 0 {
|
||||
t.Fatal(protowire.ParseError(m))
|
||||
}
|
||||
bs = bs[n+m:]
|
||||
|
||||
var site GeoSite
|
||||
if err := proto.Unmarshal(entry, &site); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
queries := []string{"", "none"}
|
||||
for _, d := range site.Domain {
|
||||
for _, a := range d.Attribute {
|
||||
if !slices.Contains(queries, a.Key) {
|
||||
queries = append(queries, a.Key, a.Key+"@none")
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, attrs := range queries {
|
||||
checkDecodeSite(t, site.Code, entry, attrs)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeSiteUnusualEncodings(t *testing.T) {
|
||||
field := func(num protowire.Number, v []byte) []byte {
|
||||
return protowire.AppendBytes(protowire.AppendTag(nil, num, protowire.BytesType), v)
|
||||
}
|
||||
typ := func(v Domain_Type) []byte {
|
||||
return protowire.AppendVarint(protowire.AppendTag(nil, 1, protowire.VarintType), uint64(v))
|
||||
}
|
||||
value := func(s string) []byte { return field(2, []byte(s)) }
|
||||
attr := func(keys ...string) []byte {
|
||||
var b []byte
|
||||
for _, k := range keys {
|
||||
b = append(b, field(1, []byte(k))...)
|
||||
}
|
||||
return field(3, b)
|
||||
}
|
||||
domain := func(fields ...[]byte) []byte { return field(2, slices.Concat(fields...)) }
|
||||
unknown := protowire.AppendFixed32(protowire.AppendTag(nil, 9, protowire.Fixed32Type), 1)
|
||||
|
||||
for name, b := range map[string][]byte{
|
||||
"unknown field": domain(typ(Domain_Full), unknown, value("example.com")),
|
||||
"repeated value": domain(value("a.com"), typ(Domain_Full), value("b.com")),
|
||||
"repeated type": domain(typ(Domain_Full), value("a.com"), typ(Domain_Regex)),
|
||||
"repeated key": domain(value("a.com"), attr("cn", "ads")),
|
||||
"type as bytes": domain(field(1, []byte("x")), value("a.com")),
|
||||
"no value": domain(typ(Domain_Domain), attr("cn")),
|
||||
"truncated": domain(typ(Domain_Full), value("example.com"))[:10],
|
||||
"invalid utf8": domain(value("example.\xff")),
|
||||
"invalid key": domain(value("a.com"), attr("\xff")),
|
||||
"bad field": protowire.AppendVarint(protowire.AppendTag(nil, protowire.MaxValidNumber+1, protowire.VarintType), 1),
|
||||
"stray end group": protowire.AppendTag(nil, 5, protowire.EndGroupType),
|
||||
} {
|
||||
for _, attrs := range []string{"", "cn", "ads", "cn@ads"} {
|
||||
checkDecodeSite(t, name, b, attrs)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadSiteReadsInPieces covers what real lists never do: an entry far longer than the read
|
||||
// buffer, with a field longer than the buffer in the middle, and a file cut short.
|
||||
func TestLoadSiteReadsInPieces(t *testing.T) {
|
||||
site := &GeoSite{Code: "BIG"}
|
||||
for i := range 5000 {
|
||||
d := &Domain{Type: Domain_Domain, Value: strings.Repeat("x", i%40) + ".example.com"}
|
||||
if i%3 == 0 {
|
||||
d.Attribute = []*Domain_Attribute{{Key: "cn"}}
|
||||
}
|
||||
if i == 2500 {
|
||||
d = &Domain{Type: Domain_Regex, Value: strings.Repeat("a", 100_000)}
|
||||
}
|
||||
site.Domain = append(site.Domain, d)
|
||||
}
|
||||
list := &GeoSiteList{Entry: []*GeoSite{{Code: "SMALL", Domain: []*Domain{{Type: Domain_Full, Value: "a.com"}}}, site}}
|
||||
bs, err := proto.Marshal(list)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
entry, err := proto.Marshal(site)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
dir := t.TempDir()
|
||||
t.Setenv("xray.location.asset", dir)
|
||||
write := func(b []byte) {
|
||||
if err := os.WriteFile(filepath.Join(dir, "big.dat"), b, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
for _, attrs := range []string{"", "cn"} {
|
||||
want, _ := unmarshalSite(entry, attrs)
|
||||
var got []siteEntry
|
||||
write(bs)
|
||||
err := loadSite("big.dat", "BIG", attrs, func(typ Domain_Type, value []byte) {
|
||||
got = append(got, siteEntry{typ, string(value)})
|
||||
})
|
||||
if err != nil || !slices.Equal(got, want) {
|
||||
t.Fatalf("attrs %q: %d entries, want %d, error %v", attrs, len(got), len(want), err)
|
||||
}
|
||||
for _, cut := range []int{30_000, len(bs) - 150_000, len(bs) - 1} {
|
||||
write(bs[:cut])
|
||||
if err := loadSite("big.dat", "BIG", attrs, func(Domain_Type, []byte) {}); err == nil {
|
||||
t.Fatalf("file cut at %d of %d: no error", cut, len(bs))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// oneEntryGeoSiteFile wraps an encoded GeoSite as a one-entry GeoSiteList, the file loadSite reads.
|
||||
func oneEntryGeoSiteFile(entry []byte) []byte {
|
||||
return protowire.AppendBytes(protowire.AppendTag(nil, 1, protowire.BytesType), entry)
|
||||
}
|
||||
|
||||
// TestLoadSiteWindowedMatchesSingleShot checks that the windowed reader in loadSite (its Peek/Discard
|
||||
// loop, the more-break when a field is cut by a window edge, the used==0 fallback for a field longer
|
||||
// than the buffer, and the tail path) reaches exactly the same result as decoding the whole entry at
|
||||
// once, for a category several 64 KiB windows long, valid and then mutated near a window edge and
|
||||
// early in the file: same error-or-not, and the same emitted (type, value) sequence when both accept.
|
||||
func TestLoadSiteWindowedMatchesSingleShot(t *testing.T) {
|
||||
const window = 64 * 1024
|
||||
site := &GeoSite{Code: "BIG"}
|
||||
for i := range 12000 { // ~250 KiB, four windows
|
||||
d := &Domain{Type: Domain_Domain, Value: fmt.Sprintf("host%d.%s.example.com", i, strings.Repeat("y", i%30))}
|
||||
if i%3 == 0 {
|
||||
d.Attribute = []*Domain_Attribute{{Key: "cn"}}
|
||||
}
|
||||
site.Domain = append(site.Domain, d)
|
||||
}
|
||||
// a field longer than the buffer, straddling the third window, to force the used==0 fallback
|
||||
site.Domain = slices.Insert(site.Domain, 8000, &Domain{Type: Domain_Regex, Value: strings.Repeat("a", 90_000)})
|
||||
entry, err := proto.Marshal(site)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
dir := t.TempDir()
|
||||
t.Setenv("xray.location.asset", dir)
|
||||
|
||||
// mutations of the encoded entry: unchanged, a byte flipped at several offsets (early windows and
|
||||
// either side of a window edge), and truncations at the same places.
|
||||
type mut struct {
|
||||
name string
|
||||
make func([]byte) []byte
|
||||
}
|
||||
muts := []mut{{"valid", func(b []byte) []byte { return b }}}
|
||||
for _, off := range []int{3, 40, 4000, window - 1, window, window + 1, 2*window - 2, 2 * window} {
|
||||
if off < len(entry) {
|
||||
off := off
|
||||
muts = append(muts, mut{fmt.Sprintf("flip@%d", off), func(b []byte) []byte {
|
||||
c := slices.Clone(b)
|
||||
c[off] ^= 0xff
|
||||
return c
|
||||
}})
|
||||
muts = append(muts, mut{fmt.Sprintf("cut@%d", off), func(b []byte) []byte { return slices.Clone(b[:off]) }})
|
||||
}
|
||||
}
|
||||
|
||||
for _, attrs := range []string{"", "cn"} {
|
||||
for _, m := range muts {
|
||||
e := m.make(entry)
|
||||
// single-shot reference: decode the whole entry in one call
|
||||
var want []siteEntry
|
||||
_, wantErr := newSiteDecoder(attrs, func(typ Domain_Type, value []byte) {
|
||||
want = append(want, siteEntry{typ, string(value)})
|
||||
}).decode(e, false)
|
||||
// windowed: loadSite reads the file 64 KiB at a time
|
||||
if err := os.WriteFile(filepath.Join(dir, "w.dat"), oneEntryGeoSiteFile(e), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var got []siteEntry
|
||||
gotErr := loadSite("w.dat", "BIG", attrs, func(typ Domain_Type, value []byte) {
|
||||
got = append(got, siteEntry{typ, string(value)})
|
||||
})
|
||||
if (gotErr == nil) != (wantErr == nil) {
|
||||
t.Fatalf("%s attrs=%q: windowed err %v, single-shot err %v", m.name, attrs, gotErr, wantErr)
|
||||
}
|
||||
if gotErr == nil && !slices.Equal(got, want) {
|
||||
t.Fatalf("%s attrs=%q: windowed got %d entries, single-shot %d", m.name, attrs, len(got), len(want))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadSiteLongCode covers a geosite entry whose code is longer than the 64 KiB read buffer. seek
|
||||
// must skip it (find compares a single length byte, so it never matches such a code) and still find a
|
||||
// later entry, and looking the long code up must fail cleanly, like a missing code, not panic.
|
||||
func TestLoadSiteLongCode(t *testing.T) {
|
||||
longCode := strings.Repeat("Z", 70000)
|
||||
list := &GeoSiteList{Entry: []*GeoSite{
|
||||
{Code: "FIRST", Domain: []*Domain{{Type: Domain_Full, Value: "first.com"}}},
|
||||
{Code: longCode, Domain: []*Domain{{Type: Domain_Full, Value: "huge.com"}}},
|
||||
{Code: "AFTER", Domain: []*Domain{{Type: Domain_Domain, Value: "after.com"}}},
|
||||
}}
|
||||
bs, err := proto.Marshal(list)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
dir := t.TempDir()
|
||||
t.Setenv("xray.location.asset", dir)
|
||||
if err := os.WriteFile(filepath.Join(dir, "lc.dat"), bs, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
collect := func(code string) ([]siteEntry, error) {
|
||||
var got []siteEntry
|
||||
err := loadSite("lc.dat", code, "", func(typ Domain_Type, value []byte) {
|
||||
got = append(got, siteEntry{typ, string(value)})
|
||||
})
|
||||
return got, err
|
||||
}
|
||||
if got, err := collect("FIRST"); err != nil || !slices.Equal(got, []siteEntry{{Domain_Full, "first.com"}}) {
|
||||
t.Fatalf("FIRST: %v %v", got, err)
|
||||
}
|
||||
if got, err := collect("AFTER"); err != nil || !slices.Equal(got, []siteEntry{{Domain_Domain, "after.com"}}) {
|
||||
t.Fatalf("AFTER (past the oversized entry): %v %v", got, err)
|
||||
}
|
||||
if _, err := collect(longCode); err == nil {
|
||||
t.Fatal("oversized code: expected a not-found error, got nil")
|
||||
}
|
||||
}
|
||||
@@ -52,7 +52,9 @@ func (g *MphIndexMatcher) Add(matcher Matcher) uint32 {
|
||||
func (g *MphIndexMatcher) Build() error {
|
||||
if g.mph != nil {
|
||||
runtime.GC() // peak mem
|
||||
g.mph.Build()
|
||||
if err := g.mph.Build(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
runtime.GC() // peak mem
|
||||
if g.ac != nil {
|
||||
@@ -64,23 +66,17 @@ func (g *MphIndexMatcher) Build() error {
|
||||
|
||||
// Match implements IndexMatcher.Match.
|
||||
func (g *MphIndexMatcher) Match(input string) []uint32 {
|
||||
result := make([][]uint32, 0, 5)
|
||||
var result []uint32
|
||||
if g.mph != nil {
|
||||
if matches := g.mph.Match(input); len(matches) > 0 {
|
||||
result = append(result, matches)
|
||||
}
|
||||
result = g.mph.Match(input) // a new slice, returned without another copy
|
||||
}
|
||||
if g.ac != nil {
|
||||
if matches := g.ac.Match(input); len(matches) > 0 {
|
||||
result = append(result, matches)
|
||||
}
|
||||
result = append(result, g.ac.Match(input)...)
|
||||
}
|
||||
if g.regex != nil {
|
||||
if matches := g.regex.Match(input); len(matches) > 0 {
|
||||
result = append(result, matches)
|
||||
}
|
||||
result = append(result, g.regex.Match(input)...)
|
||||
}
|
||||
return CompositeMatches(result)
|
||||
return result
|
||||
}
|
||||
|
||||
// MatchAny implements IndexMatcher.MatchAny.
|
||||
|
||||
@@ -78,6 +78,10 @@ func TestMphIndexMatcher(t *testing.T) {
|
||||
Input: "example.com",
|
||||
Output: []uint32{10, 4},
|
||||
},
|
||||
{
|
||||
Input: "apis.org",
|
||||
Output: []uint32{2, 6},
|
||||
},
|
||||
}
|
||||
matcherGroup := NewMphIndexMatcher()
|
||||
for _, rule := range rules {
|
||||
@@ -87,8 +91,13 @@ func TestMphIndexMatcher(t *testing.T) {
|
||||
}
|
||||
matcherGroup.Build()
|
||||
for _, test := range cases {
|
||||
if m := matcherGroup.Match(test.Input); !reflect.DeepEqual(m, test.Output) {
|
||||
m := matcherGroup.Match(test.Input)
|
||||
if !reflect.DeepEqual(m, test.Output) {
|
||||
t.Error("unexpected output: ", m, " for test case ", test)
|
||||
}
|
||||
clear(m) // the caller owns the result, so this must not change the next one
|
||||
if m := matcherGroup.Match(test.Input); !reflect.DeepEqual(m, test.Output) {
|
||||
t.Error("unexpected output after clearing the previous one: ", m, " for test case ", test)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,231 +1,440 @@
|
||||
package strmatcher
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"cmp"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"math"
|
||||
"math/bits"
|
||||
"runtime"
|
||||
"sort"
|
||||
"slices"
|
||||
"strings"
|
||||
"unsafe"
|
||||
)
|
||||
|
||||
// PrimeRK is the prime base used in Rabin-Karp algorithm.
|
||||
const PrimeRK = 16777619
|
||||
|
||||
// RollingHash calculates the rolling murmurHash of given string based on a provided suffix hash.
|
||||
func RollingHash(hash uint32, input string) uint32 {
|
||||
for i := len(input) - 1; i >= 0; i-- {
|
||||
hash = hash*PrimeRK + uint32(input[i])
|
||||
}
|
||||
return hash
|
||||
}
|
||||
|
||||
// MemHash is the hash function used by go map, it utilizes available hardware instructions(behaves
|
||||
// as aeshash if aes instruction is available).
|
||||
// With different seed, each MemHash<seed> performs as distinct hash functions.
|
||||
func MemHash(seed uint32, input string) uint32 {
|
||||
return uint32(strhash(unsafe.Pointer(&input), uintptr(seed))) // nosemgrep
|
||||
}
|
||||
|
||||
// Flags of a level1 slot, stored above the record offset.
|
||||
const (
|
||||
mphMatchTypeCount = 2 // Full and Domain
|
||||
mphDomain = 1 << 31 // matches the pattern and its subdomains
|
||||
mphFull = 1 << 30 // matches the pattern only
|
||||
mphParent = 1 << 29 // matches subdomains only, from a pattern with a leading dot
|
||||
mphOffMask = mphParent - 1
|
||||
)
|
||||
|
||||
type mphRuleInfo struct {
|
||||
rollingHash uint32
|
||||
matchers [mphMatchTypeCount][]uint32
|
||||
// Kinds of an added pattern, indexes of mphKinds.
|
||||
const (
|
||||
mphKindFull = iota
|
||||
mphKindParent
|
||||
mphKindDomain
|
||||
)
|
||||
|
||||
// mphKinds are the slot flags in the order Match reports their values.
|
||||
var mphKinds = [...]uint32{mphFull, mphParent, mphDomain}
|
||||
|
||||
// mphMultipliers are odd multipliers for the suffix hash. Build moves to the next one if two patterns collide.
|
||||
var mphMultipliers = [...]uint64{0x9e3779b97f4a7c15, 0xc2b2ae3d27d4eb4f, 0x165667b19e3779f9, 0x27d4eb2f165667c5}
|
||||
|
||||
var (
|
||||
errMphCollision = errors.New("strmatcher: suffix hash collision in MphMatcherGroup")
|
||||
errMphBuilt = errors.New("strmatcher: MphMatcherGroup is already built")
|
||||
)
|
||||
|
||||
type mphEntry struct {
|
||||
off uint32 // pattern start in buf
|
||||
value uint32
|
||||
n uint32 // pattern length
|
||||
kind uint8
|
||||
}
|
||||
|
||||
// MphMatcherGroup is an implementation of MatcherGroup.
|
||||
// It implements Rabin-Karp algorithm and minimal perfect hash table for Full and Domain matcher.
|
||||
// MphMatcherGroup is an implementation of MatcherGroup for Full and Domain matchers.
|
||||
// Each distinct pattern is stored once as a record in arena: its length (255 means a uvarint length follows),
|
||||
// its bytes and, if the group holds more than one distinct value, its values. A minimal perfect hash table
|
||||
// built with hash, displace and compress (http://cmph.sourceforge.net/papers/esa09.pdf) maps a pattern to its
|
||||
// record. Patterns are hashed from the right, so one pass over the input hashes all its parent domains.
|
||||
type MphMatcherGroup struct {
|
||||
patterns string // All rule patterns concatenated
|
||||
patternOffs []uint32 // RuleIdx -> patterns[patternOffs[i]:patternOffs[i+1]], index 0 reserved for failed lookup
|
||||
values []uint32 // All registered matcher values concatenated
|
||||
valueOffs []uint32 // RuleIdx -> values[valueOffs[i]:valueOffs[i+1]] (Full Matcher takes precedence)
|
||||
level0 []uint32 // RollingHash & Mask -> seed for Memhash
|
||||
level0Mask uint32 // Mask restricting RollingHash to 0 ~ len(level0)
|
||||
level1 []uint32 // Memhash<seed> & Mask -> stored index for rules
|
||||
level1Mask uint32 // Mask for restricting Memhash<seed> to 0 ~ len(level1)
|
||||
rules []string // RuleIdx -> pattern string, only used for building
|
||||
ruleInfos *map[string]mphRuleInfo
|
||||
arena string
|
||||
level0 []uint16 // bucket -> seed
|
||||
level1 []uint32 // slot -> flags | record offset
|
||||
fp []uint8 // slot -> low byte of its pattern's hash, rejects most misses without reading arena
|
||||
n0, n1 uint32
|
||||
mul uint64 // multiplier of the suffix hash
|
||||
single uint32 // the only value if !multi
|
||||
multi bool
|
||||
|
||||
buf []byte // build only, patterns in Add order
|
||||
entries []mphEntry
|
||||
}
|
||||
|
||||
func NewMphMatcherGroup() *MphMatcherGroup {
|
||||
return &MphMatcherGroup{
|
||||
rules: []string{""},
|
||||
level0: nil,
|
||||
level0Mask: 0,
|
||||
level1: nil,
|
||||
level1Mask: 0,
|
||||
ruleInfos: &map[string]mphRuleInfo{}, // Only used for building, destroyed after build complete
|
||||
}
|
||||
return new(MphMatcherGroup)
|
||||
}
|
||||
|
||||
// AddFullMatcher implements MatcherGroupForFull.
|
||||
func (g *MphMatcherGroup) AddFullMatcher(matcher FullMatcher, value uint32) {
|
||||
pattern := strings.ToLower(matcher.Pattern())
|
||||
g.addPattern(0, "", pattern, matcher.Type(), value)
|
||||
g.add(matcher.Pattern(), mphKindFull, value)
|
||||
}
|
||||
|
||||
// AddDomainMatcher implements MatcherGroupForDomain.
|
||||
func (g *MphMatcherGroup) AddDomainMatcher(matcher DomainMatcher, value uint32) {
|
||||
pattern := strings.ToLower(matcher.Pattern())
|
||||
hash := g.addPattern(0, "", pattern, matcher.Type(), value) // For full domain match
|
||||
g.addPattern(hash, pattern, ".", matcher.Type(), value) // For partial domain match
|
||||
g.add(matcher.Pattern(), mphKindDomain, value)
|
||||
}
|
||||
|
||||
func (g *MphMatcherGroup) addPattern(suffixHash uint32, suffixPattern string, pattern string, matcherType Type, value uint32) uint32 {
|
||||
fullPattern := pattern + suffixPattern
|
||||
info, found := (*g.ruleInfos)[fullPattern]
|
||||
if !found {
|
||||
info = mphRuleInfo{rollingHash: RollingHash(suffixHash, pattern)}
|
||||
g.rules = append(g.rules, fullPattern)
|
||||
func (g *MphMatcherGroup) add(pattern string, kind uint8, value uint32) {
|
||||
if g.arena != "" {
|
||||
panic(errMphBuilt)
|
||||
}
|
||||
pattern = strings.ToLower(pattern)
|
||||
off := uint32(len(g.buf))
|
||||
g.buf = append(g.buf, pattern...)
|
||||
g.entries = append(g.entries, mphEntry{off: off, value: value, n: uint32(len(pattern)), kind: kind})
|
||||
if len(pattern) > 0 && pattern[0] == '.' {
|
||||
// ".x" has always matched "*.x" as well, so it also gets a parent-only record for "x"
|
||||
g.entries = append(g.entries, mphEntry{off: off + 1, value: value, n: uint32(len(pattern) - 1), kind: mphKindParent})
|
||||
}
|
||||
info.matchers[matcherType] = append(info.matchers[matcherType], value)
|
||||
(*g.ruleInfos)[fullPattern] = info
|
||||
return info.rollingHash
|
||||
}
|
||||
|
||||
// Build builds a minimal perfect hash table for insert rules.
|
||||
// Algorithm used: Hash, displace, and compress. See http://cmph.sourceforge.net/papers/esa09.pdf
|
||||
func (g *MphMatcherGroup) key(i uint32) []byte {
|
||||
e := &g.entries[i]
|
||||
return g.buf[e.off : e.off+e.n]
|
||||
}
|
||||
|
||||
// Build builds the hash table. It must be called once, after the last Add.
|
||||
func (g *MphMatcherGroup) Build() error {
|
||||
ruleCount := len(*g.ruleInfos)
|
||||
g.level0 = make([]uint32, nextPow2(ruleCount/4))
|
||||
g.level0Mask = uint32(len(g.level0) - 1)
|
||||
g.level1 = make([]uint32, nextPow2(ruleCount))
|
||||
g.level1Mask = uint32(len(g.level1) - 1)
|
||||
|
||||
// Flatten patterns and values so the built group has no per-rule objects
|
||||
valueCount := 0
|
||||
for _, ruleInfo := range *g.ruleInfos {
|
||||
valueCount += len(ruleInfo.matchers[Full]) + len(ruleInfo.matchers[Domain])
|
||||
if g.arena != "" {
|
||||
return errMphBuilt
|
||||
}
|
||||
g.patterns = strings.Join(g.rules, "")
|
||||
if uint64(len(g.patterns)) > math.MaxUint32 || uint64(valueCount) > math.MaxUint32 {
|
||||
if uint64(len(g.buf)) > math.MaxUint32 {
|
||||
return errors.New("too many rules for MphMatcherGroup")
|
||||
}
|
||||
g.patternOffs = make([]uint32, len(g.rules)+1)
|
||||
g.values = make([]uint32, 0, valueCount)
|
||||
g.valueOffs = make([]uint32, len(g.rules)+1)
|
||||
|
||||
// Create buckets based on all rule's rolling hash
|
||||
buckets := make([][]uint32, len(g.level0))
|
||||
for ruleIdx := 1; ruleIdx < len(g.rules); ruleIdx++ { // Traverse rules starting from index 1 (0 reserved for failed lookup)
|
||||
ruleInfo := (*g.ruleInfos)[g.rules[ruleIdx]]
|
||||
bucketIdx := ruleInfo.rollingHash & g.level0Mask
|
||||
buckets[bucketIdx] = append(buckets[bucketIdx], uint32(ruleIdx))
|
||||
g.patternOffs[ruleIdx+1] = g.patternOffs[ruleIdx] + uint32(len(g.rules[ruleIdx]))
|
||||
g.values = append(append(g.values, ruleInfo.matchers[Full]...), ruleInfo.matchers[Domain]...)
|
||||
g.valueOffs[ruleIdx+1] = uint32(len(g.values))
|
||||
recs := g.writeRecords()
|
||||
if len(g.arena) > mphOffMask {
|
||||
return errors.New("too many rules for MphMatcherGroup")
|
||||
}
|
||||
g.rules = nil
|
||||
g.ruleInfos = nil // Set ruleInfos nil to release memory
|
||||
runtime.GC() // peak mem
|
||||
|
||||
// Sort buckets in descending order with respect to each bucket's size
|
||||
bucketIdxs := make([]int, len(buckets))
|
||||
for bucketIdx := range buckets {
|
||||
bucketIdxs[bucketIdx] = bucketIdx
|
||||
hashes := make([]uint64, len(recs))
|
||||
for _, mul := range mphMultipliers {
|
||||
for i, rec := range recs {
|
||||
hashes[i] = mphMix(mphHash(mul, g.recKey(rec)))
|
||||
}
|
||||
g.mul = mul
|
||||
if err := g.place(recs, hashes); err != errMphCollision {
|
||||
return err
|
||||
}
|
||||
}
|
||||
sort.Slice(bucketIdxs, func(i, j int) bool { return len(buckets[bucketIdxs[i]]) > len(buckets[bucketIdxs[j]]) })
|
||||
return errMphCollision
|
||||
}
|
||||
|
||||
// Exercise Hash, Displace, and Compress algorithm to construct minimal perfect hash table
|
||||
occupied := make([]bool, len(g.level1)) // Whether a second-level hash has been already used
|
||||
hashedBucket := make([]uint32, 0, 4) // Second-level hashes for each rule in a specific bucket
|
||||
for _, bucketIdx := range bucketIdxs {
|
||||
bucket := buckets[bucketIdx]
|
||||
hashedBucket = hashedBucket[:0]
|
||||
seed := uint32(0)
|
||||
for len(hashedBucket) != len(bucket) {
|
||||
for _, ruleIdx := range bucket {
|
||||
memHash := MemHash(seed, g.pattern(ruleIdx)) & g.level1Mask
|
||||
if occupied[memHash] { // Collision occurred with this seed
|
||||
for _, hash := range hashedBucket { // Revert all values in this hashed bucket
|
||||
occupied[hash] = false
|
||||
g.level1[hash] = 0
|
||||
}
|
||||
hashedBucket = hashedBucket[:0]
|
||||
seed++ // Try next seed
|
||||
break
|
||||
}
|
||||
occupied[memHash] = true
|
||||
g.level1[memHash] = ruleIdx // The final value in the hash table
|
||||
hashedBucket = append(hashedBucket, memHash)
|
||||
// writeRecords writes one record per distinct pattern to arena and returns flags | offset of each.
|
||||
func (g *MphMatcherGroup) writeRecords() []uint32 {
|
||||
g.multi = false
|
||||
if len(g.entries) > 0 {
|
||||
g.single = g.entries[0].value
|
||||
for _, e := range g.entries {
|
||||
if e.value != g.single {
|
||||
g.multi = true
|
||||
break
|
||||
}
|
||||
}
|
||||
g.level0[bucketIdx] = seed // Displacement value for this bucket
|
||||
}
|
||||
// Equal patterns become neighbours in Add order, so their values keep their priority
|
||||
order := make([]uint32, len(g.entries))
|
||||
for i := range order {
|
||||
order[i] = uint32(i)
|
||||
}
|
||||
slices.SortFunc(order, func(a, b uint32) int {
|
||||
return cmp.Or(bytes.Compare(g.key(a), g.key(b)), cmp.Compare(a, b))
|
||||
})
|
||||
|
||||
size := len(g.buf) + len(g.entries) + 2
|
||||
if g.multi {
|
||||
size += 3 * len(g.entries)
|
||||
}
|
||||
arena := make([]byte, 0, size)
|
||||
recs := make([]uint32, 0, len(order))
|
||||
var vals [len(mphKinds)][]uint32
|
||||
for i := 0; i < len(order); {
|
||||
k := g.key(order[i])
|
||||
for t := range vals {
|
||||
vals[t] = vals[t][:0]
|
||||
}
|
||||
for ; i < len(order) && bytes.Equal(g.key(order[i]), k); i++ {
|
||||
e := &g.entries[order[i]]
|
||||
if !slices.Contains(vals[e.kind], e.value) {
|
||||
vals[e.kind] = append(vals[e.kind], e.value)
|
||||
}
|
||||
}
|
||||
rec := uint32(len(arena))
|
||||
if len(k) < 255 {
|
||||
arena = append(arena, byte(len(k)))
|
||||
} else {
|
||||
arena = binary.AppendUvarint(append(arena, 255), uint64(len(k)))
|
||||
}
|
||||
arena = append(arena, k...)
|
||||
for t, v := range vals {
|
||||
if len(v) == 0 {
|
||||
continue
|
||||
}
|
||||
rec |= mphKinds[t]
|
||||
if g.multi {
|
||||
arena = binary.AppendUvarint(arena, uint64(len(v)))
|
||||
for _, x := range v {
|
||||
arena = binary.AppendUvarint(arena, uint64(x))
|
||||
}
|
||||
}
|
||||
}
|
||||
recs = append(recs, rec)
|
||||
}
|
||||
// Lookups may point one byte past a pattern, and an empty group needs a record at offset 0 for empty slots
|
||||
arena = append(arena, 0)
|
||||
if len(recs) == 0 {
|
||||
arena = append(arena, 0)
|
||||
}
|
||||
g.buf, g.entries = nil, nil
|
||||
if cap(arena)-len(arena) > len(arena)/32 {
|
||||
arena = slices.Clone(arena)
|
||||
}
|
||||
g.arena = unsafe.String(unsafe.SliceData(arena), len(arena)) // arena is not written after this
|
||||
return recs
|
||||
}
|
||||
|
||||
// place fills level0, level1 and fp: records are bucketed by hash, and each bucket, largest first, gets
|
||||
// the first seed that puts all its records in free slots.
|
||||
func (g *MphMatcherGroup) place(recs []uint32, hashes []uint64) error {
|
||||
r := len(recs)
|
||||
n0, n1 := max(1, r/3), max(1, r+r/99)
|
||||
g.n0, g.n1 = uint32(n0), uint32(n1)
|
||||
g.level0 = make([]uint16, n0)
|
||||
g.level1 = make([]uint32, n1)
|
||||
g.fp = make([]uint8, n1)
|
||||
|
||||
start := make([]uint32, n0+1)
|
||||
for _, h := range hashes {
|
||||
start[g.bucket(h)+1]++
|
||||
}
|
||||
for b := range n0 {
|
||||
start[b+1] += start[b]
|
||||
}
|
||||
members := make([]uint32, r)
|
||||
fill := slices.Clone(start[:n0])
|
||||
for i, h := range hashes {
|
||||
b := g.bucket(h)
|
||||
members[fill[b]] = uint32(i)
|
||||
fill[b]++
|
||||
}
|
||||
fill = nil
|
||||
buckets := make([]uint32, n0)
|
||||
for b := range buckets {
|
||||
buckets[b] = uint32(b)
|
||||
}
|
||||
slices.SortStableFunc(buckets, func(a, b uint32) int {
|
||||
return cmp.Compare(start[b+1]-start[b], start[a+1]-start[a])
|
||||
})
|
||||
|
||||
occupied := make([]uint64, (n1+63)/64)
|
||||
var slots []uint32
|
||||
next:
|
||||
for _, b := range buckets {
|
||||
m := members[start[b]:start[b+1]]
|
||||
if len(m) == 0 {
|
||||
break
|
||||
}
|
||||
for i := range m {
|
||||
for j := range i {
|
||||
if hashes[m[i]] == hashes[m[j]] {
|
||||
return errMphCollision // no seed can separate them
|
||||
}
|
||||
}
|
||||
}
|
||||
search:
|
||||
for seed := range math.MaxUint16 + 1 {
|
||||
slots = slots[:0]
|
||||
for _, ri := range m {
|
||||
s := g.slot(hashes[ri], uint16(seed))
|
||||
if occupied[s/64]&(1<<(s%64)) != 0 || slices.Contains(slots, s) {
|
||||
continue search
|
||||
}
|
||||
slots = append(slots, s)
|
||||
}
|
||||
for k, ri := range m {
|
||||
s := slots[k]
|
||||
occupied[s/64] |= 1 << (s % 64)
|
||||
g.level1[s] = recs[ri]
|
||||
g.fp[s] = uint8(hashes[ri])
|
||||
}
|
||||
g.level0[b] = uint16(seed)
|
||||
continue next
|
||||
}
|
||||
return errors.New("strmatcher: no seed found for a bucket in MphMatcherGroup")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (g *MphMatcherGroup) pattern(ruleIdx uint32) string {
|
||||
return g.patterns[g.patternOffs[ruleIdx]:g.patternOffs[ruleIdx+1]]
|
||||
// mphHash is the suffix hash of s, taken from the right: the hash of s[i:] is the state after reading s[i].
|
||||
func mphHash(mul uint64, s string) uint64 {
|
||||
h := uint64(0)
|
||||
for i := len(s) - 1; i >= 0; i-- {
|
||||
h = h*mul + uint64(s[i])
|
||||
}
|
||||
return h
|
||||
}
|
||||
|
||||
// valuesOf caps the capacity, so appending to a Match result can't overwrite the next rule's values.
|
||||
func (g *MphMatcherGroup) valuesOf(ruleIdx uint32) []uint32 {
|
||||
start, end := g.valueOffs[ruleIdx], g.valueOffs[ruleIdx+1]
|
||||
return g.values[start:end:end]
|
||||
// mphMix spreads the weak low bits of a suffix hash.
|
||||
func mphMix(h uint64) uint64 {
|
||||
h ^= h >> 32
|
||||
h *= 0xd6e8feb86659fd93
|
||||
return h ^ h>>32
|
||||
}
|
||||
|
||||
// Lookup searches for input in minimal perfect hash table and returns its index. 0 indicates not found.
|
||||
func (g *MphMatcherGroup) Lookup(rollingHash uint32, input string) uint32 {
|
||||
i0 := rollingHash & g.level0Mask
|
||||
seed := g.level0[i0]
|
||||
i1 := MemHash(seed, input) & g.level1Mask
|
||||
n := g.level1[i1]
|
||||
// Build only puts valid rule indices in level1, so n+1 < len(patternOffs) and the span is inside patterns.
|
||||
// Skip the bounds checks, they made this hot path measurably slower than indexing a []string
|
||||
offs := (*[2]uint32)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.patternOffs)), uintptr(n)*4))
|
||||
if start := offs[0]; int(offs[1]-start) == len(input) && unsafe.String((*byte)(unsafe.Add(unsafe.Pointer(unsafe.StringData(g.patterns)), start)), len(input)) == input {
|
||||
return n
|
||||
func (g *MphMatcherGroup) bucket(f uint64) uint32 {
|
||||
return uint32(((f >> 32) * uint64(g.n0)) >> 32)
|
||||
}
|
||||
|
||||
func (g *MphMatcherGroup) slot(f uint64, seed uint16) uint32 {
|
||||
x := ((f ^ uint64(seed)*0x9e3779b97f4a7c15) * 0xc4ceb9fe1a85ec53) >> 32
|
||||
return uint32((x * uint64(g.n1)) >> 32)
|
||||
}
|
||||
|
||||
func (g *MphMatcherGroup) uvarint(p uint32) (x, next uint32) {
|
||||
for shift := 0; ; shift += 7 {
|
||||
c := g.arena[p]
|
||||
p++
|
||||
x |= uint32(c&0x7f) << shift
|
||||
if c < 0x80 {
|
||||
return x, p
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// recSpan returns where the pattern of the record at off starts and how long it is.
|
||||
func (g *MphMatcherGroup) recSpan(off uint32) (p, n uint32) {
|
||||
n, p = uint32(g.arena[off]), off+1
|
||||
if n == 255 {
|
||||
n, p = g.uvarint(p)
|
||||
}
|
||||
return p, n
|
||||
}
|
||||
|
||||
func (g *MphMatcherGroup) recKey(rec uint32) string {
|
||||
p, n := g.recSpan(rec & mphOffMask)
|
||||
return g.arena[p : p+n]
|
||||
}
|
||||
|
||||
// lookup returns the level1 entry of s, or 0 if s is not a pattern. h is the suffix hash of s.
|
||||
func (g *MphMatcherGroup) lookup(h uint64, s string) uint32 {
|
||||
f := mphMix(h)
|
||||
// bucket < n0 == len(level0) and slot < n1 == len(level1) == len(fp), skip the bounds checks
|
||||
seed := *(*uint16)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level0)), uintptr(g.bucket(f))*2))
|
||||
slot := uintptr(g.slot(f, seed))
|
||||
if *(*uint8)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.fp)), slot)) != uint8(f) {
|
||||
return 0
|
||||
}
|
||||
e := *(*uint32)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level1)), slot*4))
|
||||
if len(s) < 255 {
|
||||
// A record whose length byte is len(s) has len(s) pattern bytes after it
|
||||
p := unsafe.Add(unsafe.Pointer(unsafe.StringData(g.arena)), e&mphOffMask)
|
||||
if int(*(*byte)(p)) == len(s) && unsafe.String((*byte)(unsafe.Add(p, 1)), len(s)) == s {
|
||||
return e
|
||||
}
|
||||
return 0
|
||||
}
|
||||
if g.recKey(e) == s {
|
||||
return e
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// Match implements MatcherGroup.Match.
|
||||
func (g *MphMatcherGroup) Match(input string) []uint32 {
|
||||
matches := make([][]uint32, 0, 5)
|
||||
hash := uint32(0)
|
||||
for i := len(input) - 1; i >= 0; i-- {
|
||||
hash = hash*PrimeRK + uint32(input[i])
|
||||
if input[i] == '.' {
|
||||
if mphIdx := g.Lookup(hash, input[i:]); mphIdx != 0 {
|
||||
matches = append(matches, g.valuesOf(mphIdx))
|
||||
// appendValues appends the values of record e for the flags in want, in mphKinds order.
|
||||
func (g *MphMatcherGroup) appendValues(dst []uint32, e, want uint32) []uint32 {
|
||||
if !g.multi {
|
||||
for _, flag := range mphKinds {
|
||||
if e&want&flag != 0 {
|
||||
dst = append(dst, g.single)
|
||||
}
|
||||
}
|
||||
return dst
|
||||
}
|
||||
if e&want == 0 {
|
||||
return dst
|
||||
}
|
||||
p, n := g.recSpan(e & mphOffMask)
|
||||
p += n
|
||||
for _, flag := range mphKinds {
|
||||
if e&flag == 0 {
|
||||
continue
|
||||
}
|
||||
var count, v uint32
|
||||
for count, p = g.uvarint(p); count > 0; count-- {
|
||||
v, p = g.uvarint(p)
|
||||
if want&flag != 0 {
|
||||
dst = append(dst, v)
|
||||
}
|
||||
}
|
||||
}
|
||||
if mphIdx := g.Lookup(hash, input); mphIdx != 0 {
|
||||
matches = append(matches, g.valuesOf(mphIdx))
|
||||
return dst
|
||||
}
|
||||
|
||||
// Match implements MatcherGroup.Match. Values of an exact match come first (Full, then Domain), then those of
|
||||
// the parent domains, nearest first.
|
||||
func (g *MphMatcherGroup) Match(input string) []uint32 {
|
||||
var stack [8]uint32
|
||||
parents := stack[:0] // TLD side first
|
||||
h, mul := uint64(0), g.mul
|
||||
for i := len(input) - 1; i >= 0; i-- {
|
||||
if input[i] == '.' {
|
||||
if e := g.lookup(h, input[i+1:]); e&(mphDomain|mphParent) != 0 {
|
||||
parents = append(parents, e)
|
||||
}
|
||||
}
|
||||
h = h*mul + uint64(input[i])
|
||||
}
|
||||
return CompositeMatchesReverse(matches)
|
||||
exact := g.lookup(h, input)
|
||||
if exact&(mphFull|mphDomain) == 0 && len(parents) == 0 {
|
||||
return nil
|
||||
}
|
||||
result := g.appendValues(make([]uint32, 0, len(parents)+1), exact, mphFull|mphDomain)
|
||||
for k := len(parents) - 1; k >= 0; k-- {
|
||||
result = g.appendValues(result, parents[k], mphParent|mphDomain)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// MatchAny implements MatcherGroup.MatchAny.
|
||||
func (g *MphMatcherGroup) MatchAny(input string) bool {
|
||||
hash := uint32(0)
|
||||
h, mul := uint64(0), g.mul
|
||||
for i := len(input) - 1; i >= 0; i-- {
|
||||
if input[i] == '.' && g.lookup(h, input[i+1:])&(mphDomain|mphParent) != 0 {
|
||||
return true
|
||||
}
|
||||
h = h*mul + uint64(input[i])
|
||||
}
|
||||
return g.lookup(h, input)&(mphFull|mphDomain) != 0
|
||||
}
|
||||
|
||||
// mphSuffix is the suffix hash of input[off:], a parent domain of the input.
|
||||
type mphSuffix struct {
|
||||
h uint64
|
||||
off int
|
||||
}
|
||||
|
||||
// mphSuffixes appends the suffix hashes of the parent domains of input to dst, TLD side first, and returns them
|
||||
// with the hash of input itself: what MatchAny computes, computed once for several groups.
|
||||
func mphSuffixes(dst []mphSuffix, mul uint64, input string) ([]mphSuffix, uint64) {
|
||||
h := uint64(0)
|
||||
for i := len(input) - 1; i >= 0; i-- {
|
||||
hash = hash*PrimeRK + uint32(input[i])
|
||||
if input[i] == '.' {
|
||||
if g.Lookup(hash, input[i:]) != 0 {
|
||||
return true
|
||||
}
|
||||
dst = append(dst, mphSuffix{h, i + 1})
|
||||
}
|
||||
h = h*mul + uint64(input[i])
|
||||
}
|
||||
return dst, h
|
||||
}
|
||||
|
||||
// matchAnyHashed is MatchAny with parents and h from mphSuffixes(_, mul, input).
|
||||
func (g *MphMatcherGroup) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
|
||||
if g.mul != mul {
|
||||
return g.MatchAny(input) // built with a later multiplier after a collision
|
||||
}
|
||||
for _, p := range parents {
|
||||
if g.lookup(p.h, input[p.off:])&(mphDomain|mphParent) != 0 {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return g.Lookup(hash, input) != 0
|
||||
return g.lookup(h, input)&(mphFull|mphDomain) != 0
|
||||
}
|
||||
|
||||
func nextPow2(v int) int {
|
||||
if v <= 1 {
|
||||
return 1
|
||||
}
|
||||
const MaxUInt = ^uint(0)
|
||||
n := (MaxUInt >> bits.LeadingZeros(uint(v))) + 1
|
||||
return int(n)
|
||||
}
|
||||
|
||||
//go:noescape
|
||||
//go:linkname strhash runtime.strhash
|
||||
func strhash(p unsafe.Pointer, h uintptr) uintptr
|
||||
|
||||
@@ -0,0 +1,108 @@
|
||||
package strmatcher
|
||||
|
||||
import (
|
||||
"slices"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestMphMatcherGroupHashCollision(t *testing.T) {
|
||||
saved := mphMultipliers
|
||||
defer func() { mphMultipliers = saved }()
|
||||
|
||||
mphMultipliers[0] = 1 // anagrams collide
|
||||
g := NewMphMatcherGroup()
|
||||
g.AddFullMatcher(FullMatcher("ab.com"), 1)
|
||||
g.AddDomainMatcher(DomainMatcher("ba.com"), 2)
|
||||
g.AddDomainMatcher(DomainMatcher("com"), 3)
|
||||
if err := g.Build(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if g.mul != saved[1] {
|
||||
t.Errorf("multiplier %#x, want the second one %#x", g.mul, saved[1])
|
||||
}
|
||||
for input, want := range map[string][]uint32{"ab.com": {1, 3}, "x.ba.com": {2, 3}, "x.ab.com": {3}, "ba.com": {2, 3}} {
|
||||
if m := g.Match(input); !slices.Equal(m, want) {
|
||||
t.Errorf("Match(%q) = %v, want %v", input, m, want)
|
||||
}
|
||||
}
|
||||
|
||||
// Thue-Morse strings of 2048 bytes and their complements collide for every odd multiplier
|
||||
mphMultipliers = saved
|
||||
a, b := make([]byte, 2048), make([]byte, 2048)
|
||||
for i := range a {
|
||||
a[i], b[i] = "ab"[bitsOnes(i)%2], "ba"[bitsOnes(i)%2]
|
||||
}
|
||||
g = NewMphMatcherGroup()
|
||||
g.AddFullMatcher(FullMatcher(a), 1)
|
||||
g.AddFullMatcher(FullMatcher(b), 1)
|
||||
if err := g.Build(); err != errMphCollision {
|
||||
t.Errorf("Build() = %v, want %v", err, errMphCollision)
|
||||
}
|
||||
}
|
||||
|
||||
func bitsOnes(i int) int {
|
||||
n := 0
|
||||
for ; i > 0; i &= i - 1 {
|
||||
n++
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
func TestMphValueMatcherCombiner(t *testing.T) {
|
||||
build := func(matchers ...Matcher) *MphValueMatcher {
|
||||
m := NewMphValueMatcher()
|
||||
for _, x := range matchers {
|
||||
m.Add(x, 0)
|
||||
}
|
||||
if err := m.Build(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return m
|
||||
}
|
||||
regex, err := Regex.New(`^a\d+\.net$`)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
saved := mphMultipliers
|
||||
t.Cleanup(func() { mphMultipliers = saved })
|
||||
mphMultipliers[0] = 1 // anagrams collide, so this one falls back to its own hash pass
|
||||
collided := build(FullMatcher("ab.com"), DomainMatcher("ba.com"))
|
||||
mphMultipliers = saved
|
||||
if collided.mph.mul == mphMultipliers[0] {
|
||||
t.Fatal("collided matcher uses the first multiplier")
|
||||
}
|
||||
matchers := []*MphValueMatcher{
|
||||
build(DomainMatcher("example.com"), FullMatcher("full.org"), DomainMatcher(".dot.io")),
|
||||
collided,
|
||||
build(regex, SubstrMatcher("keyword")),
|
||||
build(),
|
||||
build(DomainMatcher("com"), DomainMatcher("a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s")),
|
||||
}
|
||||
var s MphValueMatcherCombiner
|
||||
for i, m := range matchers {
|
||||
s.Add(m, uint32(10+i))
|
||||
}
|
||||
inputs := []string{
|
||||
"", ".", "..", "com", "example.com", "www.example.com", "xexample.com", "example.com.", "full.org", "x.full.org",
|
||||
"dot.io", "x.dot.io", ".dot.io", "ab.com", "x.ab.com", "ba.com", "x.ba.com", "a12.net", "a12.net.x", "my-keyword.org",
|
||||
"a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s", "0.a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s", "b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s",
|
||||
"x.y.z.1.2.3.4.5.6.7.8.9.10.11.12.13.14.15.16.17.ab.com", "x.y.z.1.2.3.4.5.6.7.8.9.10.11.12.13.14.15.16.17.org",
|
||||
}
|
||||
for _, input := range inputs {
|
||||
var want []uint32
|
||||
for i, m := range matchers {
|
||||
if m.MatchAny(input) {
|
||||
want = append(want, uint32(10+i))
|
||||
}
|
||||
}
|
||||
if got := s.Match(input); !slices.Equal(got, want) {
|
||||
t.Errorf("Match(%q) = %v, want %v", input, got, want)
|
||||
}
|
||||
if got := s.MatchAny(input); got != (len(want) > 0) {
|
||||
t.Errorf("MatchAny(%q) = %v", input, got)
|
||||
}
|
||||
}
|
||||
if n := testing.AllocsPerRun(100, func() { s.MatchAny("www.a.b.c.example.org") }); n != 0 {
|
||||
t.Errorf("MatchAny allocates %v times", n)
|
||||
}
|
||||
}
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"math/rand"
|
||||
"reflect"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/xtls/xray-core/common"
|
||||
@@ -304,7 +305,7 @@ func TestMphMatcherGroupRandom(t *testing.T) {
|
||||
domain["."+p] = append(domain["."+p], value)
|
||||
}
|
||||
}
|
||||
g.Build()
|
||||
common.Must(g.Build())
|
||||
for _, input := range inputs {
|
||||
keys := []string{input} // Whole input first, then "." suffixes from longest to shortest
|
||||
for i := range len(input) {
|
||||
@@ -316,7 +317,10 @@ func TestMphMatcherGroupRandom(t *testing.T) {
|
||||
for _, k := range keys {
|
||||
want = append(append(want, full[k]...), domain[k]...)
|
||||
}
|
||||
if m := g.Match(input); !slices.Equal(m, want) {
|
||||
// Compared as sets: Match reports a value once per matching pattern, and orders them differently
|
||||
// from want for patterns and inputs with a leading dot
|
||||
m := g.Match(input)
|
||||
if !slices.Equal(sortedSet(m), sortedSet(want)) {
|
||||
t.Fatalf("seed %d: Match(%q) = %v, want %v", seed, input, m, want)
|
||||
}
|
||||
if m := g.MatchAny(input); m != (len(want) > 0) {
|
||||
@@ -338,3 +342,79 @@ func TestMphMatcherGroupAppend(t *testing.T) {
|
||||
t.Error("expect [2], but ", m)
|
||||
}
|
||||
}
|
||||
|
||||
func sortedSet(v []uint32) []uint32 {
|
||||
v = slices.Clone(v)
|
||||
slices.Sort(v)
|
||||
return slices.Compact(v)
|
||||
}
|
||||
|
||||
func TestMphMatcherGroupLongPattern(t *testing.T) {
|
||||
long := strings.Repeat("a", 300) + ".com"
|
||||
for _, values := range [][4]uint32{{1, 2, 3, 4}, {7, 7, 7, 7}} {
|
||||
g := NewMphMatcherGroup()
|
||||
g.AddDomainMatcher(DomainMatcher(long), values[0])
|
||||
g.AddFullMatcher(FullMatcher("x."+long), values[1])
|
||||
g.AddFullMatcher(FullMatcher(long[:255]), values[2]) // the shortest pattern stored with a long length
|
||||
g.AddFullMatcher(FullMatcher(long[:254]), values[3])
|
||||
common.Must(g.Build())
|
||||
cases := []struct {
|
||||
input string
|
||||
want []uint32
|
||||
}{
|
||||
{long, []uint32{values[0]}},
|
||||
{"www." + long, []uint32{values[0]}},
|
||||
{"x." + long, []uint32{values[1], values[0]}},
|
||||
{long[1:], nil},
|
||||
{"a" + long, nil},
|
||||
{long[:255], []uint32{values[2]}},
|
||||
{long[:254], []uint32{values[3]}},
|
||||
{long[:256], nil},
|
||||
{long[:253], nil},
|
||||
}
|
||||
for _, c := range cases {
|
||||
if m := g.Match(c.input); !slices.Equal(m, c.want) {
|
||||
t.Errorf("Match(%d bytes) = %v, want %v", len(c.input), m, c.want)
|
||||
}
|
||||
if m := g.MatchAny(c.input); m != (c.want != nil) {
|
||||
t.Errorf("MatchAny(%d bytes) = %v", len(c.input), m)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// A pattern longer than 65535 bytes builds and matches: a record's length is a uvarint,
|
||||
// so the only cap was the build-time length field, now widened to uint32.
|
||||
huge := strings.Repeat("a", 70000)
|
||||
g := NewMphMatcherGroup()
|
||||
g.AddFullMatcher(FullMatcher(strings.Repeat("a", 65535)), 1)
|
||||
g.AddDomainMatcher(DomainMatcher(huge+".com"), 2)
|
||||
g.AddFullMatcher(FullMatcher("a.com"), 3)
|
||||
common.Must(g.Build())
|
||||
if !g.MatchAny(strings.Repeat("a", 65535)) || g.MatchAny(strings.Repeat("a", 65534)) {
|
||||
t.Error("wrong answer for a 65535-byte pattern")
|
||||
}
|
||||
if m := g.Match(huge + ".com"); !slices.Equal(m, []uint32{2}) {
|
||||
t.Errorf("Match(%d-byte input) = %v, want [2]", len(huge)+4, m)
|
||||
}
|
||||
if m := g.Match("x." + huge + ".com"); !slices.Equal(m, []uint32{2}) {
|
||||
t.Errorf("Match(subdomain of a %d-byte pattern) = %v, want [2]", len(huge)+4, m)
|
||||
}
|
||||
if g.MatchAny(huge) { // the 70000-byte label on its own is not a rule
|
||||
t.Error("unexpected match for the bare 70000-byte label")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMphMatcherGroupBuildOnce(t *testing.T) {
|
||||
g := NewMphMatcherGroup()
|
||||
g.AddFullMatcher(FullMatcher("a.com"), 1)
|
||||
common.Must(g.Build())
|
||||
if err := g.Build(); err == nil || !g.MatchAny("a.com") {
|
||||
t.Errorf("second Build() = %v, MatchAny(a.com) = %v", err, g.MatchAny("a.com"))
|
||||
}
|
||||
defer func() {
|
||||
if recover() == nil {
|
||||
t.Error("Add after Build did not panic")
|
||||
}
|
||||
}()
|
||||
g.AddDomainMatcher(DomainMatcher("b.com"), 2)
|
||||
}
|
||||
|
||||
@@ -2,10 +2,12 @@ package strmatcher
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"math/bits"
|
||||
"regexp"
|
||||
"regexp/syntax"
|
||||
"slices"
|
||||
"strings"
|
||||
"unicode"
|
||||
"unicode/utf8"
|
||||
|
||||
"golang.org/x/net/idna"
|
||||
@@ -75,7 +77,9 @@ func (m SubstrMatcher) Match(s string) bool {
|
||||
// RegexMatcher is an implementation of Matcher.
|
||||
type RegexMatcher struct {
|
||||
pattern *regexp.Regexp
|
||||
literals []string // every match contains all of them, longest first
|
||||
literals []string // every match contains all of them, longest first
|
||||
tail []byteSet // tail[i] holds the bytes a matching input can have i bytes before its end
|
||||
rest *byteSet // the bytes it can have further before, nil if any
|
||||
}
|
||||
|
||||
func newRegexMatcher(pattern string) (Matcher, error) {
|
||||
@@ -87,10 +91,239 @@ func newRegexMatcher(pattern string) (Matcher, error) {
|
||||
if re, err := syntax.Parse(pattern, syntax.Perl); err == nil { // same flags as regexp.Compile
|
||||
m.literals = requiredLiterals(re, nil)
|
||||
slices.SortStableFunc(m.literals, func(a, b string) int { return len(b) - len(a) })
|
||||
m.tail, m.rest = tailGuard(re)
|
||||
}
|
||||
return m, nil
|
||||
}
|
||||
|
||||
// byteSet is a set of bytes. The bytes >= 0x80 share one bit with 0x7f.
|
||||
type byteSet [4]uint32
|
||||
|
||||
func (s *byteSet) add(c byte) { c = min(c, 0x7f); s[c>>5] |= 1 << (c & 31) }
|
||||
func (s *byteSet) has(c byte) bool { c = min(c, 0x7f); return s[c>>5]&(1<<(c&31)) != 0 }
|
||||
func (s *byteSet) or(t *byteSet) {
|
||||
for i := range s {
|
||||
s[i] |= t[i]
|
||||
}
|
||||
}
|
||||
|
||||
var allBytes = byteSet{^uint32(0), ^uint32(0), ^uint32(0), ^uint32(0)}
|
||||
|
||||
// tailLen is how many positions before the end of the input tailGuard tells apart.
|
||||
const tailLen = 8
|
||||
|
||||
// tailBudget caps how many repetition steps tailGuard walks. Only nested repeats can make the
|
||||
// walk explode, so only they are charged: a flat pattern, however long, is walked once and keeps
|
||||
// its guard.
|
||||
const tailBudget = 100000
|
||||
|
||||
// tailWalk is a set of positions in the input, counted in bytes before its end.
|
||||
type tailWalk struct {
|
||||
at uint32 // bit i: exactly i bytes before the end, for i < tailLen
|
||||
far bool // tailLen or more bytes before the end
|
||||
free bool // not tied to the end of the input yet
|
||||
}
|
||||
|
||||
func (w tailWalk) union(v tailWalk) tailWalk {
|
||||
return tailWalk{w.at | v.at, w.far || v.far, w.free || v.free}
|
||||
}
|
||||
|
||||
type tailBuilder struct {
|
||||
tail [tailLen]byteSet
|
||||
rest byteSet
|
||||
void bool
|
||||
work int
|
||||
}
|
||||
|
||||
// tailGuard walks re backwards from the end of the input and collects the bytes an input
|
||||
// matching re can have at each position before its end. It returns nil, nil when a branch
|
||||
// of re does not end with $ or when nested repeats push the walk past tailBudget.
|
||||
func tailGuard(re *syntax.Regexp) ([]byteSet, *byteSet) {
|
||||
var b tailBuilder
|
||||
w := b.walk(re, tailWalk{free: true})
|
||||
b.stop(w)
|
||||
if b.void {
|
||||
return nil, nil
|
||||
}
|
||||
if w.at != 0 { // a match can start here, so any bytes can come before
|
||||
for i := bits.TrailingZeros32(w.at); i < tailLen; i++ {
|
||||
b.tail[i] = allBytes
|
||||
}
|
||||
}
|
||||
if w.at != 0 || w.far {
|
||||
b.rest = allBytes
|
||||
}
|
||||
n := tailLen
|
||||
for n > 0 && b.tail[n-1] == b.rest {
|
||||
n--
|
||||
}
|
||||
var tail []byteSet
|
||||
if n > 0 {
|
||||
tail = slices.Clone(b.tail[:n])
|
||||
}
|
||||
if b.rest != allBytes {
|
||||
rest := b.rest
|
||||
return tail, &rest
|
||||
}
|
||||
return tail, nil
|
||||
}
|
||||
|
||||
// stop ends the paths of w. One that never met $ lets its match be followed by anything.
|
||||
func (b *tailBuilder) stop(w tailWalk) {
|
||||
if w.free {
|
||||
b.void = true
|
||||
}
|
||||
}
|
||||
|
||||
func (b *tailBuilder) walk(re *syntax.Regexp, w tailWalk) tailWalk {
|
||||
if w == (tailWalk{}) || b.void {
|
||||
return w
|
||||
}
|
||||
switch re.Op {
|
||||
case syntax.OpNoMatch:
|
||||
return tailWalk{}
|
||||
case syntax.OpLiteral:
|
||||
for i := len(re.Rune) - 1; i >= 0; i-- {
|
||||
var set byteSet
|
||||
set.add(byte(min(re.Rune[i], utf8.RuneSelf)))
|
||||
if re.Flags&syntax.FoldCase != 0 {
|
||||
for f := unicode.SimpleFold(re.Rune[i]); f != re.Rune[i]; f = unicode.SimpleFold(f) {
|
||||
set.add(byte(min(f, utf8.RuneSelf)))
|
||||
}
|
||||
}
|
||||
w = b.step(w, &set)
|
||||
}
|
||||
return w
|
||||
case syntax.OpCharClass:
|
||||
var set byteSet
|
||||
for i := 0; i+1 < len(re.Rune); i += 2 {
|
||||
for r := min(re.Rune[i], utf8.RuneSelf); r <= min(re.Rune[i+1], utf8.RuneSelf); r++ {
|
||||
set.add(byte(r))
|
||||
}
|
||||
}
|
||||
return b.step(w, &set)
|
||||
case syntax.OpAnyChar, syntax.OpAnyCharNotNL: // a domain has no \n to reject
|
||||
return b.step(w, &allBytes)
|
||||
case syntax.OpBeginText: // nothing comes before
|
||||
b.stop(w)
|
||||
return tailWalk{}
|
||||
case syntax.OpEndText:
|
||||
out := tailWalk{at: w.at & 1}
|
||||
if w.free {
|
||||
out.at = 1
|
||||
}
|
||||
return out
|
||||
case syntax.OpCapture:
|
||||
return b.walk(re.Sub[0], w)
|
||||
case syntax.OpConcat:
|
||||
for i := len(re.Sub) - 1; i >= 0; i-- {
|
||||
w = b.walk(re.Sub[i], w)
|
||||
}
|
||||
return w
|
||||
case syntax.OpAlternate:
|
||||
var out tailWalk
|
||||
for _, sub := range re.Sub {
|
||||
out = out.union(b.walk(sub, w))
|
||||
}
|
||||
return out
|
||||
case syntax.OpQuest:
|
||||
return b.repeat(re.Sub[0], w, 1)
|
||||
case syntax.OpStar:
|
||||
return b.repeat(re.Sub[0], w, -1)
|
||||
case syntax.OpPlus:
|
||||
return b.repeat(re.Sub[0], b.walk(re.Sub[0], w), -1)
|
||||
case syntax.OpRepeat:
|
||||
for i := 0; i < re.Min; i++ {
|
||||
if b.charge() {
|
||||
return w
|
||||
}
|
||||
w = b.walk(re.Sub[0], w)
|
||||
}
|
||||
if re.Max < 0 {
|
||||
return b.repeat(re.Sub[0], w, -1)
|
||||
}
|
||||
return b.repeat(re.Sub[0], w, re.Max-re.Min)
|
||||
}
|
||||
return w // empty match, line and word boundaries: no constraint
|
||||
}
|
||||
|
||||
// charge counts one repetition step and reports whether the walk has run out of budget. Only
|
||||
// repeats re-walk their body, so charging them alone bounds the blow-up of nested repeats while
|
||||
// leaving a single linear pass, of any length, free.
|
||||
func (b *tailBuilder) charge() bool {
|
||||
b.work++
|
||||
if b.work > tailBudget {
|
||||
b.void = true
|
||||
}
|
||||
return b.void
|
||||
}
|
||||
|
||||
// repeat walks back over up to n more repetitions of re, any number if n < 0.
|
||||
func (b *tailBuilder) repeat(re *syntax.Regexp, w tailWalk, n int) tailWalk {
|
||||
for ; n != 0; n-- {
|
||||
if b.charge() {
|
||||
return w
|
||||
}
|
||||
next := w.union(b.walk(re, w))
|
||||
if next == w {
|
||||
break
|
||||
}
|
||||
w = next
|
||||
}
|
||||
return w
|
||||
}
|
||||
|
||||
// step walks back over one character whose last byte is in set. A character that can be
|
||||
// non-ASCII can take up to 4 bytes, all >= 0x80; regexp matches an invalid byte as U+FFFD.
|
||||
func (b *tailBuilder) step(w tailWalk, set *byteSet) tailWalk {
|
||||
out := tailWalk{far: w.far, free: w.free}
|
||||
if w.far {
|
||||
b.rest.or(set)
|
||||
}
|
||||
width := 1
|
||||
if set.has(0x80) {
|
||||
width = utf8.UTFMax
|
||||
}
|
||||
for i := 0; i < tailLen; i++ {
|
||||
if w.at&(1<<i) == 0 {
|
||||
continue
|
||||
}
|
||||
b.tail[i].or(set)
|
||||
for n := 1; n <= width; n++ {
|
||||
if j := i + n; j < tailLen {
|
||||
out.at |= 1 << j
|
||||
if n < width {
|
||||
b.tail[j].add(0x80)
|
||||
}
|
||||
} else {
|
||||
out.far = true
|
||||
if n < width {
|
||||
b.rest.add(0x80)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// mayMatch reports whether s passes the tail guard.
|
||||
func (m *RegexMatcher) mayMatch(s string) bool {
|
||||
n := len(s)
|
||||
if m.rest == nil {
|
||||
n = min(n, len(m.tail))
|
||||
}
|
||||
for i := 0; i < n; i++ {
|
||||
set := m.rest
|
||||
if i < len(m.tail) {
|
||||
set = &m.tail[i]
|
||||
}
|
||||
if !set.has(s[len(s)-1-i]) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// requiredLiterals appends to dst the case-sensitive strings that every match of re contains.
|
||||
func requiredLiterals(re *syntax.Regexp, dst []string) []string {
|
||||
switch re.Op {
|
||||
@@ -126,6 +359,9 @@ func (m *RegexMatcher) String() string {
|
||||
}
|
||||
|
||||
func (m *RegexMatcher) Match(s string) bool {
|
||||
if !m.mayMatch(s) {
|
||||
return false
|
||||
}
|
||||
for _, l := range m.literals {
|
||||
if !strings.Contains(s, l) {
|
||||
return false
|
||||
|
||||
@@ -1,9 +1,16 @@
|
||||
package strmatcher
|
||||
|
||||
import (
|
||||
"hash/fnv"
|
||||
"math/rand/v2"
|
||||
"regexp"
|
||||
"regexp/syntax"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"unicode"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
var regexLiteralCases = []struct {
|
||||
@@ -37,6 +44,147 @@ func TestRegexRequiredLiterals(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
var regexTailCases = []struct {
|
||||
pattern string
|
||||
guard bool
|
||||
match []string // inputs the pattern matches
|
||||
reject []string // inputs the tail guard alone rejects
|
||||
}{
|
||||
{`^[a-z]([a-z0-9-]{0,61}[a-z0-9])?$`, true, []string{"a", "localhost", "x-1"}, []string{"www.example.com", "localhost.", "LOCALHOST", "a b"}},
|
||||
{`(^|\.)[a-z][1-9][0-9][a-z]\.com$`, true, []string{"a12b.com", "x.q10z.com"}, []string{"google.com", "a12b.co", "a12b.com.", "ab12.com"}},
|
||||
{`^hses[1-7]?\.akamaized\.net$`, true, []string{"hses.akamaized.net", "hses3.akamaized.net"}, []string{"xhses.akamaized.net", "www.hses.akamaized.net"}},
|
||||
{`(?i)k\.net$`, true, []string{"k.net", "K.NET", "\u212a.net"}, []string{"x.net", "k.nex"}},
|
||||
{`[^.]+\.cn$`, true, []string{"a.cn", "\xff.cn", "\u4e2d.cn"}, []string{"a.cnn", "a.c"}},
|
||||
{`\x{FFFD}$`, true, []string{"\xff", "a\xc3", "\uFFFD"}, []string{"a", "\xff."}},
|
||||
{`^.\.cn$`, true, []string{"a.cn", "\u4E2D.cn", "\xff.cn"}, []string{"ab.cn"}},
|
||||
{`^$`, true, []string{""}, []string{"a"}},
|
||||
{`(^|\.)youyuapi\..+$`, false, []string{"youyuapi.com"}, nil},
|
||||
{`abc`, false, []string{"abc", "xabcx"}, nil},
|
||||
{`^ab`, false, []string{"ab", "abc"}, nil},
|
||||
{`a$|b`, false, []string{"a", "bx"}, nil},
|
||||
{`(?m)a$`, false, []string{"a", "a\nb"}, nil},
|
||||
{strings.Repeat(`(?:abcdefgh(?:a`, 20) + strings.Repeat(`)*)*`, 20) + `\.com$`, false, []string{".com", "abcdefgha.com"}, nil}, // over tailBudget
|
||||
}
|
||||
|
||||
func TestRegexTailGuard(t *testing.T) {
|
||||
for _, test := range regexTailCases {
|
||||
m, err := newRegexMatcher(test.pattern)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rm := m.(*RegexMatcher)
|
||||
if guard := rm.tail != nil || rm.rest != nil; guard != test.guard {
|
||||
t.Errorf("%s: guard %v, want %v", test.pattern, guard, test.guard)
|
||||
}
|
||||
for _, s := range test.match {
|
||||
if !rm.pattern.MatchString(s) || !rm.Match(s) {
|
||||
t.Errorf("%s: %q does not match", test.pattern, s)
|
||||
}
|
||||
}
|
||||
for _, s := range test.reject {
|
||||
if rm.pattern.MatchString(s) || rm.mayMatch(s) {
|
||||
t.Errorf("%s: %q passes the guard", test.pattern, s)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestRegexTailGuardFlatAlternation checks that a long but non-recursive pattern keeps its
|
||||
// guard. Only nested repeats are charged against tailBudget, so a flat alternation of many
|
||||
// names, however large, is walked once and guarded; its guard is checked against regexp.
|
||||
func TestRegexTailGuardFlatAlternation(t *testing.T) {
|
||||
var sb strings.Builder
|
||||
sb.WriteString("(?:")
|
||||
for i := 0; i < 20000; i++ {
|
||||
if i > 0 {
|
||||
sb.WriteByte('|')
|
||||
}
|
||||
sb.WriteString("name")
|
||||
sb.WriteString(strconv.Itoa(i))
|
||||
}
|
||||
sb.WriteString(`)\.example\.com$`)
|
||||
m, err := newRegexMatcher(sb.String())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rm := m.(*RegexMatcher)
|
||||
if rm.tail == nil && rm.rest == nil {
|
||||
t.Fatal("flat alternation of 20000 names lost its guard")
|
||||
}
|
||||
for _, s := range []string{"name0.example.com", "name19999.example.com", "x.name12345.example.com"} {
|
||||
if !rm.pattern.MatchString(s) || !rm.Match(s) {
|
||||
t.Errorf("%q should match", s)
|
||||
}
|
||||
}
|
||||
for _, s := range []string{"name0.example.org", "name0.example.com.", "name0.example.con", "google.com"} {
|
||||
if rm.pattern.MatchString(s) {
|
||||
t.Fatalf("test bug: %q matches the pattern", s)
|
||||
}
|
||||
if rm.mayMatch(s) {
|
||||
t.Errorf("%q should be rejected by the guard", s)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// sampleMatch appends a string that re matches, assertions aside, unless it runs out of
|
||||
// budget, which it spends one per call so that nested repeats stay cheap.
|
||||
func sampleMatch(sb *strings.Builder, re *syntax.Regexp, rnd *rand.Rand, budget *int) {
|
||||
if *budget <= 0 {
|
||||
return
|
||||
}
|
||||
*budget--
|
||||
switch re.Op {
|
||||
case syntax.OpLiteral:
|
||||
for _, r := range re.Rune {
|
||||
if re.Flags&syntax.FoldCase != 0 {
|
||||
for n := rnd.IntN(4); n > 0; n-- {
|
||||
r = unicode.SimpleFold(r)
|
||||
}
|
||||
}
|
||||
sampleRune(sb, r, rnd)
|
||||
}
|
||||
case syntax.OpCharClass:
|
||||
if len(re.Rune) > 0 {
|
||||
i := rnd.IntN(len(re.Rune)/2) * 2
|
||||
sampleRune(sb, re.Rune[i]+rnd.Int32N(min(re.Rune[i+1]-re.Rune[i]+1, 300)), rnd)
|
||||
}
|
||||
case syntax.OpAnyChar, syntax.OpAnyCharNotNL:
|
||||
sampleRune(sb, []rune{'a', '.', '\n', 0xe9, 0x212a, utf8.RuneError}[rnd.IntN(6)], rnd)
|
||||
case syntax.OpCapture:
|
||||
sampleMatch(sb, re.Sub[0], rnd, budget)
|
||||
case syntax.OpConcat:
|
||||
for _, sub := range re.Sub {
|
||||
sampleMatch(sb, sub, rnd, budget)
|
||||
}
|
||||
case syntax.OpAlternate:
|
||||
sampleMatch(sb, re.Sub[rnd.IntN(len(re.Sub))], rnd, budget)
|
||||
case syntax.OpQuest, syntax.OpStar, syntax.OpPlus, syntax.OpRepeat:
|
||||
lo, hi := 0, 3
|
||||
switch re.Op {
|
||||
case syntax.OpQuest:
|
||||
hi = 1
|
||||
case syntax.OpPlus:
|
||||
lo = 1
|
||||
case syntax.OpRepeat:
|
||||
lo, hi = re.Min, re.Min+3
|
||||
if re.Max >= 0 {
|
||||
hi = min(hi, re.Max)
|
||||
}
|
||||
}
|
||||
for n := lo + rnd.IntN(hi-lo+1); n > 0; n-- {
|
||||
sampleMatch(sb, re.Sub[0], rnd, budget)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func sampleRune(sb *strings.Builder, r rune, rnd *rand.Rand) {
|
||||
if r == utf8.RuneError && rnd.IntN(2) == 0 {
|
||||
sb.WriteByte(0x80 | byte(rnd.IntN(0x80))) // regexp matches an invalid byte as U+FFFD
|
||||
return
|
||||
}
|
||||
sb.WriteRune(r)
|
||||
}
|
||||
|
||||
func FuzzRegexMatcher(f *testing.F) {
|
||||
inputs := []string{
|
||||
"", "x", "yy", "abd", "ccc", "ABC", "abCDef", "abcdef", "abababcc", "a.b", "a\xffb", "a\uFFFDb",
|
||||
@@ -47,14 +195,39 @@ func FuzzRegexMatcher(f *testing.F) {
|
||||
f.Add(test.pattern, s)
|
||||
}
|
||||
}
|
||||
for _, test := range regexTailCases {
|
||||
for _, s := range append(test.match, test.reject...) {
|
||||
f.Add(test.pattern, s)
|
||||
}
|
||||
}
|
||||
f.Fuzz(func(t *testing.T, pattern, s string) {
|
||||
re, err := regexp.Compile(pattern)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
m, _ := newRegexMatcher(pattern)
|
||||
if got, want := m.Match(s), re.MatchString(s); got != want {
|
||||
t.Errorf("pattern %q, input %q: got %v, want %v", pattern, s, got, want)
|
||||
check := func(s string) {
|
||||
if got, want := m.Match(s), re.MatchString(s); got != want {
|
||||
t.Errorf("pattern %q, input %q: got %v, want %v", pattern, s, got, want)
|
||||
}
|
||||
}
|
||||
check(s)
|
||||
// random inputs seldom match, so also try strings built from the pattern
|
||||
parsed, _ := syntax.Parse(pattern, syntax.Perl)
|
||||
h := fnv.New64a()
|
||||
h.Write([]byte(s))
|
||||
rnd := rand.New(rand.NewPCG(h.Sum64(), 1))
|
||||
for range 8 {
|
||||
var sb strings.Builder
|
||||
budget := 256
|
||||
sampleMatch(&sb, parsed, rnd, &budget)
|
||||
sample := sb.String()
|
||||
check(sample)
|
||||
check(s + sample)
|
||||
if len(sample) > 0 && len(s) > 0 {
|
||||
i := rnd.IntN(len(sample))
|
||||
check(sample[:i] + s[:1] + sample[i+1:])
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -46,7 +46,9 @@ func (g *MphValueMatcher) Add(matcher Matcher, value uint32) {
|
||||
func (g *MphValueMatcher) Build() error {
|
||||
if g.mph != nil {
|
||||
runtime.GC() // peak mem
|
||||
g.mph.Build()
|
||||
if err := g.mph.Build(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
runtime.GC() // peak mem
|
||||
if g.ac != nil {
|
||||
@@ -58,23 +60,17 @@ func (g *MphValueMatcher) Build() error {
|
||||
|
||||
// Match implements ValueMatcher.Match.
|
||||
func (g *MphValueMatcher) Match(input string) []uint32 {
|
||||
result := make([][]uint32, 0, 5)
|
||||
var result []uint32
|
||||
if g.mph != nil {
|
||||
if matches := g.mph.Match(input); len(matches) > 0 {
|
||||
result = append(result, matches)
|
||||
}
|
||||
result = g.mph.Match(input) // a new slice, returned without another copy
|
||||
}
|
||||
if g.ac != nil {
|
||||
if matches := g.ac.Match(input); len(matches) > 0 {
|
||||
result = append(result, matches)
|
||||
}
|
||||
result = append(result, g.ac.Match(input)...)
|
||||
}
|
||||
if g.regex != nil {
|
||||
if matches := g.regex.Match(input); len(matches) > 0 {
|
||||
result = append(result, matches)
|
||||
}
|
||||
result = append(result, g.regex.Match(input)...)
|
||||
}
|
||||
return CompositeMatches(result)
|
||||
return result
|
||||
}
|
||||
|
||||
// MatchAny implements ValueMatcher.MatchAny.
|
||||
@@ -87,3 +83,62 @@ func (g *MphValueMatcher) MatchAny(input string) bool {
|
||||
}
|
||||
return g.regex != nil && g.regex.MatchAny(input)
|
||||
}
|
||||
|
||||
func (g *MphValueMatcher) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
|
||||
if g.mph != nil && g.mph.matchAnyHashed(input, parents, h, mul) {
|
||||
return true
|
||||
}
|
||||
if g.ac != nil && g.ac.MatchAny(input) {
|
||||
return true
|
||||
}
|
||||
return g.regex != nil && g.regex.MatchAny(input)
|
||||
}
|
||||
|
||||
// MphValueMatcherCombiner combines several built MphValueMatchers, each bound to one value, and matches an input
|
||||
// against them as their MatchAny would, hashing the input once for all of them.
|
||||
type MphValueMatcherCombiner struct {
|
||||
matchers []*MphValueMatcher
|
||||
values []uint32
|
||||
}
|
||||
|
||||
// Add adds a built matcher that stands for value.
|
||||
func (s *MphValueMatcherCombiner) Add(m *MphValueMatcher, value uint32) {
|
||||
s.matchers = append(s.matchers, m)
|
||||
s.values = append(s.values, value)
|
||||
}
|
||||
|
||||
// Match returns the values of the matchers that match input, in Add order.
|
||||
func (s *MphValueMatcherCombiner) Match(input string) []uint32 {
|
||||
if len(s.matchers) == 0 {
|
||||
return nil
|
||||
}
|
||||
var stack [16]mphSuffix
|
||||
mul := mphMultipliers[0]
|
||||
parents, h := mphSuffixes(stack[:0], mul, input)
|
||||
var result []uint32
|
||||
for i, m := range s.matchers {
|
||||
if m.matchAnyHashed(input, parents, h, mul) {
|
||||
result = append(result, s.values[i])
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// MatchAny returns true as soon as one matcher matches input.
|
||||
func (s *MphValueMatcherCombiner) MatchAny(input string) bool {
|
||||
switch len(s.matchers) {
|
||||
case 0:
|
||||
return false
|
||||
case 1:
|
||||
return s.matchers[0].MatchAny(input) // nothing to share, and it stops at the first matching suffix
|
||||
}
|
||||
var stack [16]mphSuffix
|
||||
mul := mphMultipliers[0]
|
||||
parents, h := mphSuffixes(stack[:0], mul, input)
|
||||
for _, m := range s.matchers {
|
||||
if m.matchAnyHashed(input, parents, h, mul) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user