diff --git a/app/dns/lua_test.go b/app/dns/lua_test.go index 2d8fe6ff8..dac03a1e3 100644 --- a/app/dns/lua_test.go +++ b/app/dns/lua_test.go @@ -158,7 +158,7 @@ func TestLuaDNSServerQuery(t *testing.T) { server.RegisterLua(L) if err := L.DoString(` local server = require("xray.dns").servers[1] -local matcher = require("xray.geodata").ipMatcher({"127.0.0.0/8"}) +local matcher = require("xray.geodata").ipMatcher("127.0.0.0/8") function handleDNSQuery(domain, ipv4, ipv6, fake) local ips, ttl, err = server:query(domain, ipv4, ipv6, fake) assert(type(ips) == "userdata" and not err) diff --git a/app/dns/script_test.go b/app/dns/script_test.go index ef79cd575..28aaca9df 100644 --- a/app/dns/script_test.go +++ b/app/dns/script_test.go @@ -38,7 +38,7 @@ func TestDNSScriptGeoIPFallback(t *testing.T) { t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources")) script := ` local servers = require("xray.dns").servers -local us_ips = require("xray.geodata").ipMatcher({"geoip:us"}) +local us_ips = require("xray.geodata").ipMatcher("geoip:us") local by_id = {} for _, server in ipairs(servers) do diff --git a/common/geodata/lua.go b/common/geodata/lua.go index 0531cde55..021be233c 100644 --- a/common/geodata/lua.go +++ b/common/geodata/lua.go @@ -6,13 +6,12 @@ import ( ) // RegisterLua makes xray.geodata available to require in an LState. -// Matchers retain registry handles, so they remain usable after a reload. func RegisterLua(L *lua.LState) { L.PreloadModule("xray.geodata", func(L *lua.LState) int { module := L.NewTable() module.RawSetString("domainMatcher", L.NewFunction(func(L *lua.LState) int { - parsed, err := ParseDomainRules(luaRules(L, 1), Domain_Domain) + parsed, err := ParseDomainRules(luaRules(L), Domain_Domain) if err != nil { L.RaiseError("%v", err) return 0 @@ -27,7 +26,7 @@ func RegisterLua(L *lua.LState) { })) module.RawSetString("ipMatcher", L.NewFunction(func(L *lua.LState) int { - parsed, err := ParseIPRules(luaRules(L, 1)) + parsed, err := ParseIPRules(luaRules(L)) if err != nil { L.RaiseError("%v", err) return 0 @@ -45,11 +44,10 @@ func RegisterLua(L *lua.LState) { }) } -func luaRules(L *lua.LState, index int) []string { - table := L.CheckTable(index) - rules := make([]string, table.Len()) +func luaRules(L *lua.LState) []string { + rules := make([]string, L.GetTop()) for i := range rules { - value, ok := table.RawGetInt(i + 1).(lua.LString) + value, ok := L.Get(i + 1).(lua.LString) if !ok { L.RaiseError("geodata rules must be strings") return nil diff --git a/common/geodata/lua_test.go b/common/geodata/lua_test.go index 524dbfeb1..00bae2e53 100644 --- a/common/geodata/lua_test.go +++ b/common/geodata/lua_test.go @@ -7,34 +7,39 @@ import ( lua "github.com/yuin/gopher-lua" ) -func TestLuaIPMatcherAcceptsNativeIP(t *testing.T) { +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").ipMatcher({"127.0.0.0/8"}) + local matcher = require("xray.geodata").ipMatcher("127.0.0.0/8", "::1") assert(matcher:match(ip)) - assert(matcher:anyMatch({ip})) - assert(matcher:matches({ip})) - local matched, unmatched = matcher:filterIPs({ip}) - assert(#matched == 1 and #unmatched == 0) - assert(matcher:match(matched[1])) + 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 TestLuaDomainMatcherUsesNativeMatcher(t *testing.T) { +func TestLuaDomainMatcher(t *testing.T) { L := lua.NewState() defer L.Close() RegisterLua(L) if err := L.DoString(` - local matcher = require("xray.geodata").domainMatcher({"example.com"}) + local matcher = require("xray.geodata").domainMatcher("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) @@ -46,8 +51,8 @@ func TestLuaMatchersRejectInvalidRules(t *testing.T) { name string script string }{ - {"IP rule", `require("xray.geodata").ipMatcher({"not-an-ip"})`}, - {"non-string domain rule", `require("xray.geodata").domainMatcher({true})`}, + {"IP rule", `require("xray.geodata").ipMatcher("not-an-ip")`}, + {"non-string domain rule", `require("xray.geodata").domainMatcher("example.com", true)`}, } { t.Run(tc.name, func(t *testing.T) { L := lua.NewState()