From bbb8cb924b4443c7f4c6b13ce3045a71ddab342a Mon Sep 17 00:00:00 2001 From: wwqgtxx Date: Tue, 8 Sep 2026 10:27:49 +0800 Subject: [PATCH] chore: add standalone DomainSetBuilder --- component/fakeip/skipper_test.go | 16 +-- component/geodata/router/condition.go | 8 +- component/trie/domain_set.go | 73 ++++++++++-- component/trie/domain_set_test.go | 153 ++++++++++++++++++++++---- component/trie/domain_test.go | 12 +- config/config.go | 28 ++--- rules/provider/domain_strategy.go | 13 +-- 7 files changed, 241 insertions(+), 62 deletions(-) diff --git a/component/fakeip/skipper_test.go b/component/fakeip/skipper_test.go index 3ffffee9..be4eb26b 100644 --- a/component/fakeip/skipper_test.go +++ b/component/fakeip/skipper_test.go @@ -10,11 +10,11 @@ import ( ) func TestSkipper_BlackList(t *testing.T) { - tree := trie.New[struct{}]() - assert.NoError(t, tree.Insert("example.com", struct{}{})) - assert.False(t, tree.IsEmpty()) + var builder trie.DomainSetBuilder + assert.NoError(t, builder.Insert("example.com")) + assert.False(t, builder.IsEmpty()) skipper := &Skipper{ - Host: []C.DomainMatcher{tree.NewDomainSet()}, + Host: []C.DomainMatcher{builder.Build()}, } assert.True(t, skipper.ShouldSkipped("example.com")) assert.False(t, skipper.ShouldSkipped("foo.com")) @@ -22,11 +22,11 @@ func TestSkipper_BlackList(t *testing.T) { } func TestSkipper_WhiteList(t *testing.T) { - tree := trie.New[struct{}]() - assert.NoError(t, tree.Insert("example.com", struct{}{})) - assert.False(t, tree.IsEmpty()) + var builder trie.DomainSetBuilder + assert.NoError(t, builder.Insert("example.com")) + assert.False(t, builder.IsEmpty()) skipper := &Skipper{ - Host: []C.DomainMatcher{tree.NewDomainSet()}, + Host: []C.DomainMatcher{builder.Build()}, Mode: C.FilterWhiteList, } assert.False(t, skipper.ShouldSkipped("example.com")) diff --git a/component/geodata/router/condition.go b/component/geodata/router/condition.go index fb47e4a4..dd5b45bf 100644 --- a/component/geodata/router/condition.go +++ b/component/geodata/router/condition.go @@ -60,7 +60,7 @@ func (m *succinctDomainMatcher) Count() int { } func NewSuccinctMatcherGroup(domains []*Domain) (DomainMatcher, error) { - t := trie.New[struct{}]() + var builder trie.DomainSetBuilder m := &succinctDomainMatcher{ count: len(domains), } @@ -74,19 +74,19 @@ func NewSuccinctMatcherGroup(domains []*Domain) (DomainMatcher, error) { m.otherMatchers = append(m.otherMatchers, matcher) case Domain_Domain: - err := t.Insert("+."+d.Value, struct{}{}) + err := builder.Insert("+." + d.Value) if err != nil { return nil, err } case Domain_Full: - err := t.Insert(d.Value, struct{}{}) + err := builder.Insert(d.Value) if err != nil { return nil, err } } } - m.set = t.NewDomainSet() + m.set = builder.Build() return m, nil } diff --git a/component/trie/domain_set.go b/component/trie/domain_set.go index a5392542..6d66daee 100644 --- a/component/trie/domain_set.go +++ b/component/trie/domain_set.go @@ -10,6 +10,7 @@ import ( "github.com/metacubex/mihomo/common/utils" "github.com/openacid/low/bitmap" + "golang.org/x/exp/slices" ) const ( @@ -26,20 +27,78 @@ type DomainSet struct { type qElt struct{ s, e, col int } +// DomainSetBuilder incrementally collects domain patterns for a DomainSet. +// Its zero value is ready to use. +type DomainSetBuilder struct { + keys []string +} + +// Insert validates and adds a domain pattern to the builder. It accepts the +// same domain syntax as [DomainTrie.Insert]. +func (b *DomainSetBuilder) Insert(domain string) error { + parts, err := ValidAndSplitDomain(domain) + if err != nil { + return err + } + + if parts[0] == complexWildcard { + b.insert(parts[1:]) + b.insert(parts) + } else { + b.insert(parts) + } + return nil +} + +func (b *DomainSetBuilder) insert(parts []string) { + if parts[0] == dotWildcard { + parts[0] = complexWildcard + } + b.keys = append(b.keys, utils.Reverse(joinDomain(parts))) +} + +// IsEmpty reports whether the builder contains any domain paths. +func (b *DomainSetBuilder) IsEmpty() bool { + return b == nil || len(b.keys) == 0 +} + +// Reset discards all domains accumulated by the builder. +func (b *DomainSetBuilder) Reset() { + b.keys = nil +} + +// Build consumes the accumulated domains and creates an immutable DomainSet. +// The builder can be reused after Build returns. +func (b *DomainSetBuilder) Build() *DomainSet { + if b == nil { + return nil + } + keys := b.keys + b.keys = nil + return buildDomainSet(keys) +} + // NewDomainSet creates a new *DomainSet struct, from a DomainTrie. func (t *DomainTrie[T]) NewDomainSet() *DomainSet { - reserveDomains := make([]string, 0) - t.Foreach(func(domain string, data T) bool { - reserveDomains = append(reserveDomains, utils.Reverse(domain)) + keys := make([]string, 0) + t.Foreach(func(domain string, _ T) bool { + keys = append(keys, utils.Reverse(domain)) return true }) - // ensure that the same prefix is continuous - // and according to the ascending sequence of length - sort.Strings(reserveDomains) - keys := reserveDomains + return buildDomainSet(keys) +} + +func buildDomainSet(keys []string) *DomainSet { if len(keys) == 0 { return nil } + // ensure that the same prefix is continuous + // and according to the ascending sequence of length + sort.Strings(keys) + // The construction loop below consumes only one terminal key per node, so a + // duplicate would be indexed past its end and panic. + keys = slices.Compact(keys) + ss := &DomainSet{} lIdx := 0 diff --git a/component/trie/domain_set_test.go b/component/trie/domain_set_test.go index 0e9a7f22..2dbc280e 100644 --- a/component/trie/domain_set_test.go +++ b/component/trie/domain_set_test.go @@ -1,6 +1,7 @@ package trie_test import ( + "runtime" "strconv" "strings" "testing" @@ -29,6 +30,7 @@ func testDump(t *testing.T, tree *trie.DomainTrie[struct{}], set *trie.DomainSet func TestDomainSet(t *testing.T) { tree := trie.New[struct{}]() + var builder trie.DomainSetBuilder domainSet := []string{ "baidu.com", "google.com", @@ -42,9 +44,11 @@ func TestDomainSet(t *testing.T) { for _, domain := range domainSet { assert.NoError(t, tree.Insert(domain, struct{}{})) + assert.NoError(t, builder.Insert(domain)) } - assert.False(t, tree.IsEmpty()) - set := tree.NewDomainSet() + assert.False(t, builder.IsEmpty()) + set := builder.Build() + assert.Equal(t, tree.NewDomainSet(), set) assert.NotNil(t, set) assert.True(t, set.Has("test.cn")) assert.True(t, set.Has("cn")) @@ -57,8 +61,43 @@ func TestDomainSet(t *testing.T) { testDump(t, tree, set) } +func TestDomainSetBuilderLifecycle(t *testing.T) { + var builder trie.DomainSetBuilder + assert.True(t, builder.IsEmpty()) + assert.ErrorIs(t, builder.Insert("invalid..example"), trie.ErrInvalidDomain) + assert.True(t, builder.IsEmpty()) + assert.Nil(t, builder.Build()) + + for _, domain := range []string{"+.example.com", "example.com", "+.example.com"} { + assert.NoError(t, builder.Insert(domain)) + } + set := builder.Build() + assert.True(t, builder.IsEmpty()) + assert.True(t, set.Has("example.com")) + assert.True(t, set.Has("www.example.com")) + + var keys []string + set.Foreach(func(key string) bool { + keys = append(keys, key) + return true + }) + slices.Sort(keys) + assert.Equal(t, []string{"+.example.com", "example.com"}, keys) + + assert.NoError(t, builder.Insert("other.example")) + set = builder.Build() + assert.True(t, set.Has("other.example")) + assert.False(t, set.Has("example.com")) + + assert.NoError(t, builder.Insert("discard.example")) + builder.Reset() + assert.True(t, builder.IsEmpty()) + assert.Nil(t, builder.Build()) +} + func TestDomainSetComplexWildcard(t *testing.T) { tree := trie.New[struct{}]() + var builder trie.DomainSetBuilder domainSet := []string{ "+.baidu.com", "+.a.baidu.com", @@ -71,9 +110,11 @@ func TestDomainSetComplexWildcard(t *testing.T) { for _, domain := range domainSet { assert.NoError(t, tree.Insert(domain, struct{}{})) + assert.NoError(t, builder.Insert(domain)) } - assert.False(t, tree.IsEmpty()) - set := tree.NewDomainSet() + assert.False(t, builder.IsEmpty()) + set := builder.Build() + assert.Equal(t, tree.NewDomainSet(), set) assert.NotNil(t, set) assert.False(t, set.Has("google.com")) assert.True(t, set.Has("www.baidu.com")) @@ -83,6 +124,7 @@ func TestDomainSetComplexWildcard(t *testing.T) { func TestDomainSetWildcard(t *testing.T) { tree := trie.New[struct{}]() + var builder trie.DomainSetBuilder domainSet := []string{ "*.*.*.baidu.com", "www.baidu.*", @@ -94,9 +136,11 @@ func TestDomainSetWildcard(t *testing.T) { for _, domain := range domainSet { assert.NoError(t, tree.Insert(domain, struct{}{})) + assert.NoError(t, builder.Insert(domain)) } - assert.False(t, tree.IsEmpty()) - set := tree.NewDomainSet() + assert.False(t, builder.IsEmpty()) + set := builder.Build() + assert.Equal(t, tree.NewDomainSet(), set) assert.NotNil(t, set) assert.True(t, set.Has("www.baidu.com")) assert.True(t, set.Has("test.test.baidu.com")) @@ -112,10 +156,13 @@ func TestDomainSetWildcard(t *testing.T) { func TestDomainSetCase(t *testing.T) { tree := trie.New[struct{}]() - for _, domain := range []string{"example.com", "+.mixed.example.org"} { + var builder trie.DomainSetBuilder + for _, domain := range []string{"example.com", "EXAMPLE.COM", "+.mixed.example.org"} { assert.NoError(t, tree.Insert(domain, struct{}{})) + assert.NoError(t, builder.Insert(domain)) } - set := tree.NewDomainSet() + set := builder.Build() + assert.Equal(t, tree.NewDomainSet(), set) assert.NotNil(t, set) assert.True(t, set.Has("EXAMPLE.COM")) assert.True(t, set.Has("ExAmPlE.cOm")) @@ -127,10 +174,13 @@ func TestDomainSetCase(t *testing.T) { // path than the byte-wise one because the set is built with rune-wise reversal. func TestDomainSetUnicode(t *testing.T) { tree := trie.New[struct{}]() + var builder trie.DomainSetBuilder for _, domain := range []string{"中文.example", "+.测试.cn"} { assert.NoError(t, tree.Insert(domain, struct{}{})) + assert.NoError(t, builder.Insert(domain)) } - set := tree.NewDomainSet() + set := builder.Build() + assert.Equal(t, tree.NewDomainSet(), set) assert.NotNil(t, set) assert.True(t, set.Has("中文.example")) assert.True(t, set.Has("www.测试.cn")) @@ -139,24 +189,29 @@ func TestDomainSetUnicode(t *testing.T) { func TestDomainSetOversizedKey(t *testing.T) { tree := trie.New[struct{}]() - assert.NoError(t, tree.Insert("+.example.com", struct{}{})) - set := tree.NewDomainSet() + var builder trie.DomainSetBuilder + for _, domain := range []string{"+.example.com"} { + assert.NoError(t, tree.Insert(domain, struct{}{})) + assert.NoError(t, builder.Insert(domain)) + } + set := builder.Build() + assert.Equal(t, tree.NewDomainSet(), set) assert.NotNil(t, set) - var builder strings.Builder - for builder.Len() < 300 { - builder.WriteString("label.") + var keyBuilder strings.Builder + for keyBuilder.Len() < 300 { + keyBuilder.WriteString("label.") } - assert.True(t, set.Has(builder.String()+"example.com")) - assert.False(t, set.Has(builder.String()+"example.net")) + assert.True(t, set.Has(keyBuilder.String()+"example.com")) + assert.False(t, set.Has(keyBuilder.String()+"example.net")) } func BenchmarkDomainSetHas(b *testing.B) { - tree := trie.New[struct{}]() + var builder trie.DomainSetBuilder for i := 0; i < 10000; i++ { - assert.NoError(b, tree.Insert("+."+strconv.Itoa(i)+".example.com", struct{}{})) + assert.NoError(b, builder.Insert("+."+strconv.Itoa(i)+".example.com")) } - set := tree.NewDomainSet() + set := builder.Build() // Keys are split by length because the Go compiler only keeps a // non-constant sized allocation off the heap up to 32 bytes, so hostnames @@ -188,3 +243,63 @@ func BenchmarkDomainSetHas(b *testing.B) { }) } } + +func BenchmarkDomainSetBuild(b *testing.B) { + suffixes := [...]string{ + "google.com", + "github.io", + "cloudflare.net", + "mozilla.org", + "amazonaws.com", + "apple.com", + "telegram.org", + "example.co.uk", + } + domains := make([]string, 10000) + for i := range domains { + domain := strconv.Itoa(i) + "." + suffixes[i%len(suffixes)] + switch i % 20 { + case 0: + domains[i] = domain + case 1: + domains[i] = "." + domain + case 2: + domains[i] = "*." + domain + case 3: + domains[i] = "stun.*." + domain + default: + domains[i] = "+." + domain + } + if i%50 == 49 { + domains[i] = domains[i-1] + } + } + + b.Run("via_trie", func(b *testing.B) { + b.ReportAllocs() + for i := 0; i < b.N; i++ { + tree := trie.New[struct{}]() + for _, domain := range domains { + if err := tree.Insert(domain, struct{}{}); err != nil { + b.Fatal(err) + } + } + set := tree.NewDomainSet() + runtime.KeepAlive(set) + } + }) + + b.Run("builder", func(b *testing.B) { + b.ReportAllocs() + for i := 0; i < b.N; i++ { + var builder trie.DomainSetBuilder + for _, domain := range domains { + if err := builder.Insert(domain); err != nil { + b.Fatal(err) + } + } + set := builder.Build() + runtime.KeepAlive(set) + } + }) +} diff --git a/component/trie/domain_test.go b/component/trie/domain_test.go index 213075a7..6eef9c93 100644 --- a/component/trie/domain_test.go +++ b/component/trie/domain_test.go @@ -150,11 +150,17 @@ func TestTrie_InvalidWildcardPlacement(t *testing.T) { for _, d := range valid { tree := trie.New[netip.Addr]() assert.NoError(t, tree.Insert(d, localIP)) - set := tree.NewDomainSet() + setFromTrie := tree.NewDomainSet() + var builder trie.DomainSetBuilder + assert.NoError(t, builder.Insert(d)) + setFromBuilder := builder.Build() + assert.Equal(t, setFromTrie, setFromBuilder) for _, q := range queries { searchHit := tree.Search(q) != nil - setHit := set != nil && set.Has(q) - assert.Equalf(t, searchHit, setHit, "pattern %q query %q: Search=%v Has=%v", d, q, searchHit, setHit) + trieSetHit := setFromTrie != nil && setFromTrie.Has(q) + builderSetHit := setFromBuilder != nil && setFromBuilder.Has(q) + assert.Equalf(t, searchHit, trieSetHit, "pattern %q query %q: Search=%v TrieSet=%v", d, q, searchHit, trieSetHit) + assert.Equalf(t, searchHit, builderSetHit, "pattern %q query %q: Search=%v BuilderSet=%v", d, q, searchHit, builderSetHit) } } }) diff --git a/config/config.go b/config/config.go index b1c8a656..acdf0900 100644 --- a/config/config.go +++ b/config/config.go @@ -1499,14 +1499,14 @@ func parseDNS(rawCfg *RawConfig, ruleProviders map[string]P.RuleProvider) (*DNS, } if cfg.EnhancedMode == C.DNSFakeIP { - var fakeIPTrie *trie.DomainTrie[struct{}] - if len(dnsCfg.Fallback) != 0 { - fakeIPTrie = trie.New[struct{}]() + var fakeIPDomainSetBuilder *trie.DomainSetBuilder + if cfg.FakeIPFilterMode != C.FilterRule && len(dnsCfg.Fallback) != 0 { + fakeIPDomainSetBuilder = &trie.DomainSetBuilder{} for _, fb := range dnsCfg.Fallback { if net.ParseIP(fb.Addr) != nil { continue } - if err := fakeIPTrie.Insert(fb.Addr, struct{}{}); err != nil { + if err := fakeIPDomainSetBuilder.Insert(fb.Addr); err != nil { log.Warnln("skip fallback nameserver in fake-ip filter: %s", err) } } @@ -1521,7 +1521,7 @@ func parseDNS(rawCfg *RawConfig, ruleProviders map[string]P.RuleProvider) (*DNS, } skipper.Rules = rules } else { - host, err := parseDomain(cfg.FakeIPFilter, fakeIPTrie, "dns.fake-ip-filter", ruleProviders) + host, err := parseDomain(cfg.FakeIPFilter, fakeIPDomainSetBuilder, "dns.fake-ip-filter", ruleProviders) if err != nil { return nil, err } @@ -1584,14 +1584,14 @@ func parseDNS(rawCfg *RawConfig, ruleProviders map[string]P.RuleProvider) (*DNS, dnsCfg.FallbackIPFilter = append(dnsCfg.FallbackIPFilter, matcher) } if len(cfg.FallbackFilter.Domain) > 0 { - domainTrie := trie.New[struct{}]() + var domainSetBuilder trie.DomainSetBuilder for idx, domain := range cfg.FallbackFilter.Domain { - err = domainTrie.Insert(domain, struct{}{}) + err = domainSetBuilder.Insert(domain) if err != nil { return nil, fmt.Errorf("DNS FallbackDomain[%d] format error: %w", idx, err) } } - matcher := domainTrie.NewDomainSet() // dns.fallback-filter.domain + matcher := domainSetBuilder.Build() // dns.fallback-filter.domain dnsCfg.FallbackDomainFilter = append(dnsCfg.FallbackDomainFilter, matcher) } if len(cfg.FallbackFilter.GeoSite) > 0 { @@ -1898,7 +1898,7 @@ func parseIPCIDR(addresses []string, cidrSet *cidr.IpCidrSet, adapterName string return } -func parseDomain(domains []string, domainTrie *trie.DomainTrie[struct{}], adapterName string, ruleProviders map[string]P.RuleProvider) (matchers []C.DomainMatcher, err error) { +func parseDomain(domains []string, domainSetBuilder *trie.DomainSetBuilder, adapterName string, ruleProviders map[string]P.RuleProvider) (matchers []C.DomainMatcher, err error) { var matcher C.DomainMatcher for idx, domain := range domains { domainLower := strings.ToLower(domain) @@ -1925,17 +1925,17 @@ func parseDomain(domains []string, domainTrie *trie.DomainTrie[struct{}], adapte matchers = append(matchers, matcher) } } else { - if domainTrie == nil { - domainTrie = trie.New[struct{}]() + if domainSetBuilder == nil { + domainSetBuilder = &trie.DomainSetBuilder{} } - err = domainTrie.Insert(domain, struct{}{}) + err = domainSetBuilder.Insert(domain) if err != nil { return nil, fmt.Errorf("%s[%d]: %w", adapterName, idx, err) } } } - if !domainTrie.IsEmpty() { - matcher = domainTrie.NewDomainSet() + if !domainSetBuilder.IsEmpty() { + matcher = domainSetBuilder.Build() matchers = append(matchers, matcher) } return diff --git a/rules/provider/domain_strategy.go b/rules/provider/domain_strategy.go index 14b41955..46c63d6f 100644 --- a/rules/provider/domain_strategy.go +++ b/rules/provider/domain_strategy.go @@ -14,9 +14,9 @@ import ( ) type domainStrategy struct { - count int - domainTrie *trie.DomainTrie[struct{}] - domainSet *trie.DomainSet + count int + domainSetBuilder trie.DomainSetBuilder + domainSet *trie.DomainSet } func (d *domainStrategy) Behavior() P.RuleBehavior { @@ -32,7 +32,7 @@ func (d *domainStrategy) Count() int { } func (d *domainStrategy) Reset() { - d.domainTrie = trie.New[struct{}]() + d.domainSetBuilder.Reset() d.domainSet = nil d.count = 0 } @@ -42,7 +42,7 @@ func (d *domainStrategy) Insert(rule string) { log.Warnln("skip invalid domain from rule provider: invalid domain %q: slash is not allowed", rule) return } - err := d.domainTrie.Insert(rule, struct{}{}) + err := d.domainSetBuilder.Insert(rule) if err != nil { log.Warnln("skip invalid domain from rule provider: %s", err) } else { @@ -51,8 +51,7 @@ func (d *domainStrategy) Insert(rule string) { } func (d *domainStrategy) FinishInsert() { - d.domainSet = d.domainTrie.NewDomainSet() - d.domainTrie = nil + d.domainSet = d.domainSetBuilder.Build() } func (d *domainStrategy) FromMrs(r io.Reader, count int) error {