chore: add standalone DomainSetBuilder

This commit is contained in:
wwqgtxx committed 2026-09-08 10:27:49 +08:00
1 parent 8f54283e97
commit bbb8cb924b
7 files changed
+241 -62

No files matched your search

+8 -8
View File
@@ -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"))
+4 -4
View File
@@ -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
}
+66 -7
View File
@@ -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
+134 -19
View File
@@ -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)
}
})
}
+9 -3
View File
@@ -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)
}
}
})
+14 -14
View File
@@ -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
+6 -7
View File
@@ -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 {