mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-05 07:13:36 +00:00
173 lines
5.7 KiB
Go
173 lines
5.7 KiB
Go
package geodata
|
|
|
|
import (
|
|
"fmt"
|
|
"testing"
|
|
|
|
"github.com/xtls/xray-core/common/net"
|
|
lua "github.com/yuin/gopher-lua"
|
|
)
|
|
|
|
func TestLuaIPMatcher(t *testing.T) {
|
|
L := lua.NewState()
|
|
defer L.Close()
|
|
RegisterLua(L)
|
|
ip := L.NewUserData()
|
|
ip.Value = net.ParseIP("127.0.0.1")
|
|
L.SetGlobal("ip", ip)
|
|
ips := L.NewUserData()
|
|
ips.Value = []net.IP{ip.Value.(net.IP), net.ParseIP("8.8.8.8")}
|
|
L.SetGlobal("ips", ips)
|
|
if err := L.DoString(`
|
|
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8", "::1")
|
|
assert(matcher:Match(ip))
|
|
assert(matcher:AnyMatch(ips))
|
|
assert(not matcher:Matches(ips))
|
|
local matched, unmatched = matcher:FilterIPs(ips)
|
|
assert(type(matched) == "userdata" and type(unmatched) == "userdata")
|
|
assert(#matched == 1 and #unmatched == 1)
|
|
`); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
|
|
func TestLuaDomainMatcher(t *testing.T) {
|
|
L := lua.NewState()
|
|
defer L.Close()
|
|
RegisterLua(L)
|
|
if err := L.DoString(`
|
|
local matcher = require("xray.geodata").BuildDomainMatcher("example.com", "full:other.com")
|
|
assert(matcher:MatchAny("example.com"))
|
|
assert(matcher:MatchAny("www.example.com"))
|
|
assert(matcher:MatchAny("other.com"))
|
|
assert(not matcher:MatchAny("www.other.com"))
|
|
assert(#(matcher:Match("www.example.com")) == 1)
|
|
`); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
|
|
func TestLuaMatchersRejectInvalidRules(t *testing.T) {
|
|
for _, tc := range []struct {
|
|
name string
|
|
script string
|
|
}{
|
|
{"IP rule", `require("xray.geodata").BuildIPMatcher("not-an-ip")`},
|
|
{"non-string domain rule", `require("xray.geodata").BuildDomainMatcher("example.com", true)`},
|
|
} {
|
|
t.Run(tc.name, func(t *testing.T) {
|
|
L := lua.NewState()
|
|
defer L.Close()
|
|
RegisterLua(L)
|
|
if err := L.DoString(tc.script); err == nil {
|
|
t.Fatal("invalid geodata rule was accepted")
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestLuaMatcherArgumentsAndAliases(t *testing.T) {
|
|
L := lua.NewState()
|
|
defer L.Close()
|
|
RegisterLua(L)
|
|
ip := L.NewUserData()
|
|
ip.Value = net.ParseIP("127.0.0.1")
|
|
L.SetGlobal("ip", ip)
|
|
if err := L.DoString(`
|
|
local geodata = require("xray.geodata")
|
|
local matcher = geodata.BuildIPMatcher("127.0.0.0/8")
|
|
assert(matcher.Match == matcher.match and matcher.AnyMatch == matcher.anyMatch)
|
|
assert(matcher.Matches == matcher.matches and matcher.FilterIPs == matcher.filterIPs)
|
|
assert(matcher:match(ip))
|
|
assert(matcher:anyMatch({ip}) and matcher:matches({ip}))
|
|
assert(not matcher:AnyMatch(nil))
|
|
assert(matcher:Matches(nil) == matcher:Matches({}))
|
|
local matched, unmatched = matcher:FilterIPs({ip})
|
|
assert(#matched == 1 and matched[1]:Equal(ip))
|
|
assert(matcher:AnyMatch(matched) and matcher:Matches(matched))
|
|
local filtered, excluded = matcher:filterIPs(matched)
|
|
assert(#filtered == 1 and #excluded == 0 and filtered[1]:Equal(ip))
|
|
local emptyMatched, emptyUnmatched = matcher:FilterIPs(nil)
|
|
assert(#emptyMatched == 0 and #emptyUnmatched == 0)
|
|
matcher:SetReverse(true)
|
|
assert(not matcher:Match(ip) and not matcher:AnyMatch(matched))
|
|
matcher:ToggleReverse()
|
|
assert(matcher:Match(ip) and matcher:AnyMatch(matched))
|
|
assert(matcher.missing == nil)
|
|
|
|
local domain = geodata.BuildDomainMatcher("full:example.com")
|
|
assert(domain.Match == domain.match and domain.MatchAny == domain.matchAny)
|
|
assert(domain:matchAny("example.com"))
|
|
assert(#domain:Match("example.com") == 1)
|
|
assert(domain:match("example.com")[1] == 0)
|
|
assert(not pcall(function() matcher:AnyMatch() end))
|
|
assert(not pcall(function() matcher:AnyMatch(matched, true) end))
|
|
assert(not pcall(function() matcher.AnyMatch(ip, matched) end))
|
|
assert(not pcall(function() matcher:Match(true) end))
|
|
assert(not pcall(function() domain:MatchAny(123) end))
|
|
assert(not pcall(function() domain:MatchAny("example.com", true) end))
|
|
assert(not pcall(function() matcher:FilterIPs(true) end))
|
|
assert(not pcall(function() matcher:FilterIPs(matched, true) end))
|
|
assert(not pcall(function() domain:Match(123) end))
|
|
assert(not pcall(function() domain:Match("example.com", true) end))
|
|
`); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
|
|
// BenchmarkLuaMatcherCall measures repeated calls with prebuilt matchers and inputs.
|
|
func BenchmarkLuaMatcherCall(b *testing.B) {
|
|
L := lua.NewState()
|
|
defer L.Close()
|
|
RegisterLua(L)
|
|
ip := net.ParseIP("127.0.0.1")
|
|
for name, value := range map[string]any{"ip": ip, "ips": []net.IP{ip}} {
|
|
ud := L.NewUserData()
|
|
ud.Value = value
|
|
L.SetGlobal(name, ud)
|
|
}
|
|
if err := L.DoString(`
|
|
local geodata = require("xray.geodata")
|
|
ipMatcher = geodata.BuildIPMatcher("127.0.0.0/8")
|
|
domainMatcher = geodata.BuildDomainMatcher("full:example.com")
|
|
`); err != nil {
|
|
b.Fatal(err)
|
|
}
|
|
for _, benchmark := range []struct {
|
|
name, expression string
|
|
}{
|
|
{"ip_match", "ipMatcher:Match(ip)"},
|
|
{"ip_match_lower", "ipMatcher:match(ip)"},
|
|
{"ip_any_match", "ipMatcher:AnyMatch(ips)"},
|
|
{"ip_any_match_lower", "ipMatcher:anyMatch(ips)"},
|
|
{"ip_matches", "ipMatcher:Matches(ips)"},
|
|
{"ip_matches_lower", "ipMatcher:matches(ips)"},
|
|
{"domain_match_any", `domainMatcher:MatchAny("example.com")`},
|
|
{"domain_match_any_lower", `domainMatcher:matchAny("example.com")`},
|
|
{"ip_filter", "select(1, ipMatcher:FilterIPs(ips)) ~= nil"},
|
|
{"ip_filter_lower", "select(1, ipMatcher:filterIPs(ips)) ~= nil"},
|
|
{"domain_match", `#domainMatcher:Match("example.com") == 1`},
|
|
{"domain_match_lower", `#domainMatcher:match("example.com") == 1`},
|
|
{"ip_lua_table", "ipMatcher:AnyMatch({ip})"},
|
|
{"ip_lua_table_lower", "ipMatcher:anyMatch({ip})"},
|
|
} {
|
|
b.Run(benchmark.name, func(b *testing.B) {
|
|
if err := L.DoString(fmt.Sprintf("function benchmarkMatch() return %s end", benchmark.expression)); err != nil {
|
|
b.Fatal(err)
|
|
}
|
|
fn := L.GetGlobal("benchmarkMatch")
|
|
b.ReportAllocs()
|
|
b.ResetTimer()
|
|
for i := 0; i < b.N; i++ {
|
|
if err := L.CallByParam(lua.P{Fn: fn, NRet: 1, Protect: true}); err != nil {
|
|
b.Fatal(err)
|
|
}
|
|
if L.Get(-1) != lua.LTrue {
|
|
b.Fatal("matcher returned false")
|
|
}
|
|
L.Pop(1)
|
|
}
|
|
})
|
|
}
|
|
}
|