mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-04 14:56:49 +00:00
https://github.com/XTLS/Xray-core/pull/6821#issuecomment-5860086534
232 lines
8.4 KiB
Go
232 lines
8.4 KiB
Go
package strmatcher
|
|
|
|
import (
|
|
"errors"
|
|
"math"
|
|
"math/bits"
|
|
"runtime"
|
|
"sort"
|
|
"strings"
|
|
"unsafe"
|
|
)
|
|
|
|
// PrimeRK is the prime base used in Rabin-Karp algorithm.
|
|
const PrimeRK = 16777619
|
|
|
|
// RollingHash calculates the rolling murmurHash of given string based on a provided suffix hash.
|
|
func RollingHash(hash uint32, input string) uint32 {
|
|
for i := len(input) - 1; i >= 0; i-- {
|
|
hash = hash*PrimeRK + uint32(input[i])
|
|
}
|
|
return hash
|
|
}
|
|
|
|
// MemHash is the hash function used by go map, it utilizes available hardware instructions(behaves
|
|
// as aeshash if aes instruction is available).
|
|
// With different seed, each MemHash<seed> performs as distinct hash functions.
|
|
func MemHash(seed uint32, input string) uint32 {
|
|
return uint32(strhash(unsafe.Pointer(&input), uintptr(seed))) // nosemgrep
|
|
}
|
|
|
|
const (
|
|
mphMatchTypeCount = 2 // Full and Domain
|
|
)
|
|
|
|
type mphRuleInfo struct {
|
|
rollingHash uint32
|
|
matchers [mphMatchTypeCount][]uint32
|
|
}
|
|
|
|
// MphMatcherGroup is an implementation of MatcherGroup.
|
|
// It implements Rabin-Karp algorithm and minimal perfect hash table for Full and Domain matcher.
|
|
type MphMatcherGroup struct {
|
|
patterns string // All rule patterns concatenated
|
|
patternOffs []uint32 // RuleIdx -> patterns[patternOffs[i]:patternOffs[i+1]], index 0 reserved for failed lookup
|
|
values []uint32 // All registered matcher values concatenated
|
|
valueOffs []uint32 // RuleIdx -> values[valueOffs[i]:valueOffs[i+1]] (Full Matcher takes precedence)
|
|
level0 []uint32 // RollingHash & Mask -> seed for Memhash
|
|
level0Mask uint32 // Mask restricting RollingHash to 0 ~ len(level0)
|
|
level1 []uint32 // Memhash<seed> & Mask -> stored index for rules
|
|
level1Mask uint32 // Mask for restricting Memhash<seed> to 0 ~ len(level1)
|
|
rules []string // RuleIdx -> pattern string, only used for building
|
|
ruleInfos *map[string]mphRuleInfo
|
|
}
|
|
|
|
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
|
|
}
|
|
}
|
|
|
|
// AddFullMatcher implements MatcherGroupForFull.
|
|
func (g *MphMatcherGroup) AddFullMatcher(matcher FullMatcher, value uint32) {
|
|
pattern := strings.ToLower(matcher.Pattern())
|
|
g.addPattern(0, "", pattern, matcher.Type(), 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
|
|
}
|
|
|
|
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)
|
|
}
|
|
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) 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])
|
|
}
|
|
g.patterns = strings.Join(g.rules, "")
|
|
if uint64(len(g.patterns)) > math.MaxUint32 || uint64(valueCount) > 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))
|
|
}
|
|
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
|
|
}
|
|
sort.Slice(bucketIdxs, func(i, j int) bool { return len(buckets[bucketIdxs[i]]) > len(buckets[bucketIdxs[j]]) })
|
|
|
|
// 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)
|
|
}
|
|
}
|
|
g.level0[bucketIdx] = seed // Displacement value for this bucket
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (g *MphMatcherGroup) pattern(ruleIdx uint32) string {
|
|
return g.patterns[g.patternOffs[ruleIdx]:g.patternOffs[ruleIdx+1]]
|
|
}
|
|
|
|
// 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]
|
|
}
|
|
|
|
// 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
|
|
}
|
|
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))
|
|
}
|
|
}
|
|
}
|
|
if mphIdx := g.Lookup(hash, input); mphIdx != 0 {
|
|
matches = append(matches, g.valuesOf(mphIdx))
|
|
}
|
|
return CompositeMatchesReverse(matches)
|
|
}
|
|
|
|
// MatchAny implements MatcherGroup.MatchAny.
|
|
func (g *MphMatcherGroup) MatchAny(input string) bool {
|
|
hash := uint32(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
|
|
}
|
|
}
|
|
}
|
|
return g.Lookup(hash, input) != 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
|