From 235843c5d295229ca8cf9983fb8922da2aebc2bd Mon Sep 17 00:00:00 2001 From: Meo597 <197331664+Meo597@users.noreply.github.com> Date: Sat, 26 Sep 2026 05:37:23 +0800 Subject: [PATCH] feat(dns): add Lua scripting for DNS queries --- app/dns/config.pb.go | 31 +++++-- app/dns/config.proto | 4 + app/dns/dns.go | 16 ++++ app/dns/lua.go | 180 ++++++++++++++++++++++++++++++++++++ app/dns/lua_test.go | 54 +++++++++++ app/dns/nameserver.go | 3 +- app/dns/nameserver_local.go | 2 +- app/dns/script.go | 71 ++++++++++++++ common/geodata/lua.go | 60 ++++++++++++ common/geodata/lua_test.go | 42 +++++++++ common/lua/pool.go | 138 +++++++++++++++++++++++++++ common/lua/program.go | 56 +++++++++++ go.mod | 2 + go.sum | 9 ++ infra/conf/dns.go | 21 +++++ 15 files changed, 681 insertions(+), 8 deletions(-) create mode 100644 app/dns/lua.go create mode 100644 app/dns/lua_test.go create mode 100644 app/dns/script.go create mode 100644 common/geodata/lua.go create mode 100644 common/geodata/lua_test.go create mode 100644 common/lua/pool.go create mode 100644 common/lua/program.go diff --git a/app/dns/config.pb.go b/app/dns/config.pb.go index c0737a0d2..053721239 100644 --- a/app/dns/config.pb.go +++ b/app/dns/config.pb.go @@ -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" + diff --git a/app/dns/config.proto b/app/dns/config.proto index ddc19dc75..ca85582fe 100644 --- a/app/dns/config.proto +++ b/app/dns/config.proto @@ -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; } diff --git a/app/dns/dns.go b/app/dns/dns.go index b750fce26..6ea28ecf2 100644 --- a/app/dns/dns.go +++ b/app/dns/dns.go @@ -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 { diff --git a/app/dns/lua.go b/app/dns/lua.go new file mode 100644 index 000000000..1a429319b --- /dev/null +++ b/app/dns/lua.go @@ -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 +} diff --git a/app/dns/lua_test.go b/app/dns/lua_test.go new file mode 100644 index 000000000..610c6dcc7 --- /dev/null +++ b/app/dns/lua_test.go @@ -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) + } +} diff --git a/app/dns/nameserver.go b/app/dns/nameserver.go index 06882611f..27ce9a540 100644 --- a/app/dns/nameserver.go +++ b/app/dns/nameserver.go @@ -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) diff --git a/app/dns/nameserver_local.go b/app/dns/nameserver_local.go index 4369f89ed..54ab0311e 100644 --- a/app/dns/nameserver_local.go +++ b/app/dns/nameserver_local.go @@ -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} } diff --git a/app/dns/script.go b/app/dns/script.go new file mode 100644 index 000000000..e1b7ce3ac --- /dev/null +++ b/app/dns/script.go @@ -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 +} diff --git a/common/geodata/lua.go b/common/geodata/lua.go new file mode 100644 index 000000000..0531cde55 --- /dev/null +++ b/common/geodata/lua.go @@ -0,0 +1,60 @@ +package geodata + +import ( + lua "github.com/yuin/gopher-lua" + luar "layeh.com/gopher-luar" +) + +// RegisterLua makes xray.geodata available to require in an LState. +// Matchers retain registry handles, so they remain usable after a reload. +func RegisterLua(L *lua.LState) { + L.PreloadModule("xray.geodata", func(L *lua.LState) int { + module := L.NewTable() + + module.RawSetString("domainMatcher", L.NewFunction(func(L *lua.LState) int { + parsed, err := ParseDomainRules(luaRules(L, 1), Domain_Domain) + if err != nil { + L.RaiseError("%v", err) + return 0 + } + matcher, err := DomainReg.BuildDomainMatcher(parsed) + if err != nil { + L.RaiseError("%v", err) + return 0 + } + L.Push(luar.New(L, matcher)) + return 1 + })) + + module.RawSetString("ipMatcher", L.NewFunction(func(L *lua.LState) int { + parsed, err := ParseIPRules(luaRules(L, 1)) + if err != nil { + L.RaiseError("%v", err) + return 0 + } + matcher, err := IPReg.BuildIPMatcher(parsed) + if err != nil { + L.RaiseError("%v", err) + return 0 + } + L.Push(luar.New(L, matcher)) + return 1 + })) + L.Push(module) + return 1 + }) +} + +func luaRules(L *lua.LState, index int) []string { + table := L.CheckTable(index) + rules := make([]string, table.Len()) + for i := range rules { + value, ok := table.RawGetInt(i + 1).(lua.LString) + if !ok { + L.RaiseError("geodata rules must be strings") + return nil + } + rules[i] = string(value) + } + return rules +} diff --git a/common/geodata/lua_test.go b/common/geodata/lua_test.go new file mode 100644 index 000000000..400a6ed69 --- /dev/null +++ b/common/geodata/lua_test.go @@ -0,0 +1,42 @@ +package geodata + +import ( + "testing" + + "github.com/xtls/xray-core/common/net" + lua "github.com/yuin/gopher-lua" +) + +func TestLuaIPMatcherAcceptsNativeIP(t *testing.T) { + L := lua.NewState() + defer L.Close() + RegisterLua(L) + ip := L.NewUserData() + ip.Value = net.ParseIP("127.0.0.1") + L.SetGlobal("ip", ip) + if err := L.DoString(` + local matcher = require("xray.geodata").ipMatcher({"127.0.0.0/8"}) + assert(matcher:Match(ip)) + assert(matcher:AnyMatch({ip})) + assert(matcher:Matches({ip})) + local matched, unmatched = matcher:FilterIPs({ip}) + assert(#matched == 1 and #unmatched == 0) + assert(matcher:Match(matched[1])) + `); err != nil { + t.Fatal(err) + } +} + +func TestLuaDomainMatcherUsesNativeMatcher(t *testing.T) { + L := lua.NewState() + defer L.Close() + RegisterLua(L) + if err := L.DoString(` + local matcher = require("xray.geodata").domainMatcher({"example.com"}) + assert(matcher:MatchAny("example.com")) + assert(matcher:MatchAny("www.example.com")) + assert(#(matcher:Match("www.example.com")) == 1) + `); err != nil { + t.Fatal(err) + } +} diff --git a/common/lua/pool.go b/common/lua/pool.go new file mode 100644 index 000000000..dff82f733 --- /dev/null +++ b/common/lua/pool.go @@ -0,0 +1,138 @@ +package lua + +import ( + "context" + "errors" + "sync" + + glua "github.com/yuin/gopher-lua" +) + +const maxIdleStates = 16 + +// LStateFactory must initialize a state fully and observe ctx while doing so. +// The pool owns any non-nil state it returns, even when it also returns an error. +type LStateFactory func(ctx context.Context) (*glua.LState, error) + +// Pool lends each state to one caller at a time. It grows on contention and +// keeps up to maxIdleStates idle states until Close. Callers decide whether a +// state is reusable. +type Pool struct { + ctx context.Context + cancel context.CancelFunc + + factory LStateFactory + idle []*glua.LState + + mu sync.Mutex + active sync.WaitGroup + closed bool +} + +// NewPool initializes one state before returning, so top-level errors surface at startup. +func NewPool(ctx context.Context, factory LStateFactory) (*Pool, error) { + poolCtx, cancel := context.WithCancel(ctx) + + // Create one state now to catch factory errors at startup. + state, err := factory(poolCtx) + if err != nil { + cancel() + if state != nil { + state.Close() + } + return nil, err + } + if state == nil { + cancel() + return nil, errors.New("Lua state factory returned nil") + } + if err := poolCtx.Err(); err != nil { + state.Close() + cancel() + return nil, err + } + + return &Pool{ctx: poolCtx, cancel: cancel, factory: factory, idle: []*glua.LState{state}}, nil +} + +// Context is cancelled by Close. Query contexts should derive from it. +func (p *Pool) Context() context.Context { + return p.ctx +} + +// Acquire returns an initialized exclusive state, growing the pool if necessary. +func (p *Pool) Acquire() (*glua.LState, error) { + p.mu.Lock() + if p.closed || p.ctx.Err() != nil { + p.mu.Unlock() + return nil, p.ctx.Err() + } + + p.active.Add(1) + + n := len(p.idle) + if n != 0 { + state := p.idle[n-1] + p.idle = p.idle[:n-1] + p.mu.Unlock() + return state, nil + } + p.mu.Unlock() + + // TODO: Limit the total number of states. When the limit is reached, wait + // for a Release instead of creating another state; allow the wait to be + // cancelled by the caller or by Close. + state, err := p.factory(p.ctx) + if err == nil && state == nil { + err = errors.New("Lua state factory returned nil") + } + if err != nil { + if state != nil { + state.Close() + } + p.active.Done() + return nil, err + } + if err := p.ctx.Err(); err != nil { + state.Close() + p.active.Done() + return nil, err + } + + return state, nil +} + +// Release returns a healthy state to the pool and closes a failed or cancelled one. +func (p *Pool) Release(state *glua.LState, reusable bool) { + if reusable { + p.mu.Lock() + if !p.closed && p.ctx.Err() == nil && len(p.idle) < maxIdleStates { + p.idle = append(p.idle, state) + } else { + reusable = false + } + p.mu.Unlock() + } + + if !reusable { + state.Close() + } + + p.active.Done() +} + +// Close cancels active work, closes idle states, and waits for borrowed states. +func (p *Pool) Close() { + p.mu.Lock() + if !p.closed { + p.closed = true + p.cancel() + for _, state := range p.idle { + state.Close() + } + p.idle = nil + } + p.mu.Unlock() + + p.active.Wait() +} diff --git a/common/lua/program.go b/common/lua/program.go new file mode 100644 index 000000000..95bf10cc1 --- /dev/null +++ b/common/lua/program.go @@ -0,0 +1,56 @@ +// Package lua provides shared GopherLua programs and state management for Xray scripts. +package lua + +import ( + "bufio" + "context" + "os" + + glua "github.com/yuin/gopher-lua" + "github.com/yuin/gopher-lua/parse" +) + +// Program holds immutable bytecode that can be run by independent LStates. +type Program struct { + proto *glua.FunctionProto +} + +// CompileFile reads and compiles a Lua file once. +func CompileFile(path string) (*Program, error) { + f, err := os.Open(path) + if err != nil { + return nil, err + } + defer f.Close() + chunk, err := parse.Parse(bufio.NewReader(f), path) + if err != nil { + return nil, err + } + proto, err := glua.Compile(chunk, path) + if err != nil { + return nil, err + } + return &Program{proto: proto}, nil +} + +// NewState creates a VM, makes modules available, and executes the file top level. +// Module loaders run only when Lua calls require. Each state gets its own globals. +// The caller owns the returned state. +func (p *Program) NewState(ctx context.Context, register func(*glua.LState)) (*glua.LState, error) { + L := glua.NewState() + if register != nil { + register(L) + } + L.SetContext(ctx) + L.Push(L.NewFunctionFromProto(p.proto)) + err := L.PCall(0, 0, nil) + L.RemoveContext() + if err == nil { + err = ctx.Err() + } + if err != nil { + L.Close() + return nil, err + } + return L, nil +} diff --git a/go.mod b/go.mod index 45f565bea..6b2e5a405 100644 --- a/go.mod +++ b/go.mod @@ -23,6 +23,7 @@ require ( github.com/stretchr/testify v1.12.1 github.com/vishvananda/netlink v1.3.1 github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0 + github.com/yuin/gopher-lua v1.1.2 go4.org/netipx v0.0.0-20231129151722-fdeea329fbba golang.org/x/crypto v0.57.0 golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 @@ -36,6 +37,7 @@ require ( google.golang.org/protobuf v1.36.12 gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0 h12.io/socks v1.0.3 + layeh.com/gopher-luar v1.0.11 lukechampine.com/blake3 v1.4.1 mvdan.cc/gofumpt v0.12.0 ) diff --git a/go.sum b/go.sum index da03cef72..79eff3694 100644 --- a/go.sum +++ b/go.sum @@ -2,6 +2,9 @@ github.com/andybalholm/brotli v1.0.6 h1:Yf9fFpf49Zrxb9NlQaluyE92/+X7UVHlhMNJN2sx github.com/andybalholm/brotli v1.0.6/go.mod h1:fO7iG3H7G2nSZ7m0zPUDn85XEX2GTukHGRSepvi9Eig= github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e h1:5mgtR5gwIgBKMiGI1QdXldZZ+SNor06Nbu1wCBulQBg= github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e/go.mod h1:x7qxEvX6MCVtDuBKHj3E+88+BtrbEMuAL5qGUKItjW8= +github.com/chzyer/logex v1.1.10/go.mod h1:+Ywpsq7O8HXn0nuIou7OrIPyXbp3wmkHB+jjWRnGsAI= +github.com/chzyer/readline v0.0.0-20180603132655-2972be24d48e/go.mod h1:nSuG5e5PlCu98SY8svDHJxuZscDgtXS6KTTbou5AhLI= +github.com/chzyer/test v0.0.0-20180213035817-a1ea475d72b1/go.mod h1:Q3SI9o4m/ZMnBNeIyt5eFwwo7qiLfzFZmjNmxjkiQlU= github.com/cloudflare/circl v1.6.5 h1:O64F26HEqNhznd/hrC5KZXVKYuKM2rx4deZDTc4ihQA= github.com/cloudflare/circl v1.6.5/go.mod h1:h5LNyxAc5nTue9DS5jT+48en2PSDYt3zdGnz5OstK6c= github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344 h1:Arcl6UOIS/kgO2nW3A65HN+7CMjSDP/gofXL4CZt1V4= @@ -85,6 +88,9 @@ github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguH github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0 h1:rb+fKQFhz+5I2PPuQsNYxI5mUU840XWYtRF0ZBjvkws= github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0/go.mod h1:DsJblcWDGt76+FVqBVwbwRhxyyNJsGV48gJLch0OOWI= github.com/yuin/goldmark v1.4.1/go.mod h1:mwnBkeHKe2W/ZEtQ+71ViKU8L12m81fl3OWwC1Zlc8k= +github.com/yuin/gopher-lua v0.0.0-20190206043414-8bfc7677f583/go.mod h1:gqRgreBUhTSL0GeU64rtZ3Uq3wtjOa/TB2YfrtkCbVQ= +github.com/yuin/gopher-lua v1.1.2 h1:yF/FjE3hD65tBbt0VXLE13HWS9h34fdzJmrWRXwobGA= +github.com/yuin/gopher-lua v1.1.2/go.mod h1:7aRmXIWl37SqRf0koeyylBEzJ+aPt8A+mmkQ4f1ntR8= go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko= go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o= go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw= @@ -109,6 +115,7 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk= golang.org/x/sync v0.23.0/go.mod h1:sUUOizhqBxiL6pEWpqNLUiaJn1ShEbZ6BBqskPbjZm0= +golang.org/x/sys v0.0.0-20190204203706-41f3e6584952/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= @@ -159,6 +166,8 @@ gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0 h1:Lk6hARj5UPY47dBep70OD/TI gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0/go.mod h1:QkHjoMIBaYtpVufgwv3keYAbln78mBoCuShZrPrer1Q= h12.io/socks v1.0.3 h1:Ka3qaQewws4j4/eDQnOdpr4wXsC//dXtWvftlIcCQUo= h12.io/socks v1.0.3/go.mod h1:AIhxy1jOId/XCz9BO+EIgNL2rQiPTBNnOfnVnQ+3Eck= +layeh.com/gopher-luar v1.0.11 h1:8zJudpKI6HWkoh9eyyNFaTM79PY6CAPcIr6X/KTiliw= +layeh.com/gopher-luar v1.0.11/go.mod h1:TPnIVCZ2RJBndm7ohXyaqfhzjlZ+OA2SZR/YwL8tECk= lukechampine.com/blake3 v1.4.1 h1:I3Smz7gso8w4/TunLKec6K2fn+kyKtDxr/xcQEN84Wg= lukechampine.com/blake3 v1.4.1/go.mod h1:QFosUxmjB8mnrWFSNwKmvxHpfY72bmD2tQ0kBMM3kwo= mvdan.cc/gofumpt v0.12.0 h1:1Lbudkz2kpM9Cjz2pL4M19u7q+GaEhCTNf7N9mfpcho= diff --git a/infra/conf/dns.go b/infra/conf/dns.go index d55dada6a..a05814227 100644 --- a/infra/conf/dns.go +++ b/infra/conf/dns.go @@ -14,9 +14,11 @@ import ( "github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/geodata" "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/common/platform" ) type NameServerConfig struct { + ID string `json:"id"` Address *Address `json:"address"` ClientIP *Address `json:"clientIp"` Port uint16 `json:"port"` @@ -43,6 +45,7 @@ func (c *NameServerConfig) UnmarshalJSON(data []byte) error { } var advanced struct { + ID string `json:"id"` Address *Address `json:"address"` ClientIP *Address `json:"clientIp"` Port uint16 `json:"port"` @@ -60,6 +63,7 @@ func (c *NameServerConfig) UnmarshalJSON(data []byte) error { UnexpectedIPs StringList `json:"unexpectedIPs"` } if err := json.Unmarshal(data, &advanced); err == nil { + c.ID = advanced.ID c.Address = advanced.Address c.ClientIP = advanced.ClientIP c.Port = advanced.Port @@ -134,6 +138,7 @@ func (c *NameServerConfig) Build() (*dns.NameServer, error) { } return &dns.NameServer{ + Id: c.ID, Address: &net.Endpoint{ Network: net.Network_UDP, Address: c.Address.Build(), @@ -159,6 +164,7 @@ func (c *NameServerConfig) Build() (*dns.NameServer, error) { // DNSConfig is a JSON serializable object for dns.Config type DNSConfig struct { Servers []*NameServerConfig `json:"servers"` + Script string `json:"script"` Hosts *HostsWrapper `json:"hosts"` ClientIP *Address `json:"clientIp"` Tag string `json:"tag"` @@ -278,6 +284,21 @@ func (c *DNSConfig) Build() (*dns.Config, error) { QueryStrategy: resolveQueryStrategy(c.QueryStrategy), } + if c.Script != "" { + path := c.Script + if !filepath.IsAbs(path) { + path = filepath.Join(platform.GetConfDirPath(), path) + } + info, err := os.Stat(path) + if err != nil { + return nil, errors.New("DNS script does not exist: ", path).Base(err) + } + if !info.Mode().IsRegular() { + return nil, errors.New("DNS script is not a regular file: ", path) + } + config.Script = path + } + if c.ClientIP != nil { if !c.ClientIP.Family().IsIP() { return nil, errors.New("not an IP address:", c.ClientIP.String())