feat(dns): add Lua scripting for DNS queries

This commit is contained in:
Meo597
2026-09-26 05:37:23 +08:00
parent 60e2a0c502
commit 235843c5d2
15 changed files with 681 additions and 8 deletions
+25 -6
View File
@@ -93,6 +93,7 @@ type NameServer struct {
UnexpectedIp []*geodata.IPRule `protobuf:"bytes,13,rep,name=unexpected_ip,json=unexpectedIp,proto3" json:"unexpected_ip,omitempty"`
ActUnprior bool `protobuf:"varint,14,opt,name=actUnprior,proto3" json:"actUnprior,omitempty"`
PolicyID uint32 `protobuf:"varint,17,opt,name=policyID,proto3" json:"policyID,omitempty"`
Id string `protobuf:"bytes,18,opt,name=id,proto3" json:"id,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
@@ -239,6 +240,13 @@ func (x *NameServer) GetPolicyID() uint32 {
return 0
}
func (x *NameServer) GetId() string {
if x != nil {
return x.Id
}
return ""
}
type Config struct {
state protoimpl.MessageState `protogen:"open.v1"`
// NameServer list used by this DNS client.
@@ -258,8 +266,10 @@ type Config struct {
DisableFallback bool `protobuf:"varint,10,opt,name=disableFallback,proto3" json:"disableFallback,omitempty"`
DisableFallbackIfMatch bool `protobuf:"varint,11,opt,name=disableFallbackIfMatch,proto3" json:"disableFallbackIfMatch,omitempty"`
EnableParallelQuery bool `protobuf:"varint,14,opt,name=enableParallelQuery,proto3" json:"enableParallelQuery,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
// Absolute path to the Lua DNS query script.
Script string `protobuf:"bytes,15,opt,name=script,proto3" json:"script,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *Config) Reset() {
@@ -369,6 +379,13 @@ func (x *Config) GetEnableParallelQuery() bool {
return false
}
func (x *Config) GetScript() string {
if x != nil {
return x.Script
}
return ""
}
type Config_HostMapping struct {
state protoimpl.MessageState `protogen:"open.v1"`
Domain *geodata.DomainRule `protobuf:"bytes,2,opt,name=domain,proto3" json:"domain,omitempty"`
@@ -435,7 +452,7 @@ var File_app_dns_config_proto protoreflect.FileDescriptor
const file_app_dns_config_proto_rawDesc = "" +
"\n" +
"\x14app/dns/config.proto\x12\fxray.app.dns\x1a\x1ccommon/net/destination.proto\x1a\x1bcommon/geodata/geodat.proto\"\xde\x05\n" +
"\x14app/dns/config.proto\x12\fxray.app.dns\x1a\x1ccommon/net/destination.proto\x1a\x1bcommon/geodata/geodat.proto\"\xee\x05\n" +
"\n" +
"NameServer\x123\n" +
"\aaddress\x18\x01 \x01(\v2\x19.xray.common.net.EndpointR\aaddress\x12\x1b\n" +
@@ -461,10 +478,11 @@ const file_app_dns_config_proto_rawDesc = "" +
"\n" +
"actUnprior\x18\x0e \x01(\bR\n" +
"actUnprior\x12\x1a\n" +
"\bpolicyID\x18\x11 \x01(\rR\bpolicyIDB\x0f\n" +
"\bpolicyID\x18\x11 \x01(\rR\bpolicyID\x12\x0e\n" +
"\x02id\x18\x12 \x01(\tR\x02idB\x0f\n" +
"\r_disableCacheB\r\n" +
"\v_serveStaleB\x12\n" +
"\x10_serveExpiredTTLJ\x04\b\x04\x10\x05\"\x82\x05\n" +
"\x10_serveExpiredTTLJ\x04\b\x04\x10\x05\"\x9a\x05\n" +
"\x06Config\x129\n" +
"\vname_server\x18\x05 \x03(\v2\x18.xray.app.dns.NameServerR\n" +
"nameServer\x12\x1b\n" +
@@ -480,7 +498,8 @@ const file_app_dns_config_proto_rawDesc = "" +
"\x0fdisableFallback\x18\n" +
" \x01(\bR\x0fdisableFallback\x126\n" +
"\x16disableFallbackIfMatch\x18\v \x01(\bR\x16disableFallbackIfMatch\x120\n" +
"\x13enableParallelQuery\x18\x0e \x01(\bR\x13enableParallelQuery\x1a}\n" +
"\x13enableParallelQuery\x18\x0e \x01(\bR\x13enableParallelQuery\x12\x16\n" +
"\x06script\x18\x0f \x01(\tR\x06script\x1a}\n" +
"\vHostMapping\x127\n" +
"\x06domain\x18\x02 \x01(\v2\x1f.xray.common.geodata.DomainRuleR\x06domain\x12\x0e\n" +
"\x02ip\x18\x03 \x03(\fR\x02ip\x12%\n" +
+4
View File
@@ -27,6 +27,7 @@ message NameServer {
repeated xray.common.geodata.IPRule unexpected_ip = 13;
bool actUnprior = 14;
uint32 policyID = 17;
string id = 18;
}
enum QueryStrategy {
@@ -73,4 +74,7 @@ message Config {
bool disableFallbackIfMatch = 11;
bool enableParallelQuery = 14;
// Absolute path to the Lua DNS query script.
string script = 15;
}
+16
View File
@@ -31,6 +31,8 @@ type DNS struct {
domainMatcher geodata.DomainMatcher
matcherInfos []*DomainMatcherInfo
checkSystem bool
script *scriptEngine
scriptPath string
}
// DomainMatcherInfo contains information attached to index returned by Server.domainMatcher.
@@ -180,6 +182,7 @@ func New(ctx context.Context, config *Config) (*DNS, error) {
disableFallbackIfMatch: config.DisableFallbackIfMatch,
enableParallelQuery: config.EnableParallelQuery,
checkSystem: checkSystem,
scriptPath: config.Script,
}, nil
}
@@ -190,11 +193,21 @@ func (*DNS) Type() interface{} {
// Start implements common.Runnable.
func (s *DNS) Start() error {
if s.scriptPath != "" {
engine, err := newScriptEngine(s.scriptPath, s)
if err != nil {
return errors.New("failed to initialize DNS script").Base(err)
}
s.script = engine
}
return nil
}
// Close implements common.Closable.
func (s *DNS) Close() error {
if s.script != nil {
s.script.close()
}
return nil
}
@@ -257,6 +270,9 @@ func (s *DNS) LookupIP(domain string, option dns.IPOption) ([]net.IP, uint32, er
}
// Name servers lookup
if s.script != nil {
return s.script.query(domain, option)
}
if s.enableParallelQuery {
return s.parallelQuery(domain, option)
} else {
+180
View File
@@ -0,0 +1,180 @@
package dns
import (
"context"
"math"
"strings"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
featureDNS "github.com/xtls/xray-core/features/dns"
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.
func (s *DNS) RegisterLua(L *lua.LState) {
L.PreloadModule("xray.dns", func(L *lua.LState) int {
servers := L.NewTable()
for i, client := range s.clients {
server := L.NewTable()
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)
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,
}
ctx := L.Context()
if ctx == nil {
L.RaiseError("server:query requires an active DNS query")
return 0
}
var ips []net.IP
var ttl uint32
var err error
if !option.FakeEnable && strings.EqualFold(client.Name(), "FakeDNS") {
err = featureDNS.ErrEmptyResponse
} else {
ips, ttl, err = client.QueryIP(ctx, string(domain), option)
}
result := L.NewTable()
addresses := L.NewTable()
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))
if err != nil {
ud := L.NewUserData()
ud.Value = err
result.RawSetString("error", ud)
}
L.Push(result)
return 1
}))
servers.RawSetInt(i+1, server)
}
module := L.NewTable()
module.RawSetString("servers", servers)
L.Push(module)
return 1
})
}
// 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.
func (s *DNS) CallLuaHook(L *lua.LState, ctx context.Context, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
q := L.NewTable()
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() {
if previous == nil {
L.RemoveContext()
} else {
L.SetContext(previous)
}
}()
fn := L.GetGlobal("handleDNSQuery")
if fn.Type() != lua.LTFunction {
return nil, 0, errors.New("DNS script must define handleDNSQuery(q)")
}
if err := L.CallByParam(lua.P{Fn: fn, NRet: 1, Protect: true}, q); err != nil {
return nil, 0, err
}
value := L.Get(-1)
L.Pop(1)
ips, ttl, err := decodeLuaDNSResult(value, option)
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 {
if err, ok := ud.Value.(error); ok {
return nil, 0, err
}
}
if s, ok := v.(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) {
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 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
}
+54
View File
@@ -0,0 +1,54 @@
package dns
import (
"context"
"testing"
"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) {
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 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)
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}
end
`); err != nil {
t.Fatal(err)
}
s := &DNS{}
if _, _, err := s.CallLuaHook(L, context.Background(), "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil {
t.Fatal(err)
}
}
+2 -1
View File
@@ -29,6 +29,7 @@ type Server interface {
// Client is the interface for DNS client.
type Client struct {
id string
server Server
skipFallback bool
expectedIPs geodata.IPMatcher
@@ -97,7 +98,7 @@ func NewClient(
ipOption dns.IPOption,
updateRules func(bool),
) (*Client, error) {
client := &Client{}
client := &Client{id: ns.Id}
err := core.RequireFeatures(ctx, func(dispatcher routing.Dispatcher) error {
// Create a new server for each client for now
server, err := NewServer(ctx, ns.Address.AsDestination(), dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP)
+1 -1
View File
@@ -49,5 +49,5 @@ func NewLocalNameServer() *LocalNameServer {
// NewLocalDNSClient creates localdns client object for directly lookup in system DNS.
func NewLocalDNSClient(ipOption dns.IPOption) *Client {
return &Client{server: NewLocalNameServer(), ipOption: &ipOption}
return &Client{id: "localhost", server: NewLocalNameServer(), ipOption: &ipOption}
}
+71
View File
@@ -0,0 +1,71 @@
package dns
import (
"context"
"time"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/geodata"
luamgr "github.com/xtls/xray-core/common/lua"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/features/dns"
lua "github.com/yuin/gopher-lua"
)
const scriptExecutionTimeout = 10 * time.Second
type scriptEngine struct {
dns *DNS
pool *luamgr.Pool
}
func newScriptEngine(path string, server *DNS) (*scriptEngine, error) {
program, err := luamgr.CompileFile(path)
if err != nil {
return nil, err
}
e := &scriptEngine{dns: server}
e.pool, err = luamgr.NewPool(server.ctx, func(poolCtx context.Context) (*lua.LState, error) {
initCtx, cancel := context.WithTimeout(poolCtx, scriptExecutionTimeout)
defer cancel()
L, err := program.NewState(initCtx, func(L *lua.LState) {
geodata.RegisterLua(L)
server.RegisterLua(L)
})
if err != nil {
return nil, err
}
if L.GetGlobal("handleDNSQuery").Type() != lua.LTFunction {
L.Close()
return nil, errors.New("DNS script must define handleDNSQuery(q)")
}
return L, nil
})
if err != nil {
return nil, err
}
errors.LogInfo(server.ctx, "DNS script initialized from ", path)
return e, nil
}
func (e *scriptEngine) close() {
e.pool.Close()
}
func (e *scriptEngine) query(domain string, option dns.IPOption) ([]net.IP, uint32, error) {
L, err := e.pool.Acquire()
if err != nil {
return nil, 0, err
}
reusable := false
defer func() {
e.pool.Release(L, reusable)
}()
queryCtx, cancel := context.WithTimeout(e.pool.Context(), scriptExecutionTimeout)
defer cancel()
ips, ttl, err := e.dns.CallLuaHook(L, queryCtx, domain, option)
if err == nil {
reusable = true
}
return ips, ttl, err
}