mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-02 13:56:39 +00:00
Expose localdns.Client in xray.dns.Servers with the ID "localhost", so Lua scripts can use the same server API for both clients.
297 lines
9.7 KiB
Go
297 lines
9.7 KiB
Go
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"
|
|
"github.com/xtls/xray-core/features/dns/localdns"
|
|
lua "github.com/yuin/gopher-lua"
|
|
)
|
|
|
|
func TestReadLuaDNSResult(t *testing.T) {
|
|
L := lua.NewState()
|
|
defer L.Close()
|
|
want := []net.IP{net.ParseIP("8.8.8.8"), {127, 0, 0, 1}, net.ParseIP("::1")}
|
|
addresses := L.NewUserData()
|
|
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)
|
|
}
|
|
for i := range want {
|
|
if !ips[i].Equal(want[i]) {
|
|
t.Fatalf("IP %d = %v, want %v", i, ips[i], want[i])
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestReadLuaDNSResultValidation(t *testing.T) {
|
|
L := lua.NewState()
|
|
defer L.Close()
|
|
|
|
for _, tc := range []struct {
|
|
name string
|
|
change func(*[3]lua.LValue)
|
|
want string
|
|
}{
|
|
{"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) {
|
|
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("readLuaDNSResult error = %v, want %q", err, tc.want)
|
|
}
|
|
})
|
|
}
|
|
|
|
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)
|
|
}
|
|
}
|
|
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(domain, ipv4, ipv6, fake) while true do end end`); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond)
|
|
defer cancel()
|
|
_, _, err := (&DNS{}).CallLuaHook(L, ctx, "example.com", featureDNS.IPOption{IPv4Enable: true})
|
|
if err == nil {
|
|
t.Fatal("CallLuaHook did not stop after context cancellation")
|
|
}
|
|
if L.Context() != nil {
|
|
t.Fatal("CallLuaHook left the canceled context on the Lua state")
|
|
}
|
|
}
|
|
|
|
func TestCallLuaHookNormalizesDomain(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(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)
|
|
}
|
|
s := &DNS{}
|
|
if _, _, err := s.CallLuaHook(L, context.Background(), "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
|
|
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").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(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 luaDNSClient struct {
|
|
featureDNS.Client
|
|
lookup func(string, featureDNS.IPOption) ([]net.IP, uint32, error)
|
|
}
|
|
|
|
func (c *luaDNSClient) LookupIP(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
|
return c.lookup(domain, option)
|
|
}
|
|
|
|
func TestLuaDNSClientQuery(t *testing.T) {
|
|
L := lua.NewState()
|
|
defer L.Close()
|
|
L.SetContext(context.Background())
|
|
defer L.RemoveContext()
|
|
geodata.RegisterLua(L)
|
|
want := []net.IP{{127, 0, 0, 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)
|
|
}
|
|
return want, 42, nil
|
|
}}
|
|
RegisterLua(L, client)
|
|
if err := L.DoString(`
|
|
local dns = require("xray.dns")
|
|
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))
|
|
`); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
got := L.GetGlobal("ips").(*lua.LUserData).Value.([]net.IP)
|
|
if &got[0] != &want[0] {
|
|
t.Fatal("dns.Query copied the IP slice")
|
|
}
|
|
}
|
|
|
|
func TestLuaDNSLocalClient(t *testing.T) {
|
|
L := lua.NewState()
|
|
defer L.Close()
|
|
L.SetContext(context.Background())
|
|
defer L.RemoveContext()
|
|
RegisterLua(L, localdns.New())
|
|
if err := L.DoString(`
|
|
local dns = require("xray.dns")
|
|
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)
|
|
`); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
for _, name := range []string{"serverIPs", "clientIPs"} {
|
|
ips := L.GetGlobal(name).(*lua.LUserData).Value.([]net.IP)
|
|
if len(ips) != 1 || !ips[0].Equal(net.ParseIP("127.0.0.1")) {
|
|
t.Fatalf("%s = %v", name, ips)
|
|
}
|
|
}
|
|
}
|
|
|
|
type benchmarkLuaNameServer struct {
|
|
ips []net.IP
|
|
}
|
|
|
|
func (*benchmarkLuaNameServer) Name() string { return "benchmark" }
|
|
func (*benchmarkLuaNameServer) IsDisableCache() bool { return true }
|
|
func (s *benchmarkLuaNameServer) QueryIP(context.Context, string, featureDNS.IPOption) ([]net.IP, uint32, error) {
|
|
return s.ips, 60, nil
|
|
}
|
|
|
|
// BenchmarkLuaDNSHookCall isolates a preloaded Lua hook and its server:Query bridge.
|
|
// The direct case measures the same DNS client without Lua.
|
|
func BenchmarkLuaDNSHookCall(b *testing.B) {
|
|
option := featureDNS.IPOption{IPv4Enable: true}
|
|
ip := net.ParseIP("127.0.0.1")
|
|
upstream := &benchmarkLuaNameServer{ips: []net.IP{ip}}
|
|
client := &Client{server: upstream, ipOption: &option, timeoutMs: time.Second}
|
|
server := &DNS{clients: []*Client{client}}
|
|
L := lua.NewState()
|
|
defer L.Close()
|
|
server.RegisterLua(L)
|
|
if err := L.DoString(`
|
|
local server = require("xray.dns").Servers[1]
|
|
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
|
return server:Query(domain, ipv4, ipv6, fake)
|
|
end
|
|
`); err != nil {
|
|
b.Fatal(err)
|
|
}
|
|
|
|
ctx := context.Background()
|
|
for _, bench := range []struct {
|
|
name string
|
|
query func() ([]net.IP, uint32, error)
|
|
}{
|
|
{"direct", func() ([]net.IP, uint32, error) { return client.QueryIP(ctx, "example.com", option) }},
|
|
{"lua_hook", func() ([]net.IP, uint32, error) { return server.CallLuaHook(L, ctx, "example.com", option) }},
|
|
} {
|
|
b.Run(bench.name, func(b *testing.B) {
|
|
b.ReportAllocs()
|
|
b.ResetTimer()
|
|
var ips []net.IP
|
|
var ttl uint32
|
|
var err error
|
|
for i := 0; i < b.N; i++ {
|
|
ips, ttl, err = bench.query()
|
|
if err != nil {
|
|
b.Fatal(err)
|
|
}
|
|
}
|
|
b.StopTimer()
|
|
if ttl != 60 || len(ips) != 1 || !ips[0].Equal(ip) {
|
|
b.Fatalf("query() = %v, TTL %d; want %v, TTL 60", ips, ttl, ip)
|
|
}
|
|
})
|
|
}
|
|
}
|