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
+39 -89
View File
@@ -11,8 +11,7 @@ import (
lua "github.com/yuin/gopher-lua"
)
// RegisterLua makes xray.dns available to require in an LState. The caller
// owns the state and registers modules before running the script top level.
// RegisterLua makes xray.dns available to require in an LState.
func (s *DNS) RegisterLua(L *lua.LState) {
L.PreloadModule("xray.dns", func(L *lua.LState) int {
servers := L.NewTable()
@@ -22,16 +21,15 @@ func (s *DNS) RegisterLua(L *lua.LState) {
server.RawSetString("id", lua.LString(client.id))
server.RawSetString("query", L.NewFunction(func(L *lua.LState) int {
q := L.CheckTable(2)
domain, ok := q.RawGetString("domain").(lua.LString)
domain, ok := L.Get(2).(lua.LString)
if !ok {
L.RaiseError("server:query requires a domain")
return 0
}
option := featureDNS.IPOption{
IPv4Enable: q.RawGetString("ipv4") == lua.LTrue,
IPv6Enable: q.RawGetString("ipv6") == lua.LTrue,
FakeEnable: q.RawGetString("fake") == lua.LTrue,
IPv4Enable: L.CheckBool(3),
IPv6Enable: L.CheckBool(4),
FakeEnable: L.CheckBool(5),
}
ctx := L.Context()
if ctx == nil {
@@ -46,22 +44,18 @@ func (s *DNS) RegisterLua(L *lua.LState) {
} else {
ips, ttl, err = client.QueryIP(ctx, string(domain), option)
}
result := L.CreateTable(0, 3)
addresses := L.CreateTable(len(ips), 0)
for j, ip := range ips {
address := L.NewUserData()
address.Value = ip
addresses.RawSetInt(j+1, address)
}
result.RawSetString("ips", addresses)
result.RawSetString("ttl", lua.LNumber(ttl))
addresses := L.NewUserData()
addresses.Value = ips
L.Push(addresses)
L.Push(lua.LNumber(ttl))
if err != nil {
ud := L.NewUserData()
ud.Value = err
result.RawSetString("error", ud)
L.Push(ud)
} else {
L.Push(lua.LNil)
}
L.Push(result)
return 1
return 3
}))
servers.RawSetInt(i+1, server)
}
@@ -72,15 +66,9 @@ func (s *DNS) RegisterLua(L *lua.LState) {
})
}
// CallLuaHook invokes handleDNSQuery on a state owned by the caller. Domain and option
// must already have passed DNS normalization, hosts, and address-family handling.
// The caller serializes access to its state; ctx cancels Lua execution and upstream calls.
// CallLuaHook invokes handleDNSQuery in the supplied state.
// Returned slices and IP bytes may share storage with DNS caches or matcher inputs.
func (s *DNS) CallLuaHook(L *lua.LState, ctx context.Context, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
q := L.CreateTable(0, 4)
q.RawSetString("domain", lua.LString(strings.ToLower(domain)))
q.RawSetString("ipv4", lua.LBool(option.IPv4Enable))
q.RawSetString("ipv6", lua.LBool(option.IPv6Enable))
q.RawSetString("fake", lua.LBool(option.FakeEnable))
previous := L.Context()
L.SetContext(ctx)
defer func() {
@@ -92,89 +80,51 @@ func (s *DNS) CallLuaHook(L *lua.LState, ctx context.Context, domain string, opt
}()
fn := L.GetGlobal("handleDNSQuery")
if fn.Type() != lua.LTFunction {
return nil, 0, errors.New("DNS script must define handleDNSQuery(q)")
return nil, 0, errors.New("DNS script must define handleDNSQuery(domain, ipv4, ipv6, fake)")
}
if err := L.CallByParam(lua.P{Fn: fn, NRet: 1, Protect: true}, q); err != nil {
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
lua.LString(strings.ToLower(domain)), lua.LBool(option.IPv4Enable),
lua.LBool(option.IPv6Enable), lua.LBool(option.FakeEnable)); err != nil {
return nil, 0, err
}
value := L.Get(-1)
L.Pop(1)
ips, ttl, err := decodeLuaDNSResult(value, option)
addresses, ttlValue, errorValue := L.Get(-3), L.Get(-2), L.Get(-1)
L.Pop(3)
ips, ttl, err := readLuaDNSResult(addresses, ttlValue, errorValue)
if ctx.Err() != nil {
return nil, 0, ctx.Err()
}
return ips, ttl, err
}
func decodeLuaDNSResult(value lua.LValue, option featureDNS.IPOption) ([]net.IP, uint32, error) {
table, ok := value.(*lua.LTable)
if !ok {
return nil, 0, errors.New("DNS script result must be a table")
}
if v := table.RawGetString("error"); v != lua.LNil {
if ud, ok := v.(*lua.LUserData); ok {
func readLuaDNSResult(addresses, ttlValue, errorValue lua.LValue) ([]net.IP, uint32, error) {
if errorValue != lua.LNil {
if ud, ok := errorValue.(*lua.LUserData); ok {
if err, ok := ud.Value.(error); ok {
return nil, 0, err
}
}
if s, ok := v.(lua.LString); ok {
if s, ok := errorValue.(lua.LString); ok {
return nil, 0, errors.New(string(s))
}
return nil, 0, errors.New("DNS script error must be an error or string")
}
ttlValue, ok := table.RawGetString("ttl").(lua.LNumber)
if !ok || ttlValue < 0 || ttlValue > math.MaxUint32 || math.Trunc(float64(ttlValue)) != float64(ttlValue) {
ttl, ok := ttlValue.(lua.LNumber)
if !ok || ttl < 0 || ttl > math.MaxUint32 || math.Trunc(float64(ttl)) != float64(ttl) {
return nil, 0, errors.New("DNS script returned invalid TTL")
}
var ips []net.IP
switch addresses := table.RawGetString("ips").(type) {
case *lua.LTable:
ips = make([]net.IP, 0, addresses.Len())
for i := 1; i <= addresses.Len(); i++ {
ip, err := decodeLuaIP(addresses.RawGetInt(i), i, option)
if err != nil {
return nil, 0, err
}
ips = append(ips, ip)
}
case *lua.LUserData:
addressesIP, ok := addresses.Value.([]net.IP)
if !ok {
return nil, 0, errors.New("DNS script result.ips must be an array")
}
ips = make([]net.IP, 0, len(addressesIP))
for i, ip := range addressesIP {
valid, err := validateLuaIP(ip, i+1, option)
if err != nil {
return nil, 0, err
}
ips = append(ips, valid)
}
default:
return nil, 0, errors.New("DNS script result.ips must be an array")
if addresses == lua.LNil {
return nil, 0, featureDNS.ErrEmptyResponse
}
ud, ok := addresses.(*lua.LUserData)
if !ok {
return nil, 0, errors.New("DNS script IPs must be native IP slice userdata")
}
ips, ok := ud.Value.([]net.IP)
if !ok {
return nil, 0, errors.New("DNS script IPs must be native IP slice userdata")
}
if len(ips) == 0 {
return nil, 0, featureDNS.ErrEmptyResponse
}
return ips, uint32(ttlValue), nil
}
func decodeLuaIP(value lua.LValue, index int, option featureDNS.IPOption) (net.IP, error) {
address, ok := value.(*lua.LUserData)
if !ok {
return nil, errors.New("DNS script returned invalid address at index ", index)
}
ip, ok := address.Value.(net.IP)
if !ok {
return nil, errors.New("DNS script returned invalid address at index ", index)
}
return validateLuaIP(ip, index, option)
}
func validateLuaIP(ip net.IP, index int, option featureDNS.IPOption) (net.IP, error) {
ip4 := ip.To4()
if ip.To16() == nil || (ip4 != nil && !option.IPv4Enable) || (ip4 == nil && !option.IPv6Enable) {
return nil, errors.New("DNS script returned invalid or disabled address at index ", index)
}
return append(net.IP(nil), ip...), nil
return ips, uint32(ttl), nil
}