mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-09-30 13:05:43 +00:00
dns: reduce Lua allocations with flat arguments and returns and native IP slice userdata
This commit is contained in:
+39
-89
@@ -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
|
||||
}
|
||||
|
||||
+110
-75
@@ -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)
|
||||
|
||||
+1
-1
@@ -39,7 +39,7 @@ func newScriptEngine(path string, server *DNS) (*scriptEngine, error) {
|
||||
}
|
||||
if L.GetGlobal("handleDNSQuery").Type() != lua.LTFunction {
|
||||
L.Close()
|
||||
return nil, errors.New("DNS script must define handleDNSQuery(q)")
|
||||
return nil, errors.New("DNS script must define handleDNSQuery(domain, ipv4, ipv6, fake)")
|
||||
}
|
||||
return L, nil
|
||||
})
|
||||
|
||||
+11
-11
@@ -46,12 +46,12 @@ for _, server in ipairs(servers) do
|
||||
end
|
||||
assert(by_id.primary and by_id.fallback, "primary and fallback DNS servers are required")
|
||||
|
||||
function handleDNSQuery(q)
|
||||
local answer = by_id.primary:query(q)
|
||||
if not answer.error and us_ips:anyMatch(answer.ips) then
|
||||
return answer
|
||||
function handleDNSQuery(domain, ipv4, ipv6, fake)
|
||||
local ips, ttl, err = by_id.primary:query(domain, ipv4, ipv6, fake)
|
||||
if not err and us_ips:anyMatch(ips) then
|
||||
return ips, ttl, nil
|
||||
end
|
||||
return by_id.fallback:query(q)
|
||||
return by_id.fallback:query(domain, ipv4, ipv6, fake)
|
||||
end
|
||||
`
|
||||
scriptPath := filepath.Join(t.TempDir(), "geoip_fallback.lua")
|
||||
@@ -144,12 +144,12 @@ func TestDNSScriptHookErrorAndFakeDNSOption(t *testing.T) {
|
||||
local server = require("xray.dns").servers[1]
|
||||
local log = require("xray.log")
|
||||
log.info("DNS script loaded")
|
||||
function handleDNSQuery(q)
|
||||
log.debug("DNS query: ", q.domain)
|
||||
if q.domain == "bad.example" then error("script failure") end
|
||||
local answer = server:query(q)
|
||||
if answer.error then log.error("DNS failed: ", answer.error) end
|
||||
return answer
|
||||
function handleDNSQuery(domain, ipv4, ipv6, fake)
|
||||
log.debug("DNS query: ", domain)
|
||||
if domain == "bad.example" then error("script failure") end
|
||||
local ips, ttl, err = server:query(domain, ipv4, ipv6, fake)
|
||||
if err then log.error("DNS failed: ", err) end
|
||||
return ips, ttl, err
|
||||
end
|
||||
`
|
||||
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
|
||||
|
||||
@@ -129,7 +129,7 @@ func TestDNSScriptConfig(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
t.Setenv("xray.location.confdir", dir)
|
||||
path := filepath.Join(dir, "lookup.lua")
|
||||
if err := os.WriteFile(path, []byte("function handleDNSQuery(q) end"), 0o600); err != nil {
|
||||
if err := os.WriteFile(path, []byte("function handleDNSQuery(domain, ipv4, ipv6, fake) end"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user