diff --git a/app/dns/lua.go b/app/dns/lua.go index 8466d1351..ab31471da 100644 --- a/app/dns/lua.go +++ b/app/dns/lua.go @@ -126,31 +126,32 @@ func newLuaClientQuery(L *lua.LState, client featureDNS.Client) *lua.LFunction { }) } -// 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, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) { - top := L.GetTop() - defer L.SetTop(top) +// callLuaQuery runs HandleDNSQuery and leaves (ips, ttl, err) on the stack. +func callLuaQuery(L *lua.LState, domain string, option featureDNS.IPOption) error { fn := L.GetGlobal("HandleDNSQuery") 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), - lua.LBool(option.IPv6Enable), lua.LBool(option.FakeEnable)); err != nil { - return nil, 0, err - } - return readLuaDNSResult(L.Get(-3), L.Get(-2), L.Get(-1)) + + return 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)) } -func readLuaDNSResult(addresses, ttlValue, errorValue lua.LValue) ([]net.IP, uint32, error) { - if err := xlua.ReadError(errorValue, "DNS script error must be an error or string"); err != nil { +// readLuaQueryResult reads (ips, ttl, err) from the stack without copying the IPs. +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 } - 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 { return nil, 0, err } + + addresses := L.Get(-3) if addresses == lua.LNil { return nil, 0, featureDNS.ErrEmptyResponse } @@ -161,5 +162,6 @@ func readLuaDNSResult(addresses, ttlValue, errorValue lua.LValue) ([]net.IP, uin if len(ips) == 0 { return nil, 0, featureDNS.ErrEmptyResponse } + return ips, ttl, nil } diff --git a/app/dns/lua_test.go b/app/dns/lua_test.go index 0648c0305..305e3588a 100644 --- a/app/dns/lua_test.go +++ b/app/dns/lua_test.go @@ -3,7 +3,6 @@ package dns import ( "context" go_errors "errors" - "math" "strings" "testing" "time" @@ -15,70 +14,74 @@ import ( 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() - +func TestReadLuaQueryResult(t *testing.T) { + wantIPs := []net.IP{net.ParseIP("8.8.8.8"), {127, 0, 0, 1}, net.ParseIP("::1")} + nativeErr := go_errors.New("upstream failed") for _, tc := range []struct { - name string - change func(*[3]lua.LValue) - want string + name, values string + wantIPs []net.IP + wantTTL uint32 + wantErr error + wantMessage 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"}, + {name: "IPs", values: `ips, 45`, wantIPs: wantIPs, wantTTL: 45}, + {name: "nil IPs", values: `nil, 0`, wantErr: featureDNS.ErrEmptyResponse}, + {name: "empty IPs", values: `emptyIPs, 0`, wantErr: featureDNS.ErrEmptyResponse}, + {name: "native error", values: `nil, nil, nativeError`, wantErr: nativeErr}, + {name: "string error", values: `nil, nil, "blocked"`, wantMessage: "blocked"}, + {name: "fractional TTL", values: `ips, 1.5`, wantMessage: "invalid TTL"}, + {name: "oversized TTL", values: `ips, 4294967296`, wantMessage: "invalid TTL"}, + {name: "negative TTL", values: `ips, -1`, wantMessage: "invalid TTL"}, + {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) { - 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) + L := lua.NewState() + defer L.Close() + for name, value := range map[string]any{"ips": wantIPs, "ip": wantIPs[0], "emptyIPs": []net.IP(nil), "nativeError": nativeErr} { + ud := L.NewUserData() + ud.Value = value + L.SetGlobal(name, ud) + } + 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() defer L.Close() 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()) cancel() 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 { - t.Fatal("CallLuaHook did not stop after context cancellation") + t.Fatal("callLuaQuery did not stop after context cancellation") } 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() defer L.Close() addresses := L.NewUserData() @@ -111,39 +114,11 @@ func TestCallLuaHookNormalizesDomain(t *testing.T) { `); err != nil { t.Fatal(err) } - s := &DNS{} - if _, _, err := s.callLuaHook(L, "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil { + if err := callLuaQuery(L, "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil { t.Fatal(err) } -} - -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") - } - }) + 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") } } @@ -169,7 +144,10 @@ end t.Fatal(err) } 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]) { 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 } -// 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. -func BenchmarkLuaDNSHookCall(b *testing.B) { +func BenchmarkLuaDNSQuery(b *testing.B) { option := featureDNS.IPOption{IPv4Enable: true} ip := net.ParseIP("127.0.0.1") upstream := &benchmarkLuaNameServer{ips: []net.IP{ip}} @@ -271,7 +249,14 @@ end 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, "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.ReportAllocs() diff --git a/app/dns/script.go b/app/dns/script.go index 7fc80507f..012d44f7a 100644 --- a/app/dns/script.go +++ b/app/dns/script.go @@ -15,7 +15,6 @@ import ( const scriptExecutionTimeout = 6 * time.Second type scriptEngine struct { - dns *DNS pool *xlua.Pool } @@ -24,8 +23,8 @@ func newScriptEngine(path string, server *DNS) (*scriptEngine, error) { if err != nil { 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, func(L *lua.LState) { geodata.RegisterLua(L) @@ -41,19 +40,24 @@ func newScriptEngine(path string, server *DNS) (*scriptEngine, error) { if err != nil { return nil, err } + errors.LogInfo(server.ctx, "DNS script initialized from ", path) - return e, nil + return &scriptEngine{pool: pool}, nil } func (e *scriptEngine) close() { e.pool.Close() } -func (e *scriptEngine) query(domain string, option dns.IPOption) (ips []net.IP, ttl uint32, err error) { - err = e.pool.WithState(nil, 0, func(L *lua.LState) error { - var hookErr error - ips, ttl, hookErr = e.dns.callLuaHook(L, domain, option) - return hookErr - }) - return +func (e *scriptEngine) query(domain string, option dns.IPOption) (ips []net.IP, ttl uint32, queryErr error) { + if err := e.pool.WithState(nil, 0, func(L *lua.LState) error { + if err := callLuaQuery(L, domain, option); err != nil { + return err + } + ips, ttl, queryErr = readLuaQueryResult(L) + return nil + }); err != nil { + return nil, 0, err + } + return ips, ttl, queryErr } diff --git a/app/dns/script_test.go b/app/dns/script_test.go index 4ddcf8ef5..cea0bb1c7 100644 --- a/app/dns/script_test.go +++ b/app/dns/script_test.go @@ -2,6 +2,7 @@ package dns import ( "context" + go_errors "errors" "os" "path/filepath" "strings" @@ -12,21 +13,25 @@ import ( featureDNS "github.com/xtls/xray-core/features/dns" ) -type geoIPScriptNameServer struct { +type scriptNameServer struct { name string answers map[string]net.IP + errors map[string]error ttl uint32 calls int } -func (s *geoIPScriptNameServer) Name() string { return s.name } -func (s *geoIPScriptNameServer) IsDisableCache() bool { return true } +func (s *scriptNameServer) Name() string { return s.name } +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 { return nil, 0, err } s.calls++ + if err := s.errors[domain]; err != nil { + return nil, 0, err + } ip, ok := s.answers[domain] if !ok { 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 } +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) { t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources")) script := ` @@ -59,7 +146,7 @@ end t.Fatal(err) } - primary := &geoIPScriptNameServer{ + primary := &scriptNameServer{ name: "primary", answers: map[string]net.IP{ "us.example": net.ParseIP("2001:4860:4860::8888"), @@ -67,7 +154,7 @@ end }, ttl: 30, } - fallback := &geoIPScriptNameServer{ + fallback := &scriptNameServer{ name: "fallback", answers: map[string]net.IP{"other.example": net.ParseIP("9.9.9.9")}, 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") script := ` local server = require("xray.dns").Servers[1] @@ -146,7 +233,6 @@ local log = require("xray.log") log.Info("DNS script loaded") 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 @@ -160,7 +246,7 @@ end if err != nil { t.Fatal(err) } - upstream := &geoIPScriptNameServer{ + upstream := &scriptNameServer{ name: "FakeDNS", answers: map[string]net.IP{"good.example": net.ParseIP("198.18.0.1")}, ttl: 30, @@ -177,9 +263,6 @@ end } 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 { t.Fatalf("FakeDNS without FakeEnable = %v, want ErrEmptyResponse", err) }