diff --git a/app/dns/lua.go b/app/dns/lua.go index ab31471da..3c5194baf 100644 --- a/app/dns/lua.go +++ b/app/dns/lua.go @@ -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 diff --git a/app/dns/lua_test.go b/app/dns/lua_test.go index 305e3588a..18900361d 100644 --- a/app/dns/lua_test.go +++ b/app/dns/lua_test.go @@ -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 } diff --git a/app/router/lua.go b/app/router/lua.go index c09106228..6cf55d5f1 100644 --- a/app/router/lua.go +++ b/app/router/lua.go @@ -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 { diff --git a/app/router/lua_test.go b/app/router/lua_test.go index f90b1b867..bfe475435 100644 --- a/app/router/lua_test.go +++ b/app/router/lua_test.go @@ -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 { diff --git a/common/geodata/lua.go b/common/geodata/lua.go index d76534962..43f7363ff 100644 --- a/common/geodata/lua.go +++ b/common/geodata/lua.go @@ -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) { diff --git a/common/lua/luar.go b/common/lua/luar.go index 6494187f6..072ee663f 100644 --- a/common/lua/luar.go +++ b/common/lua/luar.go @@ -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. diff --git a/common/lua/luar_test.go b/common/lua/luar_test.go new file mode 100644 index 000000000..cad82b6c6 --- /dev/null +++ b/common/lua/luar_test.go @@ -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) + } + }) + } +}