dns: reduce Lua allocations with flat arguments and returns and native IP slice userdata

This commit is contained in:
Meo597
2026-09-28 22:12:47 +08:00
parent 459301d42e
commit 5e1bb92b98
5 changed files with 162 additions and 177 deletions
+110 -75
View File
@@ -3,107 +3,84 @@ package dns
import (
"context"
go_errors "errors"
"math"
"strings"
"testing"
"time"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/net"
featureDNS "github.com/xtls/xray-core/features/dns"
lua "github.com/yuin/gopher-lua"
)
func TestDecodeLuaDNSResultNativeIP(t *testing.T) {
func TestReadLuaDNSResult(t *testing.T) {
L := lua.NewState()
defer L.Close()
ip := net.ParseIP("127.0.0.1")
address := L.NewUserData()
address.Value = ip
addresses := L.NewTable()
addresses.RawSetInt(1, address)
result := L.NewTable()
result.RawSetString("ips", addresses)
result.RawSetString("ttl", lua.LNumber(60))
got, ttl, err := decodeLuaDNSResult(result, featureDNS.IPOption{IPv4Enable: true})
if err != nil || ttl != 60 || len(got) != 1 || !got[0].Equal(ip) {
t.Fatalf("decodeLuaDNSResult() = %v, %d, %v", got, ttl, err)
}
addresses.RawSetInt(1, lua.LString("127.0.0.1"))
if _, _, err := decodeLuaDNSResult(result, featureDNS.IPOption{IPv4Enable: true}); err == nil {
t.Fatal("decodeLuaDNSResult accepted a string IP")
}
}
func TestDecodeLuaDNSResultNativeSliceCopiesIP(t *testing.T) {
L := lua.NewState()
defer L.Close()
original := net.ParseIP("8.8.8.8")
want := []net.IP{net.ParseIP("8.8.8.8"), {127, 0, 0, 1}, net.ParseIP("::1")}
addresses := L.NewUserData()
addresses.Value = []net.IP{original}
result := L.NewTable()
result.RawSetString("ips", addresses)
result.RawSetString("ttl", lua.LNumber(45))
ips, ttl, err := decodeLuaDNSResult(result, featureDNS.IPOption{IPv4Enable: true})
if err != nil || ttl != 45 || len(ips) != 1 || !ips[0].Equal(original) {
t.Fatalf("decodeLuaDNSResult() = %v, TTL %d, %v", ips, ttl, err)
addresses.Value = want
ips, ttl, err := readLuaDNSResult(addresses, lua.LNumber(45), lua.LNil)
if err != nil || ttl != 45 || len(ips) != len(want) {
t.Fatalf("readLuaDNSResult() = %v, TTL %d, %v", ips, ttl, err)
}
original[len(original)-1] = 9
if !ips[0].Equal(net.ParseIP("8.8.8.8")) {
t.Fatalf("decoded IP changed with input: %v", ips[0])
for i := range want {
if !ips[i].Equal(want[i]) {
t.Fatalf("IP %d = %v, want %v", i, ips[i], want[i])
}
}
}
func TestDecodeLuaDNSResultValidation(t *testing.T) {
func TestReadLuaDNSResultValidation(t *testing.T) {
L := lua.NewState()
defer L.Close()
option := featureDNS.IPOption{IPv4Enable: true}
for _, tc := range []struct {
name string
change func(*lua.LTable, *lua.LTable)
change func(*[3]lua.LValue)
want string
}{
{"fractional TTL", func(result, _ *lua.LTable) { result.RawSetString("ttl", lua.LNumber(1.5)) }, "invalid TTL"},
{"oversized TTL", func(result, _ *lua.LTable) { result.RawSetString("ttl", lua.LNumber(4294967296)) }, "invalid TTL"},
{"string address", func(_, addresses *lua.LTable) { addresses.RawSetInt(1, lua.LString("127.0.0.1")) }, "invalid address"},
{"missing addresses", func(result, _ *lua.LTable) { result.RawSetString("ips", lua.LString("127.0.0.1")) }, "must be an array"},
{"script error", func(result, _ *lua.LTable) { result.RawSetString("error", lua.LString("blocked by script")) }, "blocked by script"},
{"fractional TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(1.5) }, "invalid TTL"},
{"oversized TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(4294967296) }, "invalid TTL"},
{"negative TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(-1) }, "invalid TTL"},
{"NaN TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(math.NaN()) }, "invalid TTL"},
{"missing TTL", func(v *[3]lua.LValue) { v[1] = lua.LNil }, "invalid TTL"},
{"string IPs", func(v *[3]lua.LValue) { v[0] = lua.LString("127.0.0.1") }, "native IP slice"},
{"wrong userdata", func(v *[3]lua.LValue) { v[0].(*lua.LUserData).Value = net.ParseIP("127.0.0.1") }, "native IP slice"},
{"script error", func(v *[3]lua.LValue) { v[2] = lua.LString("blocked by script") }, "blocked by script"},
{"invalid error", func(v *[3]lua.LValue) { v[2] = lua.LTrue }, "error or string"},
} {
t.Run(tc.name, func(t *testing.T) {
address := L.NewUserData()
address.Value = net.ParseIP("127.0.0.1")
addresses := L.NewTable()
addresses.RawSetInt(1, address)
result := L.NewTable()
result.RawSetString("ips", addresses)
result.RawSetString("ttl", lua.LNumber(60))
tc.change(result, addresses)
_, _, err := decodeLuaDNSResult(result, option)
addresses := L.NewUserData()
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
values := [3]lua.LValue{addresses, lua.LNumber(60), lua.LNil}
tc.change(&values)
_, _, err := readLuaDNSResult(values[0], values[1], values[2])
if err == nil || !strings.Contains(err.Error(), tc.want) {
t.Fatalf("decodeLuaDNSResult error = %v, want %q", err, tc.want)
t.Fatalf("readLuaDNSResult error = %v, want %q", err, tc.want)
}
})
}
address := L.NewUserData()
address.Value = net.ParseIP("127.0.0.1")
addresses := L.NewTable()
addresses.RawSetInt(1, address)
result := L.NewTable()
result.RawSetString("ips", addresses)
result.RawSetString("ttl", lua.LNumber(60))
if _, _, err := decodeLuaDNSResult(result, featureDNS.IPOption{IPv6Enable: true}); err == nil {
t.Fatal("decodeLuaDNSResult accepted IPv4 with IPv6-only option")
addresses := L.NewUserData()
addresses.Value = []net.IP(nil)
for _, empty := range []lua.LValue{addresses, lua.LNil} {
if _, _, err := readLuaDNSResult(empty, lua.LNumber(0), lua.LNil); !go_errors.Is(err, featureDNS.ErrEmptyResponse) {
t.Fatalf("empty result error = %v, want ErrEmptyResponse", err)
}
}
result.RawSetString("ips", L.NewTable())
if _, _, err := decodeLuaDNSResult(result, option); !go_errors.Is(err, featureDNS.ErrEmptyResponse) {
t.Fatalf("empty result error = %v, want ErrEmptyResponse", err)
wantErr := go_errors.New("upstream failed")
errorValue := L.NewUserData()
errorValue.Value = wantErr
if _, _, err := readLuaDNSResult(lua.LNil, lua.LNil, errorValue); err != wantErr {
t.Fatalf("upstream error = %v, want original error %v", err, wantErr)
}
}
func TestCallLuaHookCancellation(t *testing.T) {
L := lua.NewState()
defer L.Close()
if err := L.DoString(`function handleDNSQuery(q) while true do end end`); err != nil {
if err := L.DoString(`function handleDNSQuery(domain, ipv4, ipv6, fake) while true do end end`); err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond)
@@ -120,16 +97,14 @@ func TestCallLuaHookCancellation(t *testing.T) {
func TestCallLuaHookNormalizesDomain(t *testing.T) {
L := lua.NewState()
defer L.Close()
address := L.NewUserData()
address.Value = net.ParseIP("127.0.0.1")
L.SetGlobal("ip", address)
addresses := L.NewUserData()
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
L.SetGlobal("ips", addresses)
if err := L.DoString(`
function handleDNSQuery(q)
assert(type(q) == "table")
assert(q.domain == "example.com")
assert(q.ipv4 and not q.ipv6 and not q.fake)
assert(q.ctx == nil)
return {ips = {ip}, ttl = 60}
function handleDNSQuery(domain, ipv4, ipv6, fake)
assert(domain == "example.com")
assert(ipv4 and not ipv6 and not fake)
return ips, 60, nil
end
`); err != nil {
t.Fatal(err)
@@ -140,6 +115,66 @@ func TestCallLuaHookNormalizesDomain(t *testing.T) {
}
}
func TestCallLuaHookRestoresState(t *testing.T) {
for _, tc := range []struct {
name string
body string
wantErr bool
}{
{"success", `return ips, 60`, false},
{"error", `error("failed")`, true},
} {
t.Run(tc.name, func(t *testing.T) {
L := lua.NewState()
defer L.Close()
addresses := L.NewUserData()
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
L.SetGlobal("ips", addresses)
if err := L.DoString("function handleDNSQuery() " + tc.body + " end"); err != nil {
t.Fatal(err)
}
previous, cancel := context.WithCancel(context.Background())
defer cancel()
L.SetContext(previous)
L.Push(lua.LTrue)
_, _, err := (&DNS{}).CallLuaHook(L, context.Background(), "example.com", featureDNS.IPOption{IPv4Enable: true})
if (err != nil) != tc.wantErr {
t.Fatalf("hook error = %v, want error %t", err, tc.wantErr)
}
if L.Context() != previous || L.GetTop() != 1 || L.Get(1) != lua.LTrue {
t.Fatal("hook did not restore the previous context and stack")
}
})
}
}
func TestLuaDNSServerQuery(t *testing.T) {
L := lua.NewState()
defer L.Close()
geodata.RegisterLua(L)
option := featureDNS.IPOption{IPv4Enable: true}
ips := []net.IP{net.ParseIP("127.0.0.1"), net.ParseIP("8.8.8.8")}
server := &DNS{clients: []*Client{{server: &benchmarkLuaNameServer{ips: ips}, ipOption: &option, timeoutMs: time.Second}}}
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"})
function handleDNSQuery(domain, ipv4, ipv6, fake)
local ips, ttl, err = server:query(domain, ipv4, ipv6, fake)
assert(type(ips) == "userdata" and not err)
assert(matcher:anyMatch(ips))
local matched = matcher:filterIPs(ips)
return matched, ttl, err
end
`); err != nil {
t.Fatal(err)
}
got, ttl, err := server.CallLuaHook(L, context.Background(), "example.com", option)
if err != nil || ttl != 60 || len(got) != 1 || !got[0].Equal(ips[0]) {
t.Fatalf("server query = %v, TTL %d, %v", got, ttl, err)
}
}
type benchmarkLuaNameServer struct {
ips []net.IP
}
@@ -163,8 +198,8 @@ func BenchmarkLuaDNSHookCall(b *testing.B) {
server.RegisterLua(L)
if err := L.DoString(`
local server = require("xray.dns").servers[1]
function handleDNSQuery(q)
return server:query(q)
function handleDNSQuery(domain, ipv4, ipv6, fake)
return server:query(domain, ipv4, ipv6, fake)
end
`); err != nil {
b.Fatal(err)