mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-05 15:23:33 +00: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 {
|
||||
|
||||
+32
-33
@@ -4,20 +4,6 @@ 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"
|
||||
)
|
||||
|
||||
var (
|
||||
luaDomainDirectMethods = map[string]xlua.DirectMethod{
|
||||
"Match": luaDomainMatch,
|
||||
"MatchAny": luaDomainMatchAny,
|
||||
}
|
||||
luaIPDirectMethods = map[string]xlua.DirectMethod{
|
||||
"Match": luaIPMatch,
|
||||
"AnyMatch": luaIPAnyMatch,
|
||||
"Matches": luaIPMatches,
|
||||
"FilterIPs": luaIPFilterIPs,
|
||||
}
|
||||
)
|
||||
|
||||
// RegisterLua makes xray.geodata available to require in an LState.
|
||||
@@ -36,7 +22,10 @@ func RegisterLua(L *lua.LState) {
|
||||
L.RaiseError("%v", err)
|
||||
return 0
|
||||
}
|
||||
xlua.PushWithDirectMethods(L, matcher, luaDomainDirectMethods)
|
||||
xlua.PushWithDirectMethods(L, matcher, map[string]xlua.DirectMethod{
|
||||
"Match": newLuaDomainMatch(xlua.NewSlicePusher[uint32](L)),
|
||||
"MatchAny": luaDomainMatchAny,
|
||||
})
|
||||
return 1
|
||||
}))
|
||||
|
||||
@@ -51,9 +40,15 @@ func RegisterLua(L *lua.LState) {
|
||||
L.RaiseError("%v", err)
|
||||
return 0
|
||||
}
|
||||
xlua.PushWithDirectMethods(L, matcher, luaIPDirectMethods)
|
||||
xlua.PushWithDirectMethods(L, matcher, map[string]xlua.DirectMethod{
|
||||
"Match": luaIPMatch,
|
||||
"AnyMatch": luaIPAnyMatch,
|
||||
"Matches": luaIPMatches,
|
||||
"FilterIPs": newLuaIPFilterIPs(xlua.NewSlicePusher[net.IP](L)),
|
||||
})
|
||||
return 1
|
||||
}))
|
||||
|
||||
L.Push(module)
|
||||
return 1
|
||||
})
|
||||
@@ -111,29 +106,33 @@ func luaIPMatches(L *lua.LState) (int, bool) {
|
||||
return 1, true
|
||||
}
|
||||
|
||||
func luaIPFilterIPs(L *lua.LState) (int, bool) {
|
||||
matcher, ips, ok := readLuaIPMatcherArgs[[]net.IP](L)
|
||||
if !ok {
|
||||
return 0, false
|
||||
func newLuaIPFilterIPs(pushIPs func(*lua.LState, []net.IP)) xlua.DirectMethod {
|
||||
return func(L *lua.LState) (int, bool) {
|
||||
matcher, ips, ok := readLuaIPMatcherArgs[[]net.IP](L)
|
||||
if !ok {
|
||||
return 0, false
|
||||
}
|
||||
matched, unmatched := matcher.FilterIPs(ips)
|
||||
pushIPs(L, matched)
|
||||
pushIPs(L, unmatched)
|
||||
return 2, true
|
||||
}
|
||||
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
|
||||
func newLuaDomainMatch(pushMatches func(*lua.LState, []uint32)) xlua.DirectMethod {
|
||||
return func(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 {
|
||||
pushMatches(L, matcher.Match(string(domain)))
|
||||
return 1, true
|
||||
}
|
||||
}
|
||||
}
|
||||
return 0, false
|
||||
}
|
||||
return 0, false
|
||||
}
|
||||
|
||||
func luaDomainMatchAny(L *lua.LState) (int, bool) {
|
||||
|
||||
@@ -5,6 +5,23 @@ import (
|
||||
luar "layeh.com/gopher-luar"
|
||||
)
|
||||
|
||||
// NewSlicePusher captures luar's slice metatable during state initialization.
|
||||
// The returned function wraps slices without reflection or metatable lookup,
|
||||
// and pushes nil for nil slices. Use it with this state or its coroutines.
|
||||
func NewSlicePusher[T any](L *glua.LState) func(*glua.LState, []T) {
|
||||
metatable := luar.New(L, []T{}).(*glua.LUserData).Metatable
|
||||
return func(L *glua.LState, values []T) {
|
||||
if values == nil {
|
||||
L.Push(glua.LNil)
|
||||
return
|
||||
}
|
||||
userdata := L.NewUserData()
|
||||
userdata.Value = values
|
||||
userdata.Metatable = metatable
|
||||
L.Push(userdata)
|
||||
}
|
||||
}
|
||||
|
||||
// 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.
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
luar "layeh.com/gopher-luar"
|
||||
)
|
||||
|
||||
func TestSlicePusher(t *testing.T) {
|
||||
L := glua.NewState()
|
||||
defer L.Close()
|
||||
push := NewSlicePusher[int](L)
|
||||
values := []int{3, 5}
|
||||
L.SetGlobal("getValues", L.NewFunction(func(L *glua.LState) int {
|
||||
push(L, values)
|
||||
return 1
|
||||
}))
|
||||
if err := L.DoString(`
|
||||
local values = getValues()
|
||||
assert(#values == 2 and values[1] == 3 and values[2] == 5)
|
||||
values[2] = 7
|
||||
local co = coroutine.create(function()
|
||||
local values = getValues()
|
||||
assert(#values == 2 and values[1] == 3 and values[2] == 7)
|
||||
return true
|
||||
end)
|
||||
local ok, result = coroutine.resume(co)
|
||||
assert(ok and result == true)
|
||||
`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if values[1] != 7 {
|
||||
t.Fatal("slice storage was copied")
|
||||
}
|
||||
push(L, nil)
|
||||
if L.Get(-1) != glua.LNil {
|
||||
t.Fatal("nil slice must push Lua nil")
|
||||
}
|
||||
L.Pop(1)
|
||||
push(L, []int{})
|
||||
L.SetGlobal("empty", L.Get(-1))
|
||||
L.Pop(1)
|
||||
if err := L.DoString(`assert(type(empty) == "userdata" and #empty == 0)`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSlicePusherMetatablePerState(t *testing.T) {
|
||||
first := glua.NewState()
|
||||
defer first.Close()
|
||||
second := glua.NewState()
|
||||
defer second.Close()
|
||||
NewSlicePusher[int](first)(first, []int{1})
|
||||
NewSlicePusher[int](second)(second, []int{1})
|
||||
if first.Get(-1).(*glua.LUserData).Metatable == second.Get(-1).(*glua.LUserData).Metatable {
|
||||
t.Fatal("independent states share a slice metatable")
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkSlicePusher(b *testing.B) {
|
||||
L := glua.NewState()
|
||||
defer L.Close()
|
||||
ips := []net.IP{net.ParseIP("127.0.0.1")}
|
||||
pushIPs := NewSlicePusher[net.IP](L)
|
||||
for _, benchmark := range []struct {
|
||||
name string
|
||||
push func(*glua.LState, []net.IP)
|
||||
}{
|
||||
{"bare", func(L *glua.LState, ips []net.IP) { PushUserData(L, ips) }},
|
||||
{"luar", func(L *glua.LState, ips []net.IP) { L.Push(luar.New(L, ips)) }},
|
||||
{"cached", pushIPs},
|
||||
} {
|
||||
b.Run(benchmark.name, func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
benchmark.push(L, ips)
|
||||
L.Pop(1)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user