mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-05 05:18:15 +03:00
lua: unify IP slice access and cache slice metatables
This commit is contained in:
+6
-4
@@ -52,6 +52,8 @@ func luaServers(s *DNS) []luaDNSServer {
|
||||
|
||||
func registerLua(L *lua.LState, servers []luaDNSServer, client featureDNS.Client) {
|
||||
L.PreloadModule("xray.dns", func(L *lua.LState) int {
|
||||
pushIPs := xlua.NewSlicePusher[net.IP](L)
|
||||
|
||||
serverList := L.CreateTable(len(servers), 0)
|
||||
for i, client := range servers {
|
||||
server := L.CreateTable(0, 2)
|
||||
@@ -82,7 +84,7 @@ func registerLua(L *lua.LState, servers []luaDNSServer, client featureDNS.Client
|
||||
} else {
|
||||
ips, ttl, err = client.query(ctx, string(domain), option)
|
||||
}
|
||||
xlua.PushUserData(L, ips)
|
||||
pushIPs(L, ips)
|
||||
xlua.PushNumber(L, ttl)
|
||||
xlua.PushError(L, err)
|
||||
return 3
|
||||
@@ -95,14 +97,14 @@ func registerLua(L *lua.LState, servers []luaDNSServer, client featureDNS.Client
|
||||
module.RawSetString("Servers", serverList)
|
||||
}
|
||||
if client != nil {
|
||||
module.RawSetString("Query", newLuaClientQuery(L, client))
|
||||
module.RawSetString("Query", newLuaClientQuery(L, client, pushIPs))
|
||||
}
|
||||
L.Push(module)
|
||||
return 1
|
||||
})
|
||||
}
|
||||
|
||||
func newLuaClientQuery(L *lua.LState, client featureDNS.Client) *lua.LFunction {
|
||||
func newLuaClientQuery(L *lua.LState, client featureDNS.Client, pushIPs func(*lua.LState, []net.IP)) *lua.LFunction {
|
||||
return L.NewFunction(func(L *lua.LState) int {
|
||||
domain, ok := L.Get(1).(lua.LString)
|
||||
if !ok {
|
||||
@@ -119,7 +121,7 @@ func newLuaClientQuery(L *lua.LState, client featureDNS.Client) *lua.LFunction {
|
||||
return 0
|
||||
}
|
||||
ips, ttl, err := client.LookupIP(string(domain), option)
|
||||
xlua.PushUserData(L, ips)
|
||||
pushIPs(L, ips)
|
||||
xlua.PushNumber(L, ttl)
|
||||
xlua.PushError(L, err)
|
||||
return 3
|
||||
|
||||
+51
-2
@@ -136,8 +136,12 @@ local matcher = require("xray.geodata").BuildIPMatcher("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)
|
||||
assert(#ips == 2 and ips[1]:String() == "127.0.0.1" and ips[2]:String() == "8.8.8.8")
|
||||
assert(matcher:Match(ips[1]) and not matcher:Match(ips[2]))
|
||||
assert(matcher:AnyMatch(ips))
|
||||
local matched = matcher:FilterIPs(ips)
|
||||
local matched, unmatched = matcher:FilterIPs(ips)
|
||||
assert(#matched == 1 and #unmatched == 1)
|
||||
assert(matched[1]:Equal(ips[1]) and unmatched[1]:Equal(ips[2]))
|
||||
return matched, ttl, err
|
||||
end
|
||||
`); err != nil {
|
||||
@@ -167,7 +171,7 @@ func TestLuaDNSClientQuery(t *testing.T) {
|
||||
defer L.Close()
|
||||
L.SetContext(context.Background())
|
||||
geodata.RegisterLua(L)
|
||||
want := []net.IP{{127, 0, 0, 1}}
|
||||
want := []net.IP{{127, 0, 0, 1}, net.ParseIP("::1")}
|
||||
client := &luaDNSClient{lookup: func(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||
if domain != "MiXeD.Example." || !option.IPv4Enable || option.IPv6Enable || !option.FakeEnable {
|
||||
t.Fatalf("dns.Query arguments = %q, %+v", domain, option)
|
||||
@@ -181,6 +185,8 @@ local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.1")
|
||||
assert(dns.Servers == nil)
|
||||
ips, ttl, err = dns.Query("MiXeD.Example.", true, false, true)
|
||||
assert(not err and ttl == 42 and matcher:AnyMatch(ips))
|
||||
assert(#ips == 2 and ips[1]:String() == "127.0.0.1" and ips[2]:String() == "::1")
|
||||
assert(matcher:Match(ips[1]) and not matcher:Match(ips[2]))
|
||||
`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -201,6 +207,8 @@ assert(dns.Servers[1].ID == "localhost")
|
||||
serverIPs, _, serverErr = dns.Servers[1]:Query("127.0.0.1", true, false, false)
|
||||
clientIPs, _, clientErr = dns.Query("127.0.0.1", true, false, false)
|
||||
assert(not serverErr and not clientErr)
|
||||
assert(#serverIPs == 1 and #clientIPs == 1)
|
||||
assert(serverIPs[1]:String() == "127.0.0.1" and serverIPs[1]:Equal(clientIPs[1]))
|
||||
`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -212,6 +220,47 @@ assert(not serverErr and not clientErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLuaDNSQueryEmptyIPs(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
ips []net.IP
|
||||
}{
|
||||
{"nil", nil},
|
||||
{"empty", []net.IP{}},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
L := lua.NewState()
|
||||
defer L.Close()
|
||||
L.SetContext(context.Background())
|
||||
L.SetGlobal("expectNil", lua.LBool(tc.ips == nil))
|
||||
client := &luaDNSClient{lookup: func(string, featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||
return tc.ips, 0, featureDNS.ErrEmptyResponse
|
||||
}}
|
||||
registerLua(L, []luaDNSServer{{query: func(_ context.Context, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||
return client.LookupIP(domain, option)
|
||||
}}}, client)
|
||||
if err := L.DoString(`
|
||||
local dns = require("xray.dns")
|
||||
for _, query in ipairs({
|
||||
function() return dns.Servers[1]:Query("empty.example", true, false, false) end,
|
||||
function() return dns.Query("empty.example", true, false, false) end,
|
||||
}) do
|
||||
local ips, ttl, err = query()
|
||||
assert(ttl == 0 and err)
|
||||
if expectNil then
|
||||
assert(ips == nil)
|
||||
else
|
||||
assert(type(ips) == "userdata" and #ips == 0)
|
||||
assert(not pcall(function() return ips[1] end))
|
||||
end
|
||||
end
|
||||
`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
type benchmarkLuaNameServer struct {
|
||||
ips []net.IP
|
||||
}
|
||||
|
||||
+4
-3
@@ -62,6 +62,7 @@ func (r *Router) RegisterLua(L *lua.LState) {
|
||||
}
|
||||
|
||||
func registerLuaContext(L *lua.LState) {
|
||||
pushIPs := xlua.NewSlicePusher[net.IP](L)
|
||||
attributes := L.NewTypeMetatable(luaAttributesType)
|
||||
L.SetField(attributes, "__index", L.NewFunction(func(L *lua.LState) int {
|
||||
values := L.CheckUserData(1).Value.(map[string]string)
|
||||
@@ -76,15 +77,15 @@ func registerLuaContext(L *lua.LState) {
|
||||
methods := L.CreateTable(0, 4)
|
||||
L.SetFuncs(methods, map[string]lua.LGFunction{
|
||||
"GetSourceIPs": func(L *lua.LState) int {
|
||||
xlua.PushUserData(L, checkLuaContext(L).GetSourceIPs())
|
||||
pushIPs(L, checkLuaContext(L).GetSourceIPs())
|
||||
return 1
|
||||
},
|
||||
"GetTargetIPs": func(L *lua.LState) int {
|
||||
xlua.PushUserData(L, checkLuaContext(L).GetTargetIPs())
|
||||
pushIPs(L, checkLuaContext(L).GetTargetIPs())
|
||||
return 1
|
||||
},
|
||||
"GetLocalIPs": func(L *lua.LState) int {
|
||||
xlua.PushUserData(L, checkLuaContext(L).GetLocalIPs())
|
||||
pushIPs(L, checkLuaContext(L).GetLocalIPs())
|
||||
return 1
|
||||
},
|
||||
"GetAttributes": func(L *lua.LState) int {
|
||||
|
||||
@@ -81,7 +81,13 @@ function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort,
|
||||
savedContext = ctx
|
||||
sourceIPs, targetIPs, localIPs = ctx:GetSourceIPs(), ctx:GetTargetIPs(), ctx:GetLocalIPs()
|
||||
attributes = ctx:GetAttributes()
|
||||
assert(#sourceIPs == 1 and #targetIPs == 1 and #localIPs == 1)
|
||||
assert(sourceIPs[1]:String() == "127.0.0.2" and targetIPs[1]:String() == "127.0.0.3")
|
||||
assert(localIPs[1]:String() == "127.0.0.1")
|
||||
assert(matcher:Match(sourceIPs[1]) and matcher:Match(targetIPs[1]) and matcher:Match(localIPs[1]))
|
||||
assert(matcher:AnyMatch(sourceIPs) and matcher:AnyMatch(targetIPs) and matcher:AnyMatch(localIPs))
|
||||
local matched = matcher:FilterIPs(targetIPs)
|
||||
assert(#matched == 1 and matched[1]:Equal(targetIPs[1]))
|
||||
assert(attributes.key == "value" and attributes.missing == nil)
|
||||
assert(not pcall(function() attributes.key = "changed" end))
|
||||
return "out", "rule"
|
||||
@@ -119,6 +125,39 @@ assert(require("xray.router").LocalOS == expectedOS)`); err != nil {
|
||||
}
|
||||
}
|
||||
|
||||
func TestLuaRouteEmptyIPs(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
ips []net.IP
|
||||
}{
|
||||
{"nil", nil},
|
||||
{"empty", []net.IP{}},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
L := newLuaRouterState(t, `
|
||||
function HandleRoute(ctx)
|
||||
for _, name in ipairs({"GetSourceIPs", "GetTargetIPs", "GetLocalIPs"}) do
|
||||
local ips = ctx[name](ctx)
|
||||
if expectNil then
|
||||
assert(ips == nil)
|
||||
else
|
||||
assert(type(ips) == "userdata" and #ips == 0)
|
||||
assert(not pcall(function() return ips[1] end))
|
||||
end
|
||||
end
|
||||
return "out"
|
||||
end
|
||||
`)
|
||||
L.SetGlobal("expectNil", lua.LBool(tc.ips == nil))
|
||||
ctx := newLuaRouteTestContext()
|
||||
ctx.sourceIPs, ctx.targetIPs, ctx.localIPs = tc.ips, tc.ips, tc.ips
|
||||
if err := callLuaRoute(L, ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadLuaRouteResult(t *testing.T) {
|
||||
nativeErr := go_errors.New("native failure")
|
||||
for _, tc := range []struct {
|
||||
|
||||
Reference in New Issue
Block a user