diff --git a/app/dns/lua.go b/app/dns/lua.go index ed0295942..9e4a58ca2 100644 --- a/app/dns/lua.go +++ b/app/dns/lua.go @@ -8,33 +8,52 @@ import ( "github.com/xtls/xray-core/common/errors" "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" ) -// RegisterLua makes xray.dns available with Query and optional Servers. +// luaDNSServer adapts configured and local DNS to the same Lua API. +type luaDNSServer struct { + id string + name string + query func(context.Context, string, featureDNS.IPOption) ([]net.IP, uint32, error) +} + +// RegisterLua makes xray.dns available to scripts backed by client. func RegisterLua(L *lua.LState, client featureDNS.Client) { - // A configured DNS app passes its *DNS instance here. - if s, ok := client.(*DNS); ok { - registerLua(L, s, true) - return + var servers []luaDNSServer + switch client := client.(type) { + case *DNS: + servers = luaServers(client) + case *localdns.Client: + servers = []luaDNSServer{{ + id: "localhost", + name: "localhost", + query: func(_ context.Context, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) { + return client.LookupIP(domain, option) + }, + }} } - L.PreloadModule("xray.dns", func(L *lua.LState) int { - module := L.NewTable() - module.RawSetString("Query", newLuaClientQuery(L, client)) - L.Push(module) - return 1 - }) + registerLua(L, servers, client) } -// RegisterLua makes xray.dns available to DNS scripts with Servers so no Query. +// RegisterLua makes xray.dns available to DNS scripts. func (s *DNS) RegisterLua(L *lua.LState) { - registerLua(L, s, false) + registerLua(L, luaServers(s), nil) } -func registerLua(L *lua.LState, s *DNS, exposeQuery bool) { +func luaServers(s *DNS) []luaDNSServer { + servers := make([]luaDNSServer, len(s.clients)) + for i, client := range s.clients { + servers[i] = luaDNSServer{id: client.id, name: client.Name(), query: client.QueryIP} + } + return servers +} + +func registerLua(L *lua.LState, servers []luaDNSServer, client featureDNS.Client) { L.PreloadModule("xray.dns", func(L *lua.LState) int { - servers := L.NewTable() - for i, client := range s.clients { + serverList := L.NewTable() + for i, client := range servers { server := L.NewTable() server.RawSetString("ID", lua.LString(client.id)) @@ -58,10 +77,10 @@ func registerLua(L *lua.LState, s *DNS, exposeQuery bool) { var ips []net.IP var ttl uint32 var err error - if !option.FakeEnable && strings.EqualFold(client.Name(), "FakeDNS") { + if !option.FakeEnable && strings.EqualFold(client.name, "FakeDNS") { err = featureDNS.ErrEmptyResponse } else { - ips, ttl, err = client.QueryIP(ctx, string(domain), option) + ips, ttl, err = client.query(ctx, string(domain), option) } addresses := L.NewUserData() addresses.Value = ips @@ -76,13 +95,15 @@ func registerLua(L *lua.LState, s *DNS, exposeQuery bool) { } return 3 })) - servers.RawSetInt(i+1, server) + serverList.RawSetInt(i+1, server) } module := L.NewTable() - module.RawSetString("Servers", servers) - if exposeQuery { - module.RawSetString("Query", newLuaClientQuery(L, s)) + if servers != nil { + module.RawSetString("Servers", serverList) + } + if client != nil { + module.RawSetString("Query", newLuaClientQuery(L, client)) } L.Push(module) return 1 diff --git a/app/dns/lua_test.go b/app/dns/lua_test.go index a9ace300f..bf6c2feea 100644 --- a/app/dns/lua_test.go +++ b/app/dns/lua_test.go @@ -11,6 +11,7 @@ import ( "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" ) @@ -175,6 +176,67 @@ end } } +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 }