diff --git a/app/dns/lua.go b/app/dns/lua.go index ec1cff555..47d818e74 100644 --- a/app/dns/lua.go +++ b/app/dns/lua.go @@ -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 } diff --git a/app/dns/lua_test.go b/app/dns/lua_test.go index 5c8936a95..2d8fe6ff8 100644 --- a/app/dns/lua_test.go +++ b/app/dns/lua_test.go @@ -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) diff --git a/app/dns/script.go b/app/dns/script.go index 160f703ae..f76a96dc4 100644 --- a/app/dns/script.go +++ b/app/dns/script.go @@ -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 }) diff --git a/app/dns/script_test.go b/app/dns/script_test.go index 64070305f..ef79cd575 100644 --- a/app/dns/script_test.go +++ b/app/dns/script_test.go @@ -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 { diff --git a/infra/conf/dns_test.go b/infra/conf/dns_test.go index d90ddbfce..6ca9878f6 100644 --- a/infra/conf/dns_test.go +++ b/infra/conf/dns_test.go @@ -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) }