geodata: skip reflection for matcher calls

This commit is contained in:
Meo597
2026-10-05 04:01:08 +08:00
parent d74b8b1ba5
commit e73fb0d500
3 changed files with 257 additions and 2 deletions
+103 -2
View File
@@ -1,6 +1,8 @@
package geodata
import (
xlua "github.com/xtls/xray-core/common/lua"
"github.com/xtls/xray-core/common/net"
lua "github.com/yuin/gopher-lua"
luar "layeh.com/gopher-luar"
)
@@ -21,7 +23,10 @@ func RegisterLua(L *lua.LState) {
L.RaiseError("%v", err)
return 0
}
L.Push(luar.New(L, matcher))
xlua.PushWithDirectMethods(L, matcher, map[string]xlua.DirectMethod{
"Match": luaDomainMatch,
"MatchAny": luaDomainMatchAny,
})
return 1
}))
@@ -36,7 +41,12 @@ func RegisterLua(L *lua.LState) {
L.RaiseError("%v", err)
return 0
}
L.Push(luar.New(L, matcher))
xlua.PushWithDirectMethods(L, matcher, map[string]xlua.DirectMethod{
"Match": luaIPMatch,
"AnyMatch": luaIPAnyMatch,
"Matches": luaIPMatches,
"FilterIPs": luaIPFilterIPs,
})
return 1
}))
L.Push(module)
@@ -44,6 +54,97 @@ func RegisterLua(L *lua.LState) {
})
}
// Read native Go values by type assertion; slices keep their original storage.
func readLuaIPMatcherArgs[T any](L *lua.LState) (IPMatcher, T, bool) {
var input T
if L.GetTop() != 2 {
return nil, input, false
}
value, ok := L.Get(1).(*lua.LUserData)
if !ok {
return nil, input, false
}
matcher, ok := value.Value.(IPMatcher)
if !ok {
return nil, input, false
}
if L.Get(2) == lua.LNil {
return matcher, input, true
}
value, ok = L.Get(2).(*lua.LUserData)
if !ok {
return nil, input, false
}
input, ok = value.Value.(T)
return matcher, input, ok
}
func luaIPMatch(L *lua.LState) (int, bool) {
matcher, ip, ok := readLuaIPMatcherArgs[net.IP](L)
if !ok {
return 0, false
}
L.Push(lua.LBool(matcher.Match(ip)))
return 1, true
}
func luaIPAnyMatch(L *lua.LState) (int, bool) {
matcher, ips, ok := readLuaIPMatcherArgs[[]net.IP](L)
if !ok {
return 0, false
}
L.Push(lua.LBool(matcher.AnyMatch(ips)))
return 1, true
}
func luaIPMatches(L *lua.LState) (int, bool) {
matcher, ips, ok := readLuaIPMatcherArgs[[]net.IP](L)
if !ok {
return 0, false
}
L.Push(lua.LBool(matcher.Matches(ips)))
return 1, true
}
func luaIPFilterIPs(L *lua.LState) (int, bool) {
matcher, ips, ok := readLuaIPMatcherArgs[[]net.IP](L)
if !ok {
return 0, false
}
matched, unmatched := matcher.FilterIPs(ips)
L.Push(luar.New(L, matched))
L.Push(luar.New(L, unmatched))
return 2, true
}
func luaDomainMatch(L *lua.LState) (int, bool) {
if L.GetTop() == 2 {
if value, ok := L.Get(1).(*lua.LUserData); ok {
matcher, validMatcher := value.Value.(DomainMatcher)
domain, validDomain := L.Get(2).(lua.LString)
if validMatcher && validDomain {
L.Push(luar.New(L, matcher.Match(string(domain))))
return 1, true
}
}
}
return 0, false
}
func luaDomainMatchAny(L *lua.LState) (int, bool) {
if L.GetTop() == 2 {
if value, ok := L.Get(1).(*lua.LUserData); ok {
matcher, validMatcher := value.Value.(DomainMatcher)
domain, validDomain := L.Get(2).(lua.LString)
if validMatcher && validDomain {
L.Push(lua.LBool(matcher.MatchAny(string(domain))))
return 1, true
}
}
}
return 0, false
}
func luaRules(L *lua.LState) []string {
rules := make([]string, L.GetTop())
for i := range rules {
+106
View File
@@ -1,6 +1,7 @@
package geodata
import (
"fmt"
"testing"
"github.com/xtls/xray-core/common/net"
@@ -64,3 +65,108 @@ func TestLuaMatchersRejectInvalidRules(t *testing.T) {
})
}
}
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)
}
})
}
}
+48
View File
@@ -0,0 +1,48 @@
package lua
import (
glua "github.com/yuin/gopher-lua"
luar "layeh.com/gopher-luar"
)
// DirectMethod handles a Lua call without luar's reflected method invocation.
// It returns the result count and whether it handled the arguments. On false,
// it must leave the stack unchanged for the original luar wrapper.
type DirectMethod func(L *glua.LState) (nresults int, handled bool)
// PushWithDirectMethods pushes a luar userdata with typed Go method bindings.
// Handled calls bypass luar's argument conversion and reflect.Call; method lookup
// uses the methods table directly instead of luar's reflected __index handler.
// value must expose methods only. Bindings and their closures are installed once
// per Go type per LState, outside the method-call hot path.
func PushWithDirectMethods(L *glua.LState, value any, directMethods map[string]DirectMethod) {
userdata := luar.New(L, value).(*glua.LUserData)
metatable := userdata.Metatable.(*glua.LTable)
methods := metatable.RawGetString("methods").(*glua.LTable)
if metatable.RawGetString("__index") != methods {
for name, direct := range directMethods {
original := methods.RawGetString(name)
fn := L.NewFunction(func(L *glua.LState) int {
if nresults, handled := direct(L); handled {
return nresults
}
return callLuarMethod(L, original)
})
// Keep luar's method aliases on the same direct binding.
methods.ForEach(func(key, method glua.LValue) {
if method == original {
methods.RawSet(key, fn)
}
})
}
metatable.RawSetString("__index", methods)
}
L.Push(userdata)
}
func callLuarMethod(L *glua.LState, method glua.LValue) int {
nargs := L.GetTop()
L.Insert(method, 1)
L.Call(nargs, glua.MultRet)
return L.GetTop()
}