From 72d9ab50b947ae6edee055ba833ebb5e79a500ed Mon Sep 17 00:00:00 2001 From: Meo597 <197331664+Meo597@users.noreply.github.com> Date: Sat, 26 Sep 2026 06:16:54 +0800 Subject: [PATCH] add tests --- app/dns/lua_test.go | 87 ++++++++++++++++ app/dns/script_test.go | 192 ++++++++++++++++++++++++++++++++++++ common/geodata/lua_test.go | 19 ++++ common/lua/pool_test.go | 197 +++++++++++++++++++++++++++++++++++++ common/lua/program_test.go | 55 +++++++++++ infra/conf/dns_test.go | 50 ++++++++++ 6 files changed, 600 insertions(+) create mode 100644 app/dns/script_test.go create mode 100644 common/lua/pool_test.go create mode 100644 common/lua/program_test.go diff --git a/app/dns/lua_test.go b/app/dns/lua_test.go index 610c6dcc7..46e9e3f5f 100644 --- a/app/dns/lua_test.go +++ b/app/dns/lua_test.go @@ -2,7 +2,10 @@ package dns import ( "context" + go_errors "errors" + "strings" "testing" + "time" "github.com/xtls/xray-core/common/net" featureDNS "github.com/xtls/xray-core/features/dns" @@ -30,6 +33,90 @@ func TestDecodeLuaDNSResultNativeIP(t *testing.T) { } } +func TestDecodeLuaDNSResultNativeSliceCopiesIP(t *testing.T) { + L := lua.NewState() + defer L.Close() + original := net.ParseIP("8.8.8.8") + 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) + } + 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]) + } +} + +func TestDecodeLuaDNSResultValidation(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) + 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"}, + } { + 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) + if err == nil || !strings.Contains(err.Error(), tc.want) { + t.Fatalf("decodeLuaDNSResult 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") + } + 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) + } +} + +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 { + 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() diff --git a/app/dns/script_test.go b/app/dns/script_test.go new file mode 100644 index 000000000..835fe493e --- /dev/null +++ b/app/dns/script_test.go @@ -0,0 +1,192 @@ +package dns + +import ( + "context" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/xtls/xray-core/common/net" + featureDNS "github.com/xtls/xray-core/features/dns" +) + +type geoIPScriptNameServer struct { + name string + answers map[string]net.IP + ttl uint32 + calls int +} + +func (s *geoIPScriptNameServer) Name() string { return s.name } +func (s *geoIPScriptNameServer) IsDisableCache() bool { return true } + +func (s *geoIPScriptNameServer) QueryIP(ctx context.Context, domain string, _ featureDNS.IPOption) ([]net.IP, uint32, error) { + if err := ctx.Err(); err != nil { + return nil, 0, err + } + s.calls++ + ip, ok := s.answers[domain] + if !ok { + return nil, 0, featureDNS.ErrEmptyResponse + } + return []net.IP{ip}, s.ttl, nil +} + +func TestDNSScriptGeoIPFallback(t *testing.T) { + t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources")) + script := ` +local servers = require("xray.dns").servers +local us_ips = require("xray.geodata").ipMatcher({"geoip:us"}) + +local by_id = {} +for _, server in ipairs(servers) do + by_id[server.id] = server +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 + end + return by_id.fallback:query(q) +end +` + scriptPath := filepath.Join(t.TempDir(), "geoip_fallback.lua") + if err := os.WriteFile(scriptPath, []byte(script), 0o600); err != nil { + t.Fatal(err) + } + + primary := &geoIPScriptNameServer{ + name: "primary", + answers: map[string]net.IP{ + "us.example": net.ParseIP("2001:4860:4860::8888"), + "other.example": net.ParseIP("127.0.0.1"), + }, + ttl: 30, + } + fallback := &geoIPScriptNameServer{ + name: "fallback", + answers: map[string]net.IP{"other.example": net.ParseIP("9.9.9.9")}, + ttl: 60, + } + option := featureDNS.IPOption{IPv4Enable: true, IPv6Enable: true} + hosts, err := NewStaticHosts(nil) + if err != nil { + t.Fatal(err) + } + server := &DNS{ + ctx: context.Background(), + hosts: hosts, + ipOption: &option, + scriptPath: scriptPath, + clients: []*Client{ + {id: "primary", server: primary, ipOption: &option, timeoutMs: 2 * time.Second}, + {id: "fallback", server: fallback, ipOption: &option, timeoutMs: 2 * time.Second}, + }, + } + if err := server.Start(); err != nil { + t.Fatal(err) + } + defer server.Close() + + for _, tc := range []struct { + domain string + ip net.IP + ttl uint32 + }{ + {"Us.Example.", net.ParseIP("2001:4860:4860::8888"), 30}, + {"other.example", net.ParseIP("9.9.9.9"), 60}, + } { + ips, ttl, err := server.LookupIP(tc.domain, option) + if err != nil { + t.Fatalf("LookupIP(%q): %v", tc.domain, err) + } + if ttl != tc.ttl || len(ips) != 1 || !ips[0].Equal(tc.ip) { + t.Fatalf("LookupIP(%q) = %v, TTL %d; want %v, TTL %d", tc.domain, ips, ttl, tc.ip, tc.ttl) + } + } + if primary.calls != 2 || fallback.calls != 1 { + t.Fatalf("upstream calls: primary %d, fallback %d; want 2 and 1", primary.calls, fallback.calls) + } +} + +func TestDNSScriptRejectsInvalidStartup(t *testing.T) { + for _, tc := range []struct { + name string + script string + }{ + {"syntax", "function handleDNSQuery("}, + {"missing hook", "value = 1"}, + {"top-level error", `error("setup failed")`}, + } { + t.Run(tc.name, func(t *testing.T) { + path := filepath.Join(t.TempDir(), "script.lua") + if err := os.WriteFile(path, []byte(tc.script), 0o600); err != nil { + t.Fatal(err) + } + server := &DNS{ctx: context.Background(), scriptPath: path} + if err := server.Start(); err == nil { + t.Fatal("Start accepted an invalid DNS script") + } + if server.script != nil { + t.Fatal("Start retained a script engine after failure") + } + }) + } +} + +func TestDNSScriptHookErrorAndFakeDNSOption(t *testing.T) { + path := filepath.Join(t.TempDir(), "script.lua") + script := ` +local server = require("xray.dns").servers[1] +function handleDNSQuery(q) + if q.domain == "bad.example" then error("script failure") end + return server:query(q) +end +` + if err := os.WriteFile(path, []byte(script), 0o600); err != nil { + t.Fatal(err) + } + option := featureDNS.IPOption{IPv4Enable: true} + hosts, err := NewStaticHosts(nil) + if err != nil { + t.Fatal(err) + } + upstream := &geoIPScriptNameServer{ + name: "FakeDNS", + answers: map[string]net.IP{"good.example": net.ParseIP("198.18.0.1")}, + ttl: 30, + } + server := &DNS{ + ctx: context.Background(), + hosts: hosts, + ipOption: &option, + scriptPath: path, + clients: []*Client{{id: "fake", server: upstream, ipOption: &option, timeoutMs: time.Second}}, + } + if err := server.Start(); err != nil { + t.Fatal(err) + } + 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) + } + if upstream.calls != 0 { + t.Fatalf("FakeDNS was queried without FakeEnable: %d calls", upstream.calls) + } + withFake := featureDNS.IPOption{IPv4Enable: true, FakeEnable: true} + ips, ttl, err := server.LookupIP("good.example", withFake) + if err != nil || ttl != 30 || len(ips) != 1 || !ips[0].Equal(net.ParseIP("198.18.0.1")) { + t.Fatalf("FakeDNS with FakeEnable = %v, TTL %d, %v", ips, ttl, err) + } + if upstream.calls != 1 { + t.Fatalf("FakeDNS query count = %d, want 1", upstream.calls) + } +} diff --git a/common/geodata/lua_test.go b/common/geodata/lua_test.go index 400a6ed69..b1ecf8804 100644 --- a/common/geodata/lua_test.go +++ b/common/geodata/lua_test.go @@ -40,3 +40,22 @@ func TestLuaDomainMatcherUsesNativeMatcher(t *testing.T) { t.Fatal(err) } } + +func TestLuaMatchersRejectInvalidRules(t *testing.T) { + for _, tc := range []struct { + name string + script string + }{ + {"IP rule", `require("xray.geodata").ipMatcher({"not-an-ip"})`}, + {"non-string domain rule", `require("xray.geodata").domainMatcher({true})`}, + } { + t.Run(tc.name, func(t *testing.T) { + L := lua.NewState() + defer L.Close() + RegisterLua(L) + if err := L.DoString(tc.script); err == nil { + t.Fatal("invalid geodata rule was accepted") + } + }) + } +} diff --git a/common/lua/pool_test.go b/common/lua/pool_test.go new file mode 100644 index 000000000..7b4d45bbe --- /dev/null +++ b/common/lua/pool_test.go @@ -0,0 +1,197 @@ +package lua + +import ( + "context" + "errors" + "testing" + "time" + + glua "github.com/yuin/gopher-lua" +) + +func TestPoolFactoryFailureClosesReturnedState(t *testing.T) { + failure := errors.New("factory failed") + state := glua.NewState() + _, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) { + return state, failure + }) + if !errors.Is(err, failure) || !state.IsClosed() { + t.Fatalf("NewPool error = %v, state closed = %t", err, state.IsClosed()) + } + + var failedState *glua.LState + calls := 0 + pool, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) { + calls++ + if calls == 1 { + return glua.NewState(), nil + } + failedState = glua.NewState() + return failedState, failure + }) + if err != nil { + t.Fatal(err) + } + defer pool.Close() + borrowed, err := pool.Acquire() + if err != nil { + t.Fatal(err) + } + defer pool.Release(borrowed, true) + _, err = pool.Acquire() + if !errors.Is(err, failure) || !failedState.IsClosed() { + t.Fatalf("Acquire error = %v, state closed = %t", err, failedState.IsClosed()) + } +} + +func TestPoolReusesStatesAndLimitsIdle(t *testing.T) { + created := 0 + pool, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) { + created++ + return glua.NewState(), nil + }) + if err != nil { + t.Fatal(err) + } + defer pool.Close() + + states := make([]*glua.LState, maxIdleStates+3) + for i := range states { + states[i], err = pool.Acquire() + if err != nil { + t.Fatal(err) + } + } + for _, state := range states { + pool.Release(state, true) + } + for i, state := range states { + if got, want := state.IsClosed(), i >= maxIdleStates; got != want { + t.Fatalf("state %d closed = %t, want %t", i, got, want) + } + } + borrowed, err := pool.Acquire() + if err != nil { + t.Fatal(err) + } + if created != len(states) { + t.Fatalf("Acquire created %d states, want %d", created, len(states)) + } + pool.Release(borrowed, true) +} + +func TestPoolCloseCancelsAndWaitsForBorrowedState(t *testing.T) { + pool, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) { + return glua.NewState(), nil + }) + if err != nil { + t.Fatal(err) + } + state, err := pool.Acquire() + if err != nil { + t.Fatal(err) + } + done := make(chan struct{}) + go func() { + pool.Close() + close(done) + }() + select { + case <-pool.Context().Done(): + case <-time.After(time.Second): + pool.Release(state, false) + t.Fatal("Close did not cancel the pool context") + } + select { + case <-done: + pool.Release(state, false) + t.Fatal("Close returned while a state was borrowed") + default: + } + pool.Release(state, true) + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("Close did not finish after Release") + } + if !state.IsClosed() { + t.Fatal("borrowed state was not closed") + } + if _, err := pool.Acquire(); !errors.Is(err, context.Canceled) { + t.Fatalf("Acquire after Close = %v, want context.Canceled", err) + } + pool.Close() +} + +func TestPoolCloseCancelsStateCreation(t *testing.T) { + started := make(chan struct{}) + first := true + pool, err := NewPool(context.Background(), func(ctx context.Context) (*glua.LState, error) { + if first { + first = false + return glua.NewState(), nil + } + close(started) + <-ctx.Done() + return nil, ctx.Err() + }) + if err != nil { + t.Fatal(err) + } + borrowed, err := pool.Acquire() + if err != nil { + t.Fatal(err) + } + acquireDone := make(chan error, 1) + go func() { + _, err := pool.Acquire() + acquireDone <- err + }() + <-started + closeDone := make(chan struct{}) + go func() { + pool.Close() + close(closeDone) + }() + select { + case err := <-acquireDone: + if !errors.Is(err, context.Canceled) { + t.Fatalf("Acquire during Close = %v, want context.Canceled", err) + } + case <-time.After(time.Second): + t.Fatal("state creation did not stop after Close") + } + select { + case <-closeDone: + pool.Release(borrowed, false) + t.Fatal("Close returned while the initial state was borrowed") + default: + } + pool.Release(borrowed, true) + select { + case <-closeDone: + case <-time.After(time.Second): + t.Fatal("Close did not finish after Release") + } +} + +func BenchmarkPoolAcquireRelease(b *testing.B) { + pool, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) { + return glua.NewState(), nil + }) + if err != nil { + b.Fatal(err) + } + defer pool.Close() + + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + state, err := pool.Acquire() + if err != nil { + b.Fatal(err) + } + pool.Release(state, true) + } + b.StopTimer() +} diff --git a/common/lua/program_test.go b/common/lua/program_test.go new file mode 100644 index 000000000..d60b892fc --- /dev/null +++ b/common/lua/program_test.go @@ -0,0 +1,55 @@ +package lua + +import ( + "context" + "os" + "path/filepath" + "testing" + + glua "github.com/yuin/gopher-lua" +) + +func TestProgramStatesAreIndependent(t *testing.T) { + path := filepath.Join(t.TempDir(), "state.lua") + if err := os.WriteFile(path, []byte("value = (value or 0) + 1"), 0o600); err != nil { + t.Fatal(err) + } + program, err := CompileFile(path) + if err != nil { + t.Fatal(err) + } + first, err := program.NewState(context.Background(), nil) + if err != nil { + t.Fatal(err) + } + defer first.Close() + first.SetGlobal("value", glua.LNumber(42)) + second, err := program.NewState(context.Background(), nil) + if err != nil { + t.Fatal(err) + } + defer second.Close() + if got := second.GetGlobal("value"); got != glua.LNumber(1) { + t.Fatalf("second state value = %v, want 1", got) + } +} + +func TestProgramInitializationObservesCancellation(t *testing.T) { + path := filepath.Join(t.TempDir(), "loop.lua") + if err := os.WriteFile(path, []byte("while true do end"), 0o600); err != nil { + t.Fatal(err) + } + program, err := CompileFile(path) + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + cancel() + state, err := program.NewState(ctx, nil) + if err == nil || state != nil { + if state != nil { + state.Close() + } + t.Fatalf("NewState with canceled context = %v, %v; want nil state and error", state, err) + } +} diff --git a/infra/conf/dns_test.go b/infra/conf/dns_test.go index 278f34c38..d90ddbfce 100644 --- a/infra/conf/dns_test.go +++ b/infra/conf/dns_test.go @@ -2,6 +2,8 @@ package conf_test import ( "encoding/json" + "os" + "path/filepath" "testing" "github.com/google/go-cmp/cmp" @@ -122,3 +124,51 @@ func TestDNSConfigParsing(t *testing.T) { } } } + +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 { + t.Fatal(err) + } + + for _, tc := range []struct { + name string + script string + wantError bool + }{ + {"relative", "lookup.lua", false}, + {"absolute", path, false}, + {"missing", "missing.lua", true}, + {"directory", dir, true}, + } { + t.Run(tc.name, func(t *testing.T) { + built, err := (&DNSConfig{Script: tc.script}).Build() + if tc.wantError { + if err == nil { + t.Fatal("Build accepted an invalid script path") + } + return + } + if err != nil { + t.Fatal(err) + } + if built.Script != path { + t.Fatalf("script path = %q, want %q", built.Script, path) + } + }) + } + + var parsed DNSConfig + if err := json.Unmarshal([]byte(`{"servers":[{"id":"primary","address":"1.1.1.1"}]}`), &parsed); err != nil { + t.Fatal(err) + } + built, err := parsed.Build() + if err != nil { + t.Fatal(err) + } + if len(built.NameServer) != 1 || built.NameServer[0].Id != "primary" { + t.Fatalf("nameserver IDs = %v, want primary", built.NameServer) + } +}