diff --git a/common/geodata/domain_matcher.go b/common/geodata/domain_matcher.go index a4d3564e6..e7227f58d 100644 --- a/common/geodata/domain_matcher.go +++ b/common/geodata/domain_matcher.go @@ -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]()} } diff --git a/common/geodata/domain_matcher_test.go b/common/geodata/domain_matcher_test.go index 0c0c5080c..dc3558020 100644 --- a/common/geodata/domain_matcher_test.go +++ b/common/geodata/domain_matcher_test.go @@ -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() + }) + } +} diff --git a/common/geodata/geodat_loader.go b/common/geodata/geodat_loader.go index 0e12aa28b..81d9e7e90 100644 --- a/common/geodata/geodat_loader.go +++ b/common/geodata/geodat_loader.go @@ -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 } diff --git a/common/geodata/geodat_loader_test.go b/common/geodata/geodat_loader_test.go new file mode 100644 index 000000000..b6d400294 --- /dev/null +++ b/common/geodata/geodat_loader_test.go @@ -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") + } +} diff --git a/common/geodata/strmatcher/indexmatcher_mph.go b/common/geodata/strmatcher/indexmatcher_mph.go index b23f83763..23e4d71ae 100644 --- a/common/geodata/strmatcher/indexmatcher_mph.go +++ b/common/geodata/strmatcher/indexmatcher_mph.go @@ -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. diff --git a/common/geodata/strmatcher/indexmatcher_mph_test.go b/common/geodata/strmatcher/indexmatcher_mph_test.go index 2e1c70dd2..efc747a71 100644 --- a/common/geodata/strmatcher/indexmatcher_mph_test.go +++ b/common/geodata/strmatcher/indexmatcher_mph_test.go @@ -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) + } } } diff --git a/common/geodata/strmatcher/matchergroup_mph.go b/common/geodata/strmatcher/matchergroup_mph.go index 91528f26a..2db3cc9d6 100644 --- a/common/geodata/strmatcher/matchergroup_mph.go +++ b/common/geodata/strmatcher/matchergroup_mph.go @@ -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 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 & Mask -> stored index for rules - level1Mask uint32 // Mask for restricting Memhash 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 diff --git a/common/geodata/strmatcher/matchergroup_mph_internal_test.go b/common/geodata/strmatcher/matchergroup_mph_internal_test.go new file mode 100644 index 000000000..f6b320f60 --- /dev/null +++ b/common/geodata/strmatcher/matchergroup_mph_internal_test.go @@ -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) + } +} diff --git a/common/geodata/strmatcher/matchergroup_mph_test.go b/common/geodata/strmatcher/matchergroup_mph_test.go index 4def92246..acb82a72a 100644 --- a/common/geodata/strmatcher/matchergroup_mph_test.go +++ b/common/geodata/strmatcher/matchergroup_mph_test.go @@ -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) +} diff --git a/common/geodata/strmatcher/matchers.go b/common/geodata/strmatcher/matchers.go index c2f46df64..e0f752ed4 100644 --- a/common/geodata/strmatcher/matchers.go +++ b/common/geodata/strmatcher/matchers.go @@ -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< 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:]) + } } }) } diff --git a/common/geodata/strmatcher/valuematcher_mph.go b/common/geodata/strmatcher/valuematcher_mph.go index ced6f442f..f764dc4d4 100644 --- a/common/geodata/strmatcher/valuematcher_mph.go +++ b/common/geodata/strmatcher/valuematcher_mph.go @@ -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 +}