mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-11 10:05:47 +00:00
dns: preserve Lua states after recoverable query errors & refactor
This commit is contained in:
+17
-15
@@ -126,31 +126,32 @@ func newLuaClientQuery(L *lua.LState, client featureDNS.Client) *lua.LFunction {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// callLuaHook invokes HandleDNSQuery in the supplied state.
|
// callLuaQuery runs HandleDNSQuery and leaves (ips, ttl, err) on the stack.
|
||||||
// Returned slices and IP bytes may share storage with DNS caches or matcher inputs.
|
func callLuaQuery(L *lua.LState, domain string, option featureDNS.IPOption) error {
|
||||||
func (s *DNS) callLuaHook(L *lua.LState, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
|
||||||
top := L.GetTop()
|
|
||||||
defer L.SetTop(top)
|
|
||||||
fn := L.GetGlobal("HandleDNSQuery")
|
fn := L.GetGlobal("HandleDNSQuery")
|
||||||
if fn.Type() != lua.LTFunction {
|
if fn.Type() != lua.LTFunction {
|
||||||
return nil, 0, errors.New("DNS script must define HandleDNSQuery(...)")
|
return errors.New("DNS script must define HandleDNSQuery(...)")
|
||||||
}
|
}
|
||||||
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
|
|
||||||
lua.LString(strings.ToLower(domain)), lua.LBool(option.IPv4Enable),
|
return L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
|
||||||
lua.LBool(option.IPv6Enable), lua.LBool(option.FakeEnable)); err != nil {
|
lua.LString(strings.ToLower(domain)),
|
||||||
return nil, 0, err
|
lua.LBool(option.IPv4Enable),
|
||||||
}
|
lua.LBool(option.IPv6Enable),
|
||||||
return readLuaDNSResult(L.Get(-3), L.Get(-2), L.Get(-1))
|
lua.LBool(option.FakeEnable))
|
||||||
}
|
}
|
||||||
|
|
||||||
func readLuaDNSResult(addresses, ttlValue, errorValue lua.LValue) ([]net.IP, uint32, error) {
|
// readLuaQueryResult reads (ips, ttl, err) from the stack without copying the IPs.
|
||||||
if err := xlua.ReadError(errorValue, "DNS script error must be an error or string"); err != nil {
|
func readLuaQueryResult(L *lua.LState) ([]net.IP, uint32, error) {
|
||||||
|
if err := xlua.ReadError(L.Get(-1), "DNS script error must be an error or string"); err != nil {
|
||||||
return nil, 0, err
|
return nil, 0, err
|
||||||
}
|
}
|
||||||
ttl, err := xlua.ReadUint32(ttlValue, "DNS script returned invalid TTL")
|
|
||||||
|
ttl, err := xlua.ReadUint32(L.Get(-2), "DNS script returned invalid TTL")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, 0, err
|
return nil, 0, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
addresses := L.Get(-3)
|
||||||
if addresses == lua.LNil {
|
if addresses == lua.LNil {
|
||||||
return nil, 0, featureDNS.ErrEmptyResponse
|
return nil, 0, featureDNS.ErrEmptyResponse
|
||||||
}
|
}
|
||||||
@@ -161,5 +162,6 @@ func readLuaDNSResult(addresses, ttlValue, errorValue lua.LValue) ([]net.IP, uin
|
|||||||
if len(ips) == 0 {
|
if len(ips) == 0 {
|
||||||
return nil, 0, featureDNS.ErrEmptyResponse
|
return nil, 0, featureDNS.ErrEmptyResponse
|
||||||
}
|
}
|
||||||
|
|
||||||
return ips, ttl, nil
|
return ips, ttl, nil
|
||||||
}
|
}
|
||||||
|
|||||||
+80
-95
@@ -3,7 +3,6 @@ package dns
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
go_errors "errors"
|
go_errors "errors"
|
||||||
"math"
|
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
@@ -15,70 +14,74 @@ import (
|
|||||||
lua "github.com/yuin/gopher-lua"
|
lua "github.com/yuin/gopher-lua"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestReadLuaDNSResult(t *testing.T) {
|
func TestReadLuaQueryResult(t *testing.T) {
|
||||||
L := lua.NewState()
|
wantIPs := []net.IP{net.ParseIP("8.8.8.8"), {127, 0, 0, 1}, net.ParseIP("::1")}
|
||||||
defer L.Close()
|
nativeErr := go_errors.New("upstream failed")
|
||||||
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 {
|
for _, tc := range []struct {
|
||||||
name string
|
name, values string
|
||||||
change func(*[3]lua.LValue)
|
wantIPs []net.IP
|
||||||
want string
|
wantTTL uint32
|
||||||
|
wantErr error
|
||||||
|
wantMessage string
|
||||||
}{
|
}{
|
||||||
{"fractional TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(1.5) }, "invalid TTL"},
|
{name: "IPs", values: `ips, 45`, wantIPs: wantIPs, wantTTL: 45},
|
||||||
{"oversized TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(4294967296) }, "invalid TTL"},
|
{name: "nil IPs", values: `nil, 0`, wantErr: featureDNS.ErrEmptyResponse},
|
||||||
{"negative TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(-1) }, "invalid TTL"},
|
{name: "empty IPs", values: `emptyIPs, 0`, wantErr: featureDNS.ErrEmptyResponse},
|
||||||
{"NaN TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(math.NaN()) }, "invalid TTL"},
|
{name: "native error", values: `nil, nil, nativeError`, wantErr: nativeErr},
|
||||||
{"missing TTL", func(v *[3]lua.LValue) { v[1] = lua.LNil }, "invalid TTL"},
|
{name: "string error", values: `nil, nil, "blocked"`, wantMessage: "blocked"},
|
||||||
{"string IPs", func(v *[3]lua.LValue) { v[0] = lua.LString("127.0.0.1") }, "native IP slice"},
|
{name: "fractional TTL", values: `ips, 1.5`, wantMessage: "invalid TTL"},
|
||||||
{"wrong userdata", func(v *[3]lua.LValue) { v[0].(*lua.LUserData).Value = net.ParseIP("127.0.0.1") }, "native IP slice"},
|
{name: "oversized TTL", values: `ips, 4294967296`, wantMessage: "invalid TTL"},
|
||||||
{"script error", func(v *[3]lua.LValue) { v[2] = lua.LString("blocked by script") }, "blocked by script"},
|
{name: "negative TTL", values: `ips, -1`, wantMessage: "invalid TTL"},
|
||||||
{"invalid error", func(v *[3]lua.LValue) { v[2] = lua.LTrue }, "error or string"},
|
{name: "NaN TTL", values: `ips, 0/0`, wantMessage: "invalid TTL"},
|
||||||
|
{name: "missing TTL", values: `ips`, wantMessage: "invalid TTL"},
|
||||||
|
{name: "string IPs", values: `"127.0.0.1", 60`, wantMessage: "native IP slice"},
|
||||||
|
{name: "wrong userdata", values: `ip, 60`, wantMessage: "native IP slice"},
|
||||||
|
{name: "invalid error", values: `ips, 60, false`, wantMessage: "error or string"},
|
||||||
} {
|
} {
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
addresses := L.NewUserData()
|
L := lua.NewState()
|
||||||
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
|
defer L.Close()
|
||||||
values := [3]lua.LValue{addresses, lua.LNumber(60), lua.LNil}
|
for name, value := range map[string]any{"ips": wantIPs, "ip": wantIPs[0], "emptyIPs": []net.IP(nil), "nativeError": nativeErr} {
|
||||||
tc.change(&values)
|
ud := L.NewUserData()
|
||||||
_, _, err := readLuaDNSResult(values[0], values[1], values[2])
|
ud.Value = value
|
||||||
if err == nil || !strings.Contains(err.Error(), tc.want) {
|
L.SetGlobal(name, ud)
|
||||||
t.Fatalf("readLuaDNSResult error = %v, want %q", err, tc.want)
|
}
|
||||||
|
fn, err := L.LoadString("return " + tc.values)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
ips, ttl, err := readLuaQueryResult(L)
|
||||||
|
switch {
|
||||||
|
case tc.wantErr != nil:
|
||||||
|
if err != tc.wantErr {
|
||||||
|
t.Fatalf("error = %v, want original error %v", err, tc.wantErr)
|
||||||
|
}
|
||||||
|
case tc.wantMessage != "":
|
||||||
|
if err == nil || !strings.Contains(err.Error(), tc.wantMessage) {
|
||||||
|
t.Fatalf("error = %v, want %q", err, tc.wantMessage)
|
||||||
|
}
|
||||||
|
case err != nil:
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if ttl != tc.wantTTL || len(ips) != len(tc.wantIPs) {
|
||||||
|
t.Fatalf("result = %v, TTL %d; want %v, TTL %d", ips, ttl, tc.wantIPs, tc.wantTTL)
|
||||||
|
}
|
||||||
|
for i := range ips {
|
||||||
|
if !ips[i].Equal(tc.wantIPs[i]) {
|
||||||
|
t.Fatalf("IP %d = %v, want %v", i, ips[i], tc.wantIPs[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(ips) != 0 && &ips[0] != &tc.wantIPs[0] {
|
||||||
|
t.Fatal("result copied the IP slice")
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
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) {
|
func TestCallLuaQueryCancellation(t *testing.T) {
|
||||||
L := lua.NewState()
|
L := lua.NewState()
|
||||||
defer L.Close()
|
defer L.Close()
|
||||||
if err := L.DoString(`function HandleDNSQuery(domain, ipv4, ipv6, fake) while true do end end`); err != nil {
|
if err := L.DoString(`function HandleDNSQuery(domain, ipv4, ipv6, fake) while true do end end`); err != nil {
|
||||||
@@ -87,16 +90,16 @@ func TestCallLuaHookCancellation(t *testing.T) {
|
|||||||
ctx, cancel := context.WithCancel(context.Background())
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
cancel()
|
cancel()
|
||||||
L.SetContext(ctx)
|
L.SetContext(ctx)
|
||||||
_, _, err := (&DNS{}).callLuaHook(L, "example.com", featureDNS.IPOption{IPv4Enable: true})
|
err := callLuaQuery(L, "example.com", featureDNS.IPOption{IPv4Enable: true})
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Fatal("CallLuaHook did not stop after context cancellation")
|
t.Fatal("callLuaQuery did not stop after context cancellation")
|
||||||
}
|
}
|
||||||
if L.Context() != ctx {
|
if L.Context() != ctx {
|
||||||
t.Fatal("CallLuaHook changed the Lua state's context")
|
t.Fatal("callLuaQuery changed the Lua state's context")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestCallLuaHookNormalizesDomain(t *testing.T) {
|
func TestCallLuaQuery(t *testing.T) {
|
||||||
L := lua.NewState()
|
L := lua.NewState()
|
||||||
defer L.Close()
|
defer L.Close()
|
||||||
addresses := L.NewUserData()
|
addresses := L.NewUserData()
|
||||||
@@ -111,39 +114,11 @@ func TestCallLuaHookNormalizesDomain(t *testing.T) {
|
|||||||
`); err != nil {
|
`); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
s := &DNS{}
|
if err := callLuaQuery(L, "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil {
|
||||||
if _, _, err := s.callLuaHook(L, "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil {
|
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
}
|
if L.GetTop() != 3 || L.Get(1) != addresses || L.Get(2) != lua.LNumber(60) || L.Get(3) != lua.LNil {
|
||||||
|
t.Fatal("callLuaQuery did not leave the three query results on the stack")
|
||||||
func TestCallLuaHookRestoresStack(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)
|
|
||||||
}
|
|
||||||
L.Push(lua.LTrue)
|
|
||||||
_, _, err := (&DNS{}).callLuaHook(L, "example.com", featureDNS.IPOption{IPv4Enable: true})
|
|
||||||
if (err != nil) != tc.wantErr {
|
|
||||||
t.Fatalf("hook error = %v, want error %t", err, tc.wantErr)
|
|
||||||
}
|
|
||||||
if L.GetTop() != 1 || L.Get(1) != lua.LTrue {
|
|
||||||
t.Fatal("hook did not restore the stack")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -169,7 +144,10 @@ end
|
|||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
L.SetContext(context.Background())
|
L.SetContext(context.Background())
|
||||||
got, ttl, err := server.callLuaHook(L, "example.com", option)
|
if err := callLuaQuery(L, "example.com", option); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
got, ttl, err := readLuaQueryResult(L)
|
||||||
if err != nil || ttl != 60 || len(got) != 1 || !got[0].Equal(ips[0]) {
|
if err != nil || ttl != 60 || len(got) != 1 || !got[0].Equal(ips[0]) {
|
||||||
t.Fatalf("server query = %v, TTL %d, %v", got, ttl, err)
|
t.Fatalf("server query = %v, TTL %d, %v", got, ttl, err)
|
||||||
}
|
}
|
||||||
@@ -244,9 +222,9 @@ func (s *benchmarkLuaNameServer) QueryIP(context.Context, string, featureDNS.IPO
|
|||||||
return s.ips, 60, nil
|
return s.ips, 60, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// BenchmarkLuaDNSHookCall isolates a preloaded Lua hook and its server:Query bridge.
|
// BenchmarkLuaDNSQuery measures a preloaded DNS script using server:Query.
|
||||||
// The direct case measures the same DNS client without Lua.
|
// The direct case measures the same DNS client without Lua.
|
||||||
func BenchmarkLuaDNSHookCall(b *testing.B) {
|
func BenchmarkLuaDNSQuery(b *testing.B) {
|
||||||
option := featureDNS.IPOption{IPv4Enable: true}
|
option := featureDNS.IPOption{IPv4Enable: true}
|
||||||
ip := net.ParseIP("127.0.0.1")
|
ip := net.ParseIP("127.0.0.1")
|
||||||
upstream := &benchmarkLuaNameServer{ips: []net.IP{ip}}
|
upstream := &benchmarkLuaNameServer{ips: []net.IP{ip}}
|
||||||
@@ -271,7 +249,14 @@ end
|
|||||||
query func() ([]net.IP, uint32, error)
|
query func() ([]net.IP, uint32, error)
|
||||||
}{
|
}{
|
||||||
{"direct", func() ([]net.IP, uint32, error) { return client.QueryIP(ctx, "example.com", option) }},
|
{"direct", func() ([]net.IP, uint32, error) { return client.QueryIP(ctx, "example.com", option) }},
|
||||||
{"lua_hook", func() ([]net.IP, uint32, error) { return server.callLuaHook(L, "example.com", option) }},
|
{"lua_script", func() ([]net.IP, uint32, error) {
|
||||||
|
if err := callLuaQuery(L, "example.com", option); err != nil {
|
||||||
|
return nil, 0, err
|
||||||
|
}
|
||||||
|
ips, ttl, err := readLuaQueryResult(L)
|
||||||
|
L.Pop(3)
|
||||||
|
return ips, ttl, err
|
||||||
|
}},
|
||||||
} {
|
} {
|
||||||
b.Run(bench.name, func(b *testing.B) {
|
b.Run(bench.name, func(b *testing.B) {
|
||||||
b.ReportAllocs()
|
b.ReportAllocs()
|
||||||
|
|||||||
+15
-11
@@ -15,7 +15,6 @@ import (
|
|||||||
const scriptExecutionTimeout = 6 * time.Second
|
const scriptExecutionTimeout = 6 * time.Second
|
||||||
|
|
||||||
type scriptEngine struct {
|
type scriptEngine struct {
|
||||||
dns *DNS
|
|
||||||
pool *xlua.Pool
|
pool *xlua.Pool
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -24,8 +23,8 @@ func newScriptEngine(path string, server *DNS) (*scriptEngine, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
e := &scriptEngine{dns: server}
|
|
||||||
e.pool, err = xlua.NewPool(server.ctx, scriptExecutionTimeout, program.NewStateFactory(
|
pool, err := xlua.NewPool(server.ctx, scriptExecutionTimeout, program.NewStateFactory(
|
||||||
scriptExecutionTimeout*20,
|
scriptExecutionTimeout*20,
|
||||||
func(L *lua.LState) {
|
func(L *lua.LState) {
|
||||||
geodata.RegisterLua(L)
|
geodata.RegisterLua(L)
|
||||||
@@ -41,19 +40,24 @@ func newScriptEngine(path string, server *DNS) (*scriptEngine, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
errors.LogInfo(server.ctx, "DNS script initialized from ", path)
|
errors.LogInfo(server.ctx, "DNS script initialized from ", path)
|
||||||
return e, nil
|
return &scriptEngine{pool: pool}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (e *scriptEngine) close() {
|
func (e *scriptEngine) close() {
|
||||||
e.pool.Close()
|
e.pool.Close()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (e *scriptEngine) query(domain string, option dns.IPOption) (ips []net.IP, ttl uint32, err error) {
|
func (e *scriptEngine) query(domain string, option dns.IPOption) (ips []net.IP, ttl uint32, queryErr error) {
|
||||||
err = e.pool.WithState(nil, 0, func(L *lua.LState) error {
|
if err := e.pool.WithState(nil, 0, func(L *lua.LState) error {
|
||||||
var hookErr error
|
if err := callLuaQuery(L, domain, option); err != nil {
|
||||||
ips, ttl, hookErr = e.dns.callLuaHook(L, domain, option)
|
return err
|
||||||
return hookErr
|
}
|
||||||
})
|
ips, ttl, queryErr = readLuaQueryResult(L)
|
||||||
return
|
return nil
|
||||||
|
}); err != nil {
|
||||||
|
return nil, 0, err
|
||||||
|
}
|
||||||
|
return ips, ttl, queryErr
|
||||||
}
|
}
|
||||||
|
|||||||
+95
-12
@@ -2,6 +2,7 @@ package dns
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
go_errors "errors"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
@@ -12,21 +13,25 @@ import (
|
|||||||
featureDNS "github.com/xtls/xray-core/features/dns"
|
featureDNS "github.com/xtls/xray-core/features/dns"
|
||||||
)
|
)
|
||||||
|
|
||||||
type geoIPScriptNameServer struct {
|
type scriptNameServer struct {
|
||||||
name string
|
name string
|
||||||
answers map[string]net.IP
|
answers map[string]net.IP
|
||||||
|
errors map[string]error
|
||||||
ttl uint32
|
ttl uint32
|
||||||
calls int
|
calls int
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *geoIPScriptNameServer) Name() string { return s.name }
|
func (s *scriptNameServer) Name() string { return s.name }
|
||||||
func (s *geoIPScriptNameServer) IsDisableCache() bool { return true }
|
func (s *scriptNameServer) IsDisableCache() bool { return true }
|
||||||
|
|
||||||
func (s *geoIPScriptNameServer) QueryIP(ctx context.Context, domain string, _ featureDNS.IPOption) ([]net.IP, uint32, error) {
|
func (s *scriptNameServer) QueryIP(ctx context.Context, domain string, _ featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||||
if err := ctx.Err(); err != nil {
|
if err := ctx.Err(); err != nil {
|
||||||
return nil, 0, err
|
return nil, 0, err
|
||||||
}
|
}
|
||||||
s.calls++
|
s.calls++
|
||||||
|
if err := s.errors[domain]; err != nil {
|
||||||
|
return nil, 0, err
|
||||||
|
}
|
||||||
ip, ok := s.answers[domain]
|
ip, ok := s.answers[domain]
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, 0, featureDNS.ErrEmptyResponse
|
return nil, 0, featureDNS.ErrEmptyResponse
|
||||||
@@ -34,6 +39,88 @@ func (s *geoIPScriptNameServer) QueryIP(ctx context.Context, domain string, _ fe
|
|||||||
return []net.IP{ip}, s.ttl, nil
|
return []net.IP{ip}, s.ttl, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestDNSScriptQuery(t *testing.T) {
|
||||||
|
wantIP := net.ParseIP("127.0.0.1")
|
||||||
|
upstreamErr := go_errors.New("upstream failed")
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name, body string
|
||||||
|
wantIPs []net.IP
|
||||||
|
wantTTL uint32
|
||||||
|
wantErr error
|
||||||
|
wantMessage string
|
||||||
|
wantCalls uint32
|
||||||
|
}{
|
||||||
|
{name: "IPs", body: `return server:Query(domain, ipv4, ipv6, fake)`, wantIPs: []net.IP{wantIP}, wantTTL: 60, wantCalls: 2},
|
||||||
|
{name: "empty result", body: `return nil, 0`, wantErr: featureDNS.ErrEmptyResponse, wantCalls: 2},
|
||||||
|
{name: "upstream error", body: `return server:Query("failed.example", ipv4, ipv6, fake)`, wantErr: upstreamErr, wantCalls: 2},
|
||||||
|
{name: "string error", body: `return nil, nil, "blocked"`, wantMessage: "blocked", wantCalls: 2},
|
||||||
|
{name: "invalid result", body: `return false, 0`, wantMessage: "native IP slice", wantCalls: 2},
|
||||||
|
{name: "execution error", body: `error("execution failed")`, wantMessage: "execution failed", wantCalls: 1},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
script := `
|
||||||
|
local server = require("xray.dns").Servers[1]
|
||||||
|
local calls = 0
|
||||||
|
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
||||||
|
calls = calls + 1
|
||||||
|
if domain == "count.example" then
|
||||||
|
local ips, _, err = server:Query("good.example", ipv4, ipv6, fake)
|
||||||
|
return ips, calls, err
|
||||||
|
end
|
||||||
|
` + tc.body + `
|
||||||
|
end
|
||||||
|
`
|
||||||
|
path := filepath.Join(t.TempDir(), "query.lua")
|
||||||
|
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
option := featureDNS.IPOption{IPv4Enable: true}
|
||||||
|
upstream := &scriptNameServer{
|
||||||
|
name: "test",
|
||||||
|
answers: map[string]net.IP{"good.example": wantIP},
|
||||||
|
errors: map[string]error{"failed.example": upstreamErr},
|
||||||
|
ttl: 60,
|
||||||
|
}
|
||||||
|
server := &DNS{
|
||||||
|
ctx: context.Background(),
|
||||||
|
clients: []*Client{{server: upstream, ipOption: &option, timeoutMs: time.Second}},
|
||||||
|
}
|
||||||
|
engine, err := newScriptEngine(path, server)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer engine.close()
|
||||||
|
|
||||||
|
ips, ttl, err := engine.query("good.example", option)
|
||||||
|
switch {
|
||||||
|
case tc.wantErr != nil:
|
||||||
|
if err != tc.wantErr {
|
||||||
|
t.Fatalf("query error = %v, want original error %v", err, tc.wantErr)
|
||||||
|
}
|
||||||
|
case tc.wantMessage != "":
|
||||||
|
if err == nil || !strings.Contains(err.Error(), tc.wantMessage) {
|
||||||
|
t.Fatalf("query error = %v, want %q", err, tc.wantMessage)
|
||||||
|
}
|
||||||
|
case err != nil:
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if ttl != tc.wantTTL || len(ips) != len(tc.wantIPs) {
|
||||||
|
t.Fatalf("query = %v, TTL %d; want %v, TTL %d", ips, ttl, tc.wantIPs, tc.wantTTL)
|
||||||
|
}
|
||||||
|
for i := range ips {
|
||||||
|
if !ips[i].Equal(tc.wantIPs[i]) {
|
||||||
|
t.Fatalf("IP %d = %v, want %v", i, ips[i], tc.wantIPs[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
ips, calls, err := engine.query("count.example", option)
|
||||||
|
if err != nil || calls != tc.wantCalls || len(ips) != 1 || !ips[0].Equal(wantIP) {
|
||||||
|
t.Fatalf("next query = %v, calls %d, %v; want %v, calls %d", ips, calls, err, wantIP, tc.wantCalls)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestDNSScriptGeoIPFallback(t *testing.T) {
|
func TestDNSScriptGeoIPFallback(t *testing.T) {
|
||||||
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
||||||
script := `
|
script := `
|
||||||
@@ -59,7 +146,7 @@ end
|
|||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
primary := &geoIPScriptNameServer{
|
primary := &scriptNameServer{
|
||||||
name: "primary",
|
name: "primary",
|
||||||
answers: map[string]net.IP{
|
answers: map[string]net.IP{
|
||||||
"us.example": net.ParseIP("2001:4860:4860::8888"),
|
"us.example": net.ParseIP("2001:4860:4860::8888"),
|
||||||
@@ -67,7 +154,7 @@ end
|
|||||||
},
|
},
|
||||||
ttl: 30,
|
ttl: 30,
|
||||||
}
|
}
|
||||||
fallback := &geoIPScriptNameServer{
|
fallback := &scriptNameServer{
|
||||||
name: "fallback",
|
name: "fallback",
|
||||||
answers: map[string]net.IP{"other.example": net.ParseIP("9.9.9.9")},
|
answers: map[string]net.IP{"other.example": net.ParseIP("9.9.9.9")},
|
||||||
ttl: 60,
|
ttl: 60,
|
||||||
@@ -138,7 +225,7 @@ func TestDNSScriptRejectsInvalidStartup(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestDNSScriptHookErrorAndFakeDNSOption(t *testing.T) {
|
func TestDNSScriptFakeDNSOption(t *testing.T) {
|
||||||
path := filepath.Join(t.TempDir(), "script.lua")
|
path := filepath.Join(t.TempDir(), "script.lua")
|
||||||
script := `
|
script := `
|
||||||
local server = require("xray.dns").Servers[1]
|
local server = require("xray.dns").Servers[1]
|
||||||
@@ -146,7 +233,6 @@ local log = require("xray.log")
|
|||||||
log.Info("DNS script loaded")
|
log.Info("DNS script loaded")
|
||||||
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
||||||
log.Debug("DNS query: ", domain)
|
log.Debug("DNS query: ", domain)
|
||||||
if domain == "bad.example" then error("script failure") end
|
|
||||||
local ips, ttl, err = server:Query(domain, ipv4, ipv6, fake)
|
local ips, ttl, err = server:Query(domain, ipv4, ipv6, fake)
|
||||||
if err then log.Error("DNS failed: ", err) end
|
if err then log.Error("DNS failed: ", err) end
|
||||||
return ips, ttl, err
|
return ips, ttl, err
|
||||||
@@ -160,7 +246,7 @@ end
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
upstream := &geoIPScriptNameServer{
|
upstream := &scriptNameServer{
|
||||||
name: "FakeDNS",
|
name: "FakeDNS",
|
||||||
answers: map[string]net.IP{"good.example": net.ParseIP("198.18.0.1")},
|
answers: map[string]net.IP{"good.example": net.ParseIP("198.18.0.1")},
|
||||||
ttl: 30,
|
ttl: 30,
|
||||||
@@ -177,9 +263,6 @@ end
|
|||||||
}
|
}
|
||||||
defer server.Close()
|
defer server.Close()
|
||||||
|
|
||||||
if _, _, err := server.LookupIP("bad.example", option); err == nil || !strings.Contains(err.Error(), "script failure") {
|
|
||||||
t.Fatalf("hook failure = %v, want script failure", err)
|
|
||||||
}
|
|
||||||
if _, _, err := server.LookupIP("good.example", option); err != featureDNS.ErrEmptyResponse {
|
if _, _, err := server.LookupIP("good.example", option); err != featureDNS.ErrEmptyResponse {
|
||||||
t.Fatalf("FakeDNS without FakeEnable = %v, want ErrEmptyResponse", err)
|
t.Fatalf("FakeDNS without FakeEnable = %v, want ErrEmptyResponse", err)
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user