diff --git a/common/geodata/lua.go b/common/geodata/lua.go index a93a4b7b5..21adf8980 100644 --- a/common/geodata/lua.go +++ b/common/geodata/lua.go @@ -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 { diff --git a/common/geodata/lua_test.go b/common/geodata/lua_test.go index dbc232e5e..f66a57bc4 100644 --- a/common/geodata/lua_test.go +++ b/common/geodata/lua_test.go @@ -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) + } + }) + } +} diff --git a/common/lua/luar.go b/common/lua/luar.go new file mode 100644 index 000000000..eef71c514 --- /dev/null +++ b/common/lua/luar.go @@ -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() +}