From 8855145e58ab066f4bd767f321cb80657000dfda Mon Sep 17 00:00:00 2001 From: Meow <197331664+Meo597@users.noreply.github.com> Date: Sat, 10 Oct 2026 17:01:16 +0800 Subject: [PATCH] Xray-core: Add Lua `script` for `dns` and `routing` (#6823) https://github.com/XTLS/Xray-core/pull/6823#issuecomment-5843754759 https://github.com/XTLS/Xray-core/pull/6823#issuecomment-5861456450 https://github.com/XTLS/Xray-core/pull/6823#issuecomment-6093596069 --- app/dispatcher/default.go | 3 + app/dns/config.pb.go | 31 +- app/dns/config.proto | 4 + app/dns/dns.go | 16 ++ app/dns/lua.go | 169 +++++++++++ app/dns/lua_benchmark_test.go | 118 ++++++++ app/dns/lua_test.go | 272 ++++++++++++++++++ app/dns/nameserver.go | 3 +- app/dns/nameserver_local.go | 2 +- app/dns/script.go | 63 +++++ app/dns/script_test.go | 280 +++++++++++++++++++ app/router/config.pb.go | 18 +- app/router/config.proto | 2 + app/router/lua.go | 175 ++++++++++++ app/router/lua_benchmark_test.go | 223 +++++++++++++++ app/router/lua_test.go | 274 ++++++++++++++++++ app/router/router.go | 17 ++ app/router/script.go | 76 +++++ app/router/script_test.go | 382 +++++++++++++++++++++++++ common/geodata/lua.go | 163 +++++++++++ common/geodata/lua_test.go | 172 ++++++++++++ common/log/lua.go | 61 ++++ common/log/lua_test.go | 213 ++++++++++++++ common/lua/lua.go | 3 + common/lua/luar.go | 65 +++++ common/lua/luar_test.go | 84 ++++++ common/lua/pool.go | 151 ++++++++++ common/lua/pool_test.go | 466 +++++++++++++++++++++++++++++++ common/lua/program.go | 77 +++++ common/lua/program_test.go | 76 +++++ common/lua/utils.go | 95 +++++++ common/lua/utils_test.go | 121 ++++++++ common/platform/platform.go | 48 ++++ common/platform/platform_test.go | 51 ++++ go.mod | 2 + go.sum | 9 + infra/conf/dns.go | 14 + infra/conf/dns_test.go | 50 ++++ infra/conf/router.go | 11 + infra/conf/router_test.go | 38 +++ 40 files changed, 4086 insertions(+), 12 deletions(-) create mode 100644 app/dns/lua.go create mode 100644 app/dns/lua_benchmark_test.go create mode 100644 app/dns/lua_test.go create mode 100644 app/dns/script.go create mode 100644 app/dns/script_test.go create mode 100644 app/router/lua.go create mode 100644 app/router/lua_benchmark_test.go create mode 100644 app/router/lua_test.go create mode 100644 app/router/script.go create mode 100644 app/router/script_test.go create mode 100644 common/geodata/lua.go create mode 100644 common/geodata/lua_test.go create mode 100644 common/log/lua.go create mode 100644 common/log/lua_test.go create mode 100644 common/lua/lua.go create mode 100644 common/lua/luar.go create mode 100644 common/lua/luar_test.go create mode 100644 common/lua/pool.go create mode 100644 common/lua/pool_test.go create mode 100644 common/lua/program.go create mode 100644 common/lua/program_test.go create mode 100644 common/lua/utils.go create mode 100644 common/lua/utils_test.go diff --git a/app/dispatcher/default.go b/app/dispatcher/default.go index 2717df5fe..cae973517 100644 --- a/app/dispatcher/default.go +++ b/app/dispatcher/default.go @@ -470,6 +470,9 @@ func (d *DefaultDispatcher) routedDispatch(ctx context.Context, link *transport. return // DO NOT CHANGE: the traffic shouldn't be processed by default outbound if the specified outbound tag doesn't exist (yet), e.g., VLESS Reverse Proxy } } else { + if err != common.ErrNoClue { + errors.LogErrorInner(ctx, err, "failed to pick route for ", destination) + } errors.LogInfo(ctx, "default route for ", destination) } } 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 6fb88d695..9943da09d 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 } @@ -279,6 +292,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..3c5194baf --- /dev/null +++ b/app/dns/lua.go @@ -0,0 +1,169 @@ +package dns + +import ( + "context" + "strings" + + "github.com/xtls/xray-core/common/errors" + xlua "github.com/xtls/xray-core/common/lua" + "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" +) + +// 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) { + 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) + }, + }} + } + registerLua(L, servers, client) +} + +// registerLua makes xray.dns available to DNS scripts. +func (s *DNS) registerLua(L *lua.LState) { + registerLua(L, luaServers(s), nil) +} + +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 { + pushIPs := xlua.NewSlicePusher[net.IP](L) + + serverList := L.CreateTable(len(servers), 0) + for i, client := range servers { + server := L.CreateTable(0, 2) + + server.RawSetString("ID", lua.LString(client.id)) + + server.RawSetString("Query", L.NewFunction(func(L *lua.LState) int { + domain, ok := L.Get(2).(lua.LString) + if !ok { + L.RaiseError("server:Query requires a domain") + return 0 + } + option := featureDNS.IPOption{ + IPv4Enable: L.CheckBool(3), + IPv6Enable: L.CheckBool(4), + FakeEnable: L.CheckBool(5), + } + 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.query(ctx, string(domain), option) + } + pushIPs(L, ips) + xlua.PushNumber(L, ttl) + xlua.PushError(L, err) + return 3 + })) + serverList.RawSetInt(i+1, server) + } + + module := L.CreateTable(0, 2) + if servers != nil { + module.RawSetString("Servers", serverList) + } + if client != nil { + module.RawSetString("Query", newLuaClientQuery(L, client, pushIPs)) + } + L.Push(module) + return 1 + }) +} + +func newLuaClientQuery(L *lua.LState, client featureDNS.Client, pushIPs func(*lua.LState, []net.IP)) *lua.LFunction { + return L.NewFunction(func(L *lua.LState) int { + domain, ok := L.Get(1).(lua.LString) + if !ok { + L.RaiseError("dns.Query requires a domain") + return 0 + } + option := featureDNS.IPOption{ + IPv4Enable: L.CheckBool(2), + IPv6Enable: L.CheckBool(3), + FakeEnable: L.CheckBool(4), + } + if L.Context() == nil { + L.RaiseError("dns.Query requires an active DNS query") + return 0 + } + ips, ttl, err := client.LookupIP(string(domain), option) + pushIPs(L, ips) + xlua.PushNumber(L, ttl) + xlua.PushError(L, err) + return 3 + }) +} + +// callLuaQuery runs HandleDNSQuery and leaves (ips, ttl, err) on the stack. +func callLuaQuery(L *lua.LState, domain string, option featureDNS.IPOption) error { + fn := L.GetGlobal("HandleDNSQuery") + if fn.Type() != lua.LTFunction { + return errors.New("DNS script must define HandleDNSQuery(...)") + } + + return 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)) +} + +// readLuaQueryResult reads (ips, ttl, err) from the stack without copying the IPs. +func readLuaQueryResult(L *lua.LState) ([]net.IP, uint32, error) { + if err := xlua.ReadError(L.Get(-1), "DNS script error must be an error or string"); err != nil { + return nil, 0, err + } + + ttl, err := xlua.ReadUint32(L.Get(-2), "DNS script returned invalid TTL") + if err != nil { + return nil, 0, err + } + + addresses := L.Get(-3) + if addresses == lua.LNil { + return nil, 0, featureDNS.ErrEmptyResponse + } + ips, err := xlua.ReadUserData[[]net.IP](addresses, "DNS script IPs must be native IP slice userdata") + if err != nil { + return nil, 0, err + } + if len(ips) == 0 { + return nil, 0, featureDNS.ErrEmptyResponse + } + + return ips, ttl, nil +} diff --git a/app/dns/lua_benchmark_test.go b/app/dns/lua_benchmark_test.go new file mode 100644 index 000000000..ac4f6b8c8 --- /dev/null +++ b/app/dns/lua_benchmark_test.go @@ -0,0 +1,118 @@ +package dns + +import ( + "context" + "os" + "path/filepath" + "testing" + "time" + + "github.com/xtls/xray-core/common/net" + featureDNS "github.com/xtls/xray-core/features/dns" + lua "github.com/yuin/gopher-lua" +) + +// BenchmarkLuaDNSHook isolates scalar argument bridging and a fixed return. +// It excludes upstream queries, result decoding, and state pool management. +func BenchmarkLuaDNSHook(b *testing.B) { + L := lua.NewState() + b.Cleanup(L.Close) + if err := L.DoString(` +function HandleDNSQuery(domain, ipv4, ipv6, fake) + return true +end +`); err != nil { + b.Fatal(err) + } + L.SetContext(context.Background()) + option := featureDNS.IPOption{IPv4Enable: true} + if err := callLuaQuery(L, "example.com", option); err != nil { + b.Fatal(err) + } + if L.Get(-3) != lua.LTrue { + b.Fatal("hook did not return true") + } + L.Pop(3) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + if err := callLuaQuery(L, "example.com", option); err != nil { + b.Fatal(err) + } + L.Pop(3) + } +} + +// BenchmarkLuaDNSQuery queries the same preselected, in-memory upstream. +// client_query compares Client.QueryIP to a preloaded server:Query hook. +// script_query additionally measures production pool and timeout management. +// These cases do not measure DNS.LookupIP server selection or network latency. +func BenchmarkLuaDNSQuery(b *testing.B) { + ctx := context.Background() + option := featureDNS.IPOption{IPv4Enable: true} + ip := net.ParseIP("127.0.0.1") + upstream := &benchmarkLuaNameServer{ips: []net.IP{ip}} + client := &Client{server: upstream, ipOption: &option, timeoutMs: time.Second} + server := &DNS{ctx: ctx, clients: []*Client{client}} + const script = ` +local server = require("xray.dns").Servers[1] +function HandleDNSQuery(domain, ipv4, ipv6, fake) + return server:Query(domain, ipv4, ipv6, fake) +end +` + L := lua.NewState() + b.Cleanup(L.Close) + server.registerLua(L) + if err := L.DoString(script); err != nil { + b.Fatal(err) + } + L.SetContext(ctx) + + path := filepath.Join(b.TempDir(), "query.lua") + if err := os.WriteFile(path, []byte(script), 0o600); err != nil { + b.Fatal(err) + } + engine, err := newScriptEngine(path, server) + if err != nil { + b.Fatal(err) + } + b.Cleanup(engine.close) + for _, bench := range []struct { + name string + query func() ([]net.IP, uint32, error) + }{ + {"client_query/native", func() ([]net.IP, uint32, error) { + return client.QueryIP(ctx, "example.com", option) + }}, + {"client_query/lua", func() ([]net.IP, uint32, error) { + if err := callLuaQuery(L, "example.com", option); err != nil { + return nil, 0, err + } + ips, ttl, err := readLuaQueryResult(L) + L.Pop(3) + return ips, ttl, err + }}, + {"script_query/lua", func() ([]net.IP, uint32, error) { + return engine.query("example.com", option) + }}, + } { + b.Run(bench.name, func(b *testing.B) { + ips, ttl, err := bench.query() + if err != nil || ttl != 60 || len(ips) != 1 || !ips[0].Equal(ip) { + b.Fatalf("query() = %v, TTL %d, %v; want %v, TTL 60", ips, ttl, err, ip) + } + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + ips, ttl, err = bench.query() + if err != nil { + b.Fatal(err) + } + } + b.StopTimer() + if ttl != 60 || len(ips) != 1 || !ips[0].Equal(ip) { + b.Fatalf("query() = %v, TTL %d; want %v, TTL 60", ips, ttl, ip) + } + }) + } +} diff --git a/app/dns/lua_test.go b/app/dns/lua_test.go new file mode 100644 index 000000000..d8c647326 --- /dev/null +++ b/app/dns/lua_test.go @@ -0,0 +1,272 @@ +package dns + +import ( + "context" + go_errors "errors" + "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" + "github.com/xtls/xray-core/features/dns/localdns" + lua "github.com/yuin/gopher-lua" +) + +func TestReadLuaQueryResult(t *testing.T) { + wantIPs := []net.IP{net.ParseIP("8.8.8.8"), {127, 0, 0, 1}, net.ParseIP("::1")} + nativeErr := go_errors.New("upstream failed") + for _, tc := range []struct { + name, values string + wantIPs []net.IP + wantTTL uint32 + wantErr error + wantMessage string + }{ + {name: "IPs", values: `ips, 45`, wantIPs: wantIPs, wantTTL: 45}, + {name: "nil IPs", values: `nil, 0`, wantErr: featureDNS.ErrEmptyResponse}, + {name: "empty IPs", values: `emptyIPs, 0`, wantErr: featureDNS.ErrEmptyResponse}, + {name: "native error", values: `nil, nil, nativeError`, wantErr: nativeErr}, + {name: "string error", values: `nil, nil, "blocked"`, wantMessage: "blocked"}, + {name: "fractional TTL", values: `ips, 1.5`, wantMessage: "invalid TTL"}, + {name: "oversized TTL", values: `ips, 4294967296`, wantMessage: "invalid TTL"}, + {name: "negative TTL", values: `ips, -1`, wantMessage: "invalid TTL"}, + {name: "NaN TTL", values: `ips, 0/0`, wantMessage: "invalid TTL"}, + {name: "missing TTL", values: `ips`, wantMessage: "invalid TTL"}, + {name: "string IPs", values: `"127.0.0.1", 60`, wantMessage: "native IP slice"}, + {name: "wrong userdata", values: `ip, 60`, wantMessage: "native IP slice"}, + {name: "invalid error", values: `ips, 60, false`, wantMessage: "error or string"}, + } { + t.Run(tc.name, func(t *testing.T) { + L := lua.NewState() + defer L.Close() + for name, value := range map[string]any{"ips": wantIPs, "ip": wantIPs[0], "emptyIPs": []net.IP(nil), "nativeError": nativeErr} { + ud := L.NewUserData() + ud.Value = value + L.SetGlobal(name, ud) + } + fn, err := L.LoadString("return " + tc.values) + if err != nil { + t.Fatal(err) + } + if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true}); err != nil { + t.Fatal(err) + } + ips, ttl, err := readLuaQueryResult(L) + switch { + case tc.wantErr != nil: + if err != tc.wantErr { + t.Fatalf("error = %v, want original error %v", err, tc.wantErr) + } + case tc.wantMessage != "": + if err == nil || !strings.Contains(err.Error(), tc.wantMessage) { + t.Fatalf("error = %v, want %q", err, tc.wantMessage) + } + case err != nil: + t.Fatal(err) + } + if ttl != tc.wantTTL || len(ips) != len(tc.wantIPs) { + t.Fatalf("result = %v, TTL %d; want %v, TTL %d", ips, ttl, tc.wantIPs, tc.wantTTL) + } + for i := range ips { + if !ips[i].Equal(tc.wantIPs[i]) { + t.Fatalf("IP %d = %v, want %v", i, ips[i], tc.wantIPs[i]) + } + } + if len(ips) != 0 && &ips[0] != &tc.wantIPs[0] { + t.Fatal("result copied the IP slice") + } + }) + } +} + +func TestCallLuaQueryCancellation(t *testing.T) { + L := lua.NewState() + defer L.Close() + if err := L.DoString(`function HandleDNSQuery(domain, ipv4, ipv6, fake) while true do end end`); err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + cancel() + L.SetContext(ctx) + err := callLuaQuery(L, "example.com", featureDNS.IPOption{IPv4Enable: true}) + if err == nil { + t.Fatal("callLuaQuery did not stop after context cancellation") + } + if L.Context() != ctx { + t.Fatal("callLuaQuery changed the Lua state's context") + } +} + +func TestCallLuaQuery(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(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) + } + if err := callLuaQuery(L, "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil { + t.Fatal(err) + } + if L.GetTop() != 3 || L.Get(1) != addresses || L.Get(2) != lua.LNumber(60) || L.Get(3) != lua.LNil { + t.Fatal("callLuaQuery did not leave the three query results on the 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").BuildIPMatcher("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(#ips == 2 and ips[1]:String() == "127.0.0.1" and ips[2]:String() == "8.8.8.8") + assert(matcher:Match(ips[1]) and not matcher:Match(ips[2])) + assert(matcher:AnyMatch(ips)) + local matched, unmatched = matcher:FilterIPs(ips) + assert(#matched == 1 and #unmatched == 1) + assert(matched[1]:Equal(ips[1]) and unmatched[1]:Equal(ips[2])) + return matched, ttl, err +end +`); err != nil { + t.Fatal(err) + } + L.SetContext(context.Background()) + if err := callLuaQuery(L, "example.com", option); err != nil { + t.Fatal(err) + } + got, ttl, err := readLuaQueryResult(L) + 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 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()) + geodata.RegisterLua(L) + want := []net.IP{{127, 0, 0, 1}, net.ParseIP("::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)) +assert(#ips == 2 and ips[1]:String() == "127.0.0.1" and ips[2]:String() == "::1") +assert(matcher:Match(ips[1]) and not matcher:Match(ips[2])) +`); 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()) + 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) +assert(#serverIPs == 1 and #clientIPs == 1) +assert(serverIPs[1]:String() == "127.0.0.1" and serverIPs[1]:Equal(clientIPs[1])) +`); 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) + } + } +} + +func TestLuaDNSQueryEmptyIPs(t *testing.T) { + for _, tc := range []struct { + name string + ips []net.IP + }{ + {"nil", nil}, + {"empty", []net.IP{}}, + } { + t.Run(tc.name, func(t *testing.T) { + L := lua.NewState() + defer L.Close() + L.SetContext(context.Background()) + L.SetGlobal("expectNil", lua.LBool(tc.ips == nil)) + client := &luaDNSClient{lookup: func(string, featureDNS.IPOption) ([]net.IP, uint32, error) { + return tc.ips, 0, featureDNS.ErrEmptyResponse + }} + registerLua(L, []luaDNSServer{{query: func(_ context.Context, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) { + return client.LookupIP(domain, option) + }}}, client) + if err := L.DoString(` +local dns = require("xray.dns") +for _, query in ipairs({ + function() return dns.Servers[1]:Query("empty.example", true, false, false) end, + function() return dns.Query("empty.example", true, false, false) end, +}) do + local ips, ttl, err = query() + assert(ttl == 0 and err) + if expectNil then + assert(ips == nil) + else + assert(type(ips) == "userdata" and #ips == 0) + assert(not pcall(function() return ips[1] end)) + end +end +`); err != nil { + t.Fatal(err) + } + }) + } +} + +type benchmarkLuaNameServer struct { + ips []net.IP +} + +func (*benchmarkLuaNameServer) Name() string { return "benchmark" } +func (*benchmarkLuaNameServer) IsDisableCache() bool { return true } +func (s *benchmarkLuaNameServer) QueryIP(context.Context, string, featureDNS.IPOption) ([]net.IP, uint32, error) { + return s.ips, 60, nil +} diff --git a/app/dns/nameserver.go b/app/dns/nameserver.go index 0882d1e7d..5112f3581 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..012d44f7a --- /dev/null +++ b/app/dns/script.go @@ -0,0 +1,63 @@ +package dns + +import ( + "time" + + "github.com/xtls/xray-core/common/errors" + "github.com/xtls/xray-core/common/geodata" + "github.com/xtls/xray-core/common/log" + xlua "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 = 6 * time.Second + +type scriptEngine struct { + pool *xlua.Pool +} + +func newScriptEngine(path string, server *DNS) (*scriptEngine, error) { + program, err := xlua.CompileFile(path) + if err != nil { + return nil, err + } + + pool, err := xlua.NewPool(server.ctx, scriptExecutionTimeout, program.NewStateFactory( + scriptExecutionTimeout*20, + func(L *lua.LState) { + geodata.RegisterLua(L) + log.RegisterLua(L) + server.registerLua(L) + }, + func(L *lua.LState) error { + if L.GetGlobal("HandleDNSQuery").Type() != lua.LTFunction { + return errors.New("DNS script must define HandleDNSQuery(...)") + } + return nil + })) + if err != nil { + return nil, err + } + + errors.LogInfo(server.ctx, "DNS script initialized from ", path) + return &scriptEngine{pool: pool}, nil +} + +func (e *scriptEngine) close() { + e.pool.Close() +} + +func (e *scriptEngine) query(domain string, option dns.IPOption) (ips []net.IP, ttl uint32, queryErr error) { + if err := e.pool.WithState(nil, 0, func(L *lua.LState) error { + if err := callLuaQuery(L, domain, option); err != nil { + return err + } + ips, ttl, queryErr = readLuaQueryResult(L) + return nil + }); err != nil { + return nil, 0, err + } + return ips, ttl, queryErr +} diff --git a/app/dns/script_test.go b/app/dns/script_test.go new file mode 100644 index 000000000..cea0bb1c7 --- /dev/null +++ b/app/dns/script_test.go @@ -0,0 +1,280 @@ +package dns + +import ( + "context" + go_errors "errors" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/xtls/xray-core/common/net" + featureDNS "github.com/xtls/xray-core/features/dns" +) + +type scriptNameServer struct { + name string + answers map[string]net.IP + errors map[string]error + ttl uint32 + calls int +} + +func (s *scriptNameServer) Name() string { return s.name } +func (s *scriptNameServer) IsDisableCache() bool { return true } + +func (s *scriptNameServer) QueryIP(ctx context.Context, domain string, _ featureDNS.IPOption) ([]net.IP, uint32, error) { + if err := ctx.Err(); err != nil { + return nil, 0, err + } + s.calls++ + if err := s.errors[domain]; err != nil { + return nil, 0, err + } + ip, ok := s.answers[domain] + if !ok { + return nil, 0, featureDNS.ErrEmptyResponse + } + return []net.IP{ip}, s.ttl, nil +} + +func TestDNSScriptQuery(t *testing.T) { + wantIP := net.ParseIP("127.0.0.1") + upstreamErr := go_errors.New("upstream failed") + for _, tc := range []struct { + name, body string + wantIPs []net.IP + wantTTL uint32 + wantErr error + wantMessage string + wantCalls uint32 + }{ + {name: "IPs", body: `return server:Query(domain, ipv4, ipv6, fake)`, wantIPs: []net.IP{wantIP}, wantTTL: 60, wantCalls: 2}, + {name: "empty result", body: `return nil, 0`, wantErr: featureDNS.ErrEmptyResponse, wantCalls: 2}, + {name: "upstream error", body: `return server:Query("failed.example", ipv4, ipv6, fake)`, wantErr: upstreamErr, wantCalls: 2}, + {name: "string error", body: `return nil, nil, "blocked"`, wantMessage: "blocked", wantCalls: 2}, + {name: "invalid result", body: `return false, 0`, wantMessage: "native IP slice", wantCalls: 2}, + {name: "execution error", body: `error("execution failed")`, wantMessage: "execution failed", wantCalls: 1}, + } { + t.Run(tc.name, func(t *testing.T) { + script := ` +local server = require("xray.dns").Servers[1] +local calls = 0 +function HandleDNSQuery(domain, ipv4, ipv6, fake) + calls = calls + 1 + if domain == "count.example" then + local ips, _, err = server:Query("good.example", ipv4, ipv6, fake) + return ips, calls, err + end + ` + tc.body + ` +end +` + path := filepath.Join(t.TempDir(), "query.lua") + if err := os.WriteFile(path, []byte(script), 0o600); err != nil { + t.Fatal(err) + } + option := featureDNS.IPOption{IPv4Enable: true} + upstream := &scriptNameServer{ + name: "test", + answers: map[string]net.IP{"good.example": wantIP}, + errors: map[string]error{"failed.example": upstreamErr}, + ttl: 60, + } + server := &DNS{ + ctx: context.Background(), + clients: []*Client{{server: upstream, ipOption: &option, timeoutMs: time.Second}}, + } + engine, err := newScriptEngine(path, server) + if err != nil { + t.Fatal(err) + } + defer engine.close() + + ips, ttl, err := engine.query("good.example", option) + switch { + case tc.wantErr != nil: + if err != tc.wantErr { + t.Fatalf("query error = %v, want original error %v", err, tc.wantErr) + } + case tc.wantMessage != "": + if err == nil || !strings.Contains(err.Error(), tc.wantMessage) { + t.Fatalf("query error = %v, want %q", err, tc.wantMessage) + } + case err != nil: + t.Fatal(err) + } + if ttl != tc.wantTTL || len(ips) != len(tc.wantIPs) { + t.Fatalf("query = %v, TTL %d; want %v, TTL %d", ips, ttl, tc.wantIPs, tc.wantTTL) + } + for i := range ips { + if !ips[i].Equal(tc.wantIPs[i]) { + t.Fatalf("IP %d = %v, want %v", i, ips[i], tc.wantIPs[i]) + } + } + + ips, calls, err := engine.query("count.example", option) + if err != nil || calls != tc.wantCalls || len(ips) != 1 || !ips[0].Equal(wantIP) { + t.Fatalf("next query = %v, calls %d, %v; want %v, calls %d", ips, calls, err, wantIP, tc.wantCalls) + } + }) + } +} + +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").BuildIPMatcher("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(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(domain, ipv4, ipv6, fake) +end +` + scriptPath := filepath.Join(t.TempDir(), "geoip_fallback.lua") + if err := os.WriteFile(scriptPath, []byte(script), 0o600); err != nil { + t.Fatal(err) + } + + primary := &scriptNameServer{ + 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 := &scriptNameServer{ + 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 TestDNSScriptFakeDNSOption(t *testing.T) { + path := filepath.Join(t.TempDir(), "script.lua") + script := ` +local server = require("xray.dns").Servers[1] +local log = require("xray.log") +log.Info("DNS script loaded") +function HandleDNSQuery(domain, ipv4, ipv6, fake) + log.Debug("DNS query: ", domain) + 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 { + t.Fatal(err) + } + option := featureDNS.IPOption{IPv4Enable: true} + hosts, err := NewStaticHosts(nil) + if err != nil { + t.Fatal(err) + } + upstream := &scriptNameServer{ + 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("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/app/router/config.pb.go b/app/router/config.pb.go index 636096a7f..2eea46d5f 100644 --- a/app/router/config.pb.go +++ b/app/router/config.pb.go @@ -587,8 +587,10 @@ type Config struct { DomainStrategy Config_DomainStrategy `protobuf:"varint,1,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.app.router.Config_DomainStrategy" json:"domain_strategy,omitempty"` Rule []*RoutingRule `protobuf:"bytes,2,rep,name=rule,proto3" json:"rule,omitempty"` BalancingRule []*BalancingRule `protobuf:"bytes,3,rep,name=balancing_rule,json=balancingRule,proto3" json:"balancing_rule,omitempty"` - unknownFields protoimpl.UnknownFields - sizeCache protoimpl.SizeCache + // Absolute path to the Lua routing script. + Script string `protobuf:"bytes,4,opt,name=script,proto3" json:"script,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache } func (x *Config) Reset() { @@ -642,6 +644,13 @@ func (x *Config) GetBalancingRule() []*BalancingRule { return nil } +func (x *Config) GetScript() string { + if x != nil { + return x.Script + } + return "" +} + var File_app_router_config_proto protoreflect.FileDescriptor const file_app_router_config_proto_rawDesc = "" + @@ -699,11 +708,12 @@ const file_app_router_config_proto_rawDesc = "" + "\tbaselines\x18\x03 \x03(\x03R\tbaselines\x12\x1a\n" + "\bexpected\x18\x04 \x01(\x05R\bexpected\x12\x16\n" + "\x06maxRTT\x18\x05 \x01(\x03R\x06maxRTT\x12\x1c\n" + - "\ttolerance\x18\x06 \x01(\x02R\ttolerance\"\x96\x02\n" + + "\ttolerance\x18\x06 \x01(\x02R\ttolerance\"\xae\x02\n" + "\x06Config\x12O\n" + "\x0fdomain_strategy\x18\x01 \x01(\x0e2&.xray.app.router.Config.DomainStrategyR\x0edomainStrategy\x120\n" + "\x04rule\x18\x02 \x03(\v2\x1c.xray.app.router.RoutingRuleR\x04rule\x12E\n" + - "\x0ebalancing_rule\x18\x03 \x03(\v2\x1e.xray.app.router.BalancingRuleR\rbalancingRule\"B\n" + + "\x0ebalancing_rule\x18\x03 \x03(\v2\x1e.xray.app.router.BalancingRuleR\rbalancingRule\x12\x16\n" + + "\x06script\x18\x04 \x01(\tR\x06script\"B\n" + "\x0eDomainStrategy\x12\b\n" + "\x04AsIs\x10\x00\x12\x10\n" + "\fIpIfNonMatch\x10\x02\x12\x0e\n" + diff --git a/app/router/config.proto b/app/router/config.proto index 67e4e47f1..09b64b0a3 100644 --- a/app/router/config.proto +++ b/app/router/config.proto @@ -110,4 +110,6 @@ message Config { DomainStrategy domain_strategy = 1; repeated RoutingRule rule = 2; repeated BalancingRule balancing_rule = 3; + // Absolute path to the Lua routing script. + string script = 4; } diff --git a/app/router/lua.go b/app/router/lua.go new file mode 100644 index 000000000..6cf55d5f1 --- /dev/null +++ b/app/router/lua.go @@ -0,0 +1,175 @@ +package router + +import ( + "runtime" + "strings" + + "github.com/xtls/xray-core/common/errors" + xlua "github.com/xtls/xray-core/common/lua" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/features/routing" + lua "github.com/yuin/gopher-lua" +) + +const ( + luaContextType = "xray.router.Context" + luaAttributesType = "xray.router.Attributes" +) + +// RegisterLua makes xray.router available to routing scripts. +func (r *Router) RegisterLua(L *lua.LState) { + registerLuaContext(L) + + L.PreloadModule("xray.router", func(L *lua.LState) int { + module := L.CreateTable(0, 7) + + module.RawSetString("NetworkUnknown", lua.LNumber(net.Network_Unknown)) + module.RawSetString("NetworkTCP", lua.LNumber(net.Network_TCP)) + module.RawSetString("NetworkUDP", lua.LNumber(net.Network_UDP)) + module.RawSetString("NetworkUNIX", lua.LNumber(net.Network_UNIX)) + module.RawSetString("LocalOS", lua.LString(runtime.GOOS)) + + module.RawSetString("PickOutbound", L.NewFunction(func(L *lua.LState) int { + tag, ok := L.Get(2).(lua.LString) + if !ok { + L.ArgError(2, "balancer tag must be a string") + return 0 + } + balancer, found := (*r.balancers.Load())[string(tag)] + if !found { + xlua.PushNil(L) + xlua.PushError(L, errors.New("balancer ", tag, " not found")) + return 2 + } + outboundTag, err := balancer.PickOutbound() + xlua.PushString(L, outboundTag) + xlua.PushError(L, err) + return 2 + })) + + module.RawSetString("FindProcess", L.NewFunction(func(L *lua.LState) int { + pid, name, path, err := findProcess(checkLuaContext(L), net.FindProcess) + xlua.PushNumber(L, pid) + xlua.PushString(L, name) + xlua.PushString(L, path) + xlua.PushError(L, err) + return 4 + })) + + L.Push(module) + return 1 + }) +} + +func registerLuaContext(L *lua.LState) { + pushIPs := xlua.NewSlicePusher[net.IP](L) + attributes := L.NewTypeMetatable(luaAttributesType) + L.SetField(attributes, "__index", L.NewFunction(func(L *lua.LState) int { + values := L.CheckUserData(1).Value.(map[string]string) + key := L.CheckString(2) + if value, found := values[key]; found { + xlua.PushString(L, value) + } else { + xlua.PushNil(L) + } + return 1 + })) + methods := L.CreateTable(0, 4) + L.SetFuncs(methods, map[string]lua.LGFunction{ + "GetSourceIPs": func(L *lua.LState) int { + pushIPs(L, checkLuaContext(L).GetSourceIPs()) + return 1 + }, + "GetTargetIPs": func(L *lua.LState) int { + pushIPs(L, checkLuaContext(L).GetTargetIPs()) + return 1 + }, + "GetLocalIPs": func(L *lua.LState) int { + pushIPs(L, checkLuaContext(L).GetLocalIPs()) + return 1 + }, + "GetAttributes": func(L *lua.LState) int { + values := L.NewUserData() + values.Value = checkLuaContext(L).GetAttributes() + L.SetMetatable(values, attributes) + L.Push(values) + return 1 + }, + }) + L.SetField(L.NewTypeMetatable(luaContextType), "__index", methods) +} + +func checkLuaContext(L *lua.LState) routing.Context { + ctx, ok := L.CheckUserData(1).Value.(routing.Context) + if !ok { + L.ArgError(1, "routing context expected") + } + return ctx +} + +// callLuaRoute runs HandleRoute and leaves (outboundTag, ruleTag, err) on the stack. +func callLuaRoute(L *lua.LState, ctx routing.Context) error { + fn := L.GetGlobal("HandleRoute") + if fn.Type() != lua.LTFunction { + return errors.New("routing script must define HandleRoute(...)") + } + + value := L.NewUserData() + value.Value = ctx + L.SetMetatable(value, L.GetTypeMetatable(luaContextType)) + + return L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true}, + value, + lua.LString(ctx.GetInboundTag()), + lua.LNumber(ctx.GetSourcePort()), + lua.LNumber(ctx.GetTargetPort()), + lua.LNumber(ctx.GetLocalPort()), + lua.LString(strings.ToLower(ctx.GetTargetDomain())), + lua.LNumber(ctx.GetNetwork()), + lua.LString(ctx.GetProtocol()), + lua.LString(ctx.GetUser()), + lua.LNumber(ctx.GetVlessRoute()), + lua.LBool(ctx.GetSkipDNSResolve())) +} + +// readLuaRouteResult reads (outboundTag, ruleTag, err) from the stack. +func readLuaRouteResult(L *lua.LState) (string, string, error) { + if err := xlua.ReadError(L.Get(-1), "routing script error must be an error or string"); err != nil { + return "", "", err + } + + outboundTag, err := xlua.ReadOptionalString(L.Get(-3), "routing script outboundTag must be a string or nil") + if err != nil || outboundTag == "" { + return "", "", err + } + + ruleTag, err := xlua.ReadOptionalString(L.Get(-2), "routing script ruleTag must be a string") + if err != nil { + return "", "", err + } + + return outboundTag, ruleTag, nil +} + +type processFinder func(string, string, uint16, string, uint16) (int, string, string, error) + +func findProcess(ctx routing.Context, finder processFinder) (int, string, string, error) { + sources := ctx.GetSourceIPs() + if len(sources) == 0 { + return 0, "", "", errors.New("process lookup requires a source IP") + } + var network string + switch ctx.GetNetwork() { + case net.Network_TCP: + network = "tcp" + case net.Network_UDP: + network = "udp" + default: + return 0, "", "", errors.New("process lookup requires TCP or UDP") + } + targetIP, targetPort := "", uint16(0) + if targets := ctx.GetTargetIPs(); len(targets) > 0 { + targetIP, targetPort = targets[0].String(), uint16(ctx.GetTargetPort()) + } + return finder(network, sources[0].String(), uint16(ctx.GetSourcePort()), targetIP, targetPort) +} diff --git a/app/router/lua_benchmark_test.go b/app/router/lua_benchmark_test.go new file mode 100644 index 000000000..a1424cee5 --- /dev/null +++ b/app/router/lua_benchmark_test.go @@ -0,0 +1,223 @@ +package router + +import ( + "context" + "fmt" + "os" + "path/filepath" + "strings" + "testing" + + "github.com/xtls/xray-core/common/geodata" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/common/session" + "github.com/xtls/xray-core/features/routing" + routing_session "github.com/xtls/xray-core/features/routing/session" + lua "github.com/yuin/gopher-lua" +) + +func benchmarkRouteContext(target net.Destination) *routing_session.Context { + // Use the production context: its IP getters construct a slice per call. + // The cached IP slices in luaRouteTestContext would undercount this cost. + return &routing_session.Context{ + Inbound: &session.Inbound{ + Tag: "in", + Source: net.TCPDestination(net.LocalHostIP, 1234), + Local: net.TCPDestination(net.LocalHostIP, 5678), + }, + Outbound: &session.Outbound{Target: target}, + Content: &session.Content{Protocol: "tls"}, + } +} + +func benchmarkRouteState(b *testing.B, r *Router, script string) *lua.LState { + b.Helper() + L := lua.NewState() + b.Cleanup(L.Close) + r.RegisterLua(L) + geodata.RegisterLua(L) + if err := L.DoString(script); err != nil { + b.Fatal(err) + } + L.SetContext(context.Background()) + return L +} + +// BenchmarkLuaRouteHook isolates argument bridging and a fixed-return hook. +// It excludes rules, result decoding, the state pool, and Route construction. +func BenchmarkLuaRouteHook(b *testing.B) { + L := benchmarkRouteState(b, new(Router), ` +function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort, + targetDomain, network, protocol, user, vlessRoute, skipDNSResolve) + return "out", "rule" +end +`) + ctx := benchmarkRouteContext(net.TCPDestination(net.LocalHostIP, 443)) + if err := callLuaRoute(L, ctx); err != nil { + b.Fatal(err) + } + outboundTag, ruleTag, err := readLuaRouteResult(L) + L.Pop(3) + if err != nil || outboundTag != "out" || ruleTag != "rule" { + b.Fatalf("hook() = %q, %q, %v", outboundTag, ruleTag, err) + } + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + if err := callLuaRoute(L, ctx); err != nil { + b.Fatal(err) + } + L.Pop(3) + } +} + +// BenchmarkLuaRoute compares equivalent ordered rules on the same session. +// rules returns tags only; pick_route uses Router.PickRoute on both sides. +// All compilation, matcher construction, and pool startup are outside timing. +func BenchmarkLuaRoute(b *testing.B) { + for _, name := range []string{"scalar", "ip", "domain", "domain_32_last"} { + b.Run(name, func(b *testing.B) { + config, script, ctx, wantTag, wantRule := benchmarkRouteFixture(b, name) + native := new(Router) + if err := native.Init(context.Background(), config, nil, nil, nil); err != nil { + b.Fatal(err) + } + L := benchmarkRouteState(b, native, script) + + path := filepath.Join(b.TempDir(), "route.lua") + if err := os.WriteFile(path, []byte(script), 0o600); err != nil { + b.Fatal(err) + } + scripted := new(Router) + if err := scripted.Init(context.Background(), &Config{Script: path}, nil, nil, nil); err != nil { + b.Fatal(err) + } + if err := scripted.Start(); err != nil { + b.Fatal(err) + } + b.Cleanup(func() { + if err := scripted.Close(); err != nil { + b.Error(err) + } + }) + + for _, bench := range []struct { + name string + route func() (string, string, error) + }{ + {"rules/native", func() (string, string, error) { + rule, _, err := native.pickRouteInternal(ctx) + if err != nil { + return "", "", err + } + tag, err := rule.GetTag() + return tag, rule.RuleTag, err + }}, + {"rules/lua", func() (string, string, error) { + if err := callLuaRoute(L, ctx); err != nil { + return "", "", err + } + tag, ruleTag, err := readLuaRouteResult(L) + L.Pop(3) + return tag, ruleTag, err + }}, + {"pick_route/native", func() (string, string, error) { + return benchmarkPickRoute(native, ctx) + }}, + {"pick_route/lua", func() (string, string, error) { + return benchmarkPickRoute(scripted, ctx) + }}, + } { + b.Run(bench.name, func(b *testing.B) { + // Validate and warm both paths before measuring steady state. + tag, ruleTag, err := bench.route() + if err != nil || tag != wantTag || ruleTag != wantRule { + b.Fatalf("route() = %q, %q, %v; want %q, %q", tag, ruleTag, err, wantTag, wantRule) + } + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + tag, ruleTag, err = bench.route() + if err != nil { + b.Fatal(err) + } + } + b.StopTimer() + if tag != wantTag || ruleTag != wantRule { + b.Fatalf("route() = %q, %q; want %q, %q", tag, ruleTag, wantTag, wantRule) + } + }) + } + }) + } +} + +func benchmarkPickRoute(r *Router, ctx routing.Context) (string, string, error) { + route, err := r.PickRoute(ctx) + if err != nil { + return "", "", err + } + return route.GetOutboundTag(), route.GetRuleTag(), nil +} + +func benchmarkRouteFixture(b *testing.B, name string) (*Config, string, routing.Context, string, string) { + b.Helper() + config := new(Config) + ctx := benchmarkRouteContext(net.TCPDestination(net.LocalHostIP, 443)) + prelude := `local router = require("xray.router") +local geodata = require("xray.geodata") +` + body := `if inboundTag == "in" and network == router.NetworkTCP then return "out", "rule" end` + wantTag, wantRule := "out", "rule" + if name == "scalar" || name == "ip" { + rule := &RoutingRule{ + TargetTag: &RoutingRule_Tag{Tag: wantTag}, + RuleTag: wantRule, + InboundTag: []string{"in"}, + Networks: []net.Network{net.Network_TCP}, + } + if name == "ip" { + var err error + rule.Ip, err = geodata.ParseIPRules([]string{"127.0.0.0/8"}) + if err != nil { + b.Fatal(err) + } + prelude += `local matcher = geodata.BuildIPMatcher("127.0.0.0/8")` + "\n" + body = `if inboundTag == "in" and network == router.NetworkTCP and matcher:AnyMatch(ctx:GetTargetIPs()) then return "out", "rule" end` + } + config.Rule = []*RoutingRule{rule} + } else { + count := 1 + if name == "domain_32_last" { + count = 32 + } + var rules strings.Builder + rules.WriteString("local rules = {\n") + for i := 0; i < count; i++ { + domain := fmt.Sprintf("route-%d.example.com", i) + tag, ruleTag := fmt.Sprintf("out-%d", i), fmt.Sprintf("rule-%d", i) + domains, err := geodata.ParseDomainRules([]string{"full:" + domain}, geodata.Domain_Domain) + if err != nil { + b.Fatal(err) + } + config.Rule = append(config.Rule, &RoutingRule{ + TargetTag: &RoutingRule_Tag{Tag: tag}, RuleTag: ruleTag, Domain: domains, + }) + fmt.Fprintf(&rules, "{geodata.BuildDomainMatcher(%q), %q, %q},\n", "full:"+domain, tag, ruleTag) + if i == count-1 { + ctx.Outbound.Target = net.TCPDestination(net.DomainAddress(domain), 443) + wantTag, wantRule = tag, ruleTag + } + } + rules.WriteString("}\n") + prelude += rules.String() + body = `for i = 1, #rules do + local rule = rules[i] + if rule[1]:MatchAny(targetDomain) then return rule[2], rule[3] end + end` + } + script := prelude + `function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort, + targetDomain, network, protocol, user, vlessRoute, skipDNSResolve) + ` + body + "\nend\n" + return config, script, ctx, wantTag, wantRule +} diff --git a/app/router/lua_test.go b/app/router/lua_test.go new file mode 100644 index 000000000..bf71e564b --- /dev/null +++ b/app/router/lua_test.go @@ -0,0 +1,274 @@ +package router + +import ( + "context" + go_errors "errors" + "runtime" + "strings" + "testing" + + "github.com/xtls/xray-core/common/geodata" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/common/protocol" + "github.com/xtls/xray-core/common/session" + "github.com/xtls/xray-core/features/routing" + routing_session "github.com/xtls/xray-core/features/routing/session" + lua "github.com/yuin/gopher-lua" +) + +type luaRouteTestContext struct { + *routing_session.Context + sourceIPs, targetIPs, localIPs []net.IP +} + +func (c *luaRouteTestContext) GetSourceIPs() []net.IP { return c.sourceIPs } +func (c *luaRouteTestContext) GetTargetIPs() []net.IP { return c.targetIPs } +func (c *luaRouteTestContext) GetLocalIPs() []net.IP { return c.localIPs } + +func newLuaRouteTestContext() *luaRouteTestContext { + return &luaRouteTestContext{ + Context: &routing_session.Context{ + Inbound: &session.Inbound{ + Tag: "in", VlessRoute: 4321, + Source: net.TCPDestination(net.LocalHostIP, 1234), + Local: net.TCPDestination(net.LocalHostIP, 5678), + User: &protocol.MemoryUser{Email: "user@example.com"}, + }, + Outbound: &session.Outbound{ + Target: net.TCPDestination(net.LocalHostIP, 443), + RouteTarget: net.TCPDestination(net.DomainAddress("MiXeD.Example."), 443), + }, + Content: &session.Content{ + Protocol: "tls", Attributes: map[string]string{"key": "value"}, SkipDNSResolve: true, + }, + }, + sourceIPs: []net.IP{{127, 0, 0, 2}}, + targetIPs: []net.IP{{127, 0, 0, 3}}, + localIPs: []net.IP{{127, 0, 0, 1}}, + } +} + +func newLuaRouterState(t *testing.T, script string) *lua.LState { + t.Helper() + r := new(Router) + if err := r.Init(context.Background(), &Config{}, nil, nil, nil); err != nil { + t.Fatal(err) + } + L := lua.NewState() + t.Cleanup(L.Close) + r.RegisterLua(L) + geodata.RegisterLua(L) + if err := L.DoString(script); err != nil { + t.Fatal(err) + } + return L +} + +func TestLuaRouteBinding(t *testing.T) { + L := newLuaRouterState(t, ` +local router = require("xray.router") +local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8") +assert(router.NetworkUnknown == 0 and router.NetworkTCP == 2) +assert(router.NetworkUDP == 3 and router.NetworkUNIX == 4) +assert(router.BuildIPMatcher == nil and router.BuildDomainMatcher == nil) +function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort, + targetDomain, network, protocol, user, vlessRoute, skipDNSResolve, ...) + assert(select("#", ...) == 0) + assert(inboundTag == "in" and sourcePort == 1234 and targetPort == 443 and localPort == 5678) + assert(targetDomain == "mixed.example." and network == router.NetworkTCP) + assert(protocol == "tls" and user == "user@example.com" and vlessRoute == 4321 and skipDNSResolve) + assert(ctx.GetNetwork == nil and ctx.Context == nil) + savedContext = ctx + sourceIPs, targetIPs, localIPs = ctx:GetSourceIPs(), ctx:GetTargetIPs(), ctx:GetLocalIPs() + attributes = ctx:GetAttributes() + assert(#sourceIPs == 1 and #targetIPs == 1 and #localIPs == 1) + assert(sourceIPs[1]:String() == "127.0.0.2" and targetIPs[1]:String() == "127.0.0.3") + assert(localIPs[1]:String() == "127.0.0.1") + assert(matcher:Match(sourceIPs[1]) and matcher:Match(targetIPs[1]) and matcher:Match(localIPs[1])) + assert(matcher:AnyMatch(sourceIPs) and matcher:AnyMatch(targetIPs) and matcher:AnyMatch(localIPs)) + local matched = matcher:FilterIPs(targetIPs) + assert(#matched == 1 and matched[1]:Equal(targetIPs[1])) + assert(attributes.key == "value" and attributes.missing == nil) + assert(not pcall(function() attributes.key = "changed" end)) + return "out", "rule" +end`) + + ctx := newLuaRouteTestContext() + if err := callLuaRoute(L, ctx); err != nil { + t.Fatal(err) + } + if L.GetTop() != 3 || L.Get(1) != lua.LString("out") || L.Get(2) != lua.LString("rule") || L.Get(3) != lua.LNil { + t.Fatal("callLuaRoute did not leave the three route results on the stack") + } + if L.GetGlobal("savedContext").(*lua.LUserData).Value != ctx { + t.Fatal("routing context was copied") + } + for _, tc := range []struct { + name string + want []net.IP + }{ + {"sourceIPs", ctx.sourceIPs}, + {"targetIPs", ctx.targetIPs}, + {"localIPs", ctx.localIPs}, + } { + got := L.GetGlobal(tc.name).(*lua.LUserData).Value.([]net.IP) + if &got[0] != &tc.want[0] { + t.Fatalf("%s storage was copied", tc.name) + } + } + ctx.Content.Attributes["key"] = "updated" + L.SetGlobal("expectedOS", lua.LString(runtime.GOOS)) + if err := L.DoString(` +assert(attributes.key == "updated") +assert(require("xray.router").LocalOS == expectedOS)`); err != nil { + t.Fatal(err) + } +} + +func TestLuaRouteEmptyIPs(t *testing.T) { + for _, tc := range []struct { + name string + ips []net.IP + }{ + {"nil", nil}, + {"empty", []net.IP{}}, + } { + t.Run(tc.name, func(t *testing.T) { + L := newLuaRouterState(t, ` +function HandleRoute(ctx) + for _, name in ipairs({"GetSourceIPs", "GetTargetIPs", "GetLocalIPs"}) do + local ips = ctx[name](ctx) + if expectNil then + assert(ips == nil) + else + assert(type(ips) == "userdata" and #ips == 0) + assert(not pcall(function() return ips[1] end)) + end + end + return "out" +end +`) + L.SetGlobal("expectNil", lua.LBool(tc.ips == nil)) + ctx := newLuaRouteTestContext() + ctx.sourceIPs, ctx.targetIPs, ctx.localIPs = tc.ips, tc.ips, tc.ips + if err := callLuaRoute(L, ctx); err != nil { + t.Fatal(err) + } + }) + } +} + +func TestReadLuaRouteResult(t *testing.T) { + nativeErr := go_errors.New("native failure") + for _, tc := range []struct { + name, values string + wantTag, wantRule string + wantErr error + wantMessage string + }{ + {name: "route", values: `"out", "rule"`, wantTag: "out", wantRule: "rule"}, + {name: "no match", values: `nil`}, + {name: "empty tag", values: `""`}, + {name: "no match ignores rule", values: `nil, false`}, + {name: "empty tag ignores rule", values: `"", false`}, + {name: "missing rule", values: `"out"`, wantTag: "out"}, + {name: "invalid tag", values: `1`, wantMessage: "outboundTag"}, + {name: "invalid rule", values: `"out", false`, wantMessage: "ruleTag"}, + {name: "string error", values: `nil, nil, "script failure"`, wantMessage: "script failure"}, + {name: "native error", values: `nil, nil, nativeError`, wantErr: nativeErr}, + {name: "error overrides invalid tags", values: `false, false, nativeError`, wantErr: nativeErr}, + {name: "invalid error", values: `"out", "rule", false`, wantMessage: "error or string"}, + {name: "wrong error userdata", values: `"out", "rule", wrongError`, wantMessage: "error or string"}, + } { + t.Run(tc.name, func(t *testing.T) { + L := lua.NewState() + defer L.Close() + for name, value := range map[string]any{"nativeError": nativeErr, "wrongError": "not a native error"} { + ud := L.NewUserData() + ud.Value = value + L.SetGlobal(name, ud) + } + fn, err := L.LoadString("return " + tc.values) + if err != nil { + t.Fatal(err) + } + if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true}); err != nil { + t.Fatal(err) + } + outboundTag, ruleTag, err := readLuaRouteResult(L) + if outboundTag != tc.wantTag || ruleTag != tc.wantRule { + t.Fatalf("result = %q, %q, %v; want %q, %q", outboundTag, ruleTag, err, tc.wantTag, tc.wantRule) + } + switch { + case tc.wantErr != nil: + if err != tc.wantErr { + t.Fatalf("error = %v, want original error", err) + } + case tc.wantMessage != "": + if err == nil || !strings.Contains(err.Error(), tc.wantMessage) { + t.Fatalf("error = %v, want %q", err, tc.wantMessage) + } + case err != nil: + t.Fatal(err) + } + }) + } +} + +func TestCallLuaRouteCancellation(t *testing.T) { + L := newLuaRouterState(t, `function HandleRoute() while true do end end`) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + L.SetContext(ctx) + if err := callLuaRoute(L, &routing_session.Context{}); err == nil { + t.Fatal("callLuaRoute did not stop after context cancellation") + } + if L.Context() != ctx { + t.Fatal("callLuaRoute changed the Lua state's context") + } +} + +func TestFindProcess(t *testing.T) { + for _, tc := range []struct { + name, network, target string + targetPort uint16 + modify func(*luaRouteTestContext) + wantErr bool + }{ + {name: "TCP", network: "tcp", target: "127.0.0.3", targetPort: 443}, + {name: "UDP", network: "udp", target: "127.0.0.3", targetPort: 443, modify: func(c *luaRouteTestContext) { + c.Outbound.Target.Network = net.Network_UDP + }}, + {name: "domain target", network: "tcp", modify: func(c *luaRouteTestContext) { c.targetIPs = nil }}, + {name: "missing source", modify: func(c *luaRouteTestContext) { c.sourceIPs = nil }, wantErr: true}, + {name: "unsupported network", modify: func(c *luaRouteTestContext) { + c.Outbound.Target.Network = net.Network_UNIX + }, wantErr: true}, + } { + t.Run(tc.name, func(t *testing.T) { + ctx := newLuaRouteTestContext() + if tc.modify != nil { + tc.modify(ctx) + } + called := false + pid, name, path, err := findProcess(ctx, func(network, source string, sourcePort uint16, target string, targetPort uint16) (int, string, string, error) { + called = true + if network != tc.network || source != "127.0.0.2" || sourcePort != 1234 || target != tc.target || targetPort != tc.targetPort { + t.Fatalf("endpoints = %s %s:%d -> %s:%d", network, source, sourcePort, target, targetPort) + } + return 42, "process", "/path/process", nil + }) + if tc.wantErr { + if err == nil || called { + t.Fatalf("findProcess = %d, %q, %q, %v", pid, name, path, err) + } + return + } + if err != nil || !called || pid != 42 || name != "process" || path != "/path/process" { + t.Fatalf("findProcess = %d, %q, %q, %v", pid, name, path, err) + } + }) + } +} + +var _ routing.Context = (*luaRouteTestContext)(nil) diff --git a/app/router/router.go b/app/router/router.go index 406f932fb..97ffba13c 100644 --- a/app/router/router.go +++ b/app/router/router.go @@ -20,6 +20,8 @@ import ( type Router struct { domainStrategy Config_DomainStrategy rules atomic.Pointer[[]*Rule] + scriptPath string + script *scriptEngine balancers atomic.Pointer[map[string]*Balancer] dns dns.Client @@ -40,6 +42,7 @@ type Route struct { // Init initializes the Router. func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm outbound.Manager, dispatcher routing.Dispatcher) error { r.domainStrategy = config.DomainStrategy + r.scriptPath = config.Script r.dns = d r.ctx = ctx r.ohm = ohm @@ -52,6 +55,10 @@ func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm out // PickRoute implements routing.Router. func (r *Router) PickRoute(ctx routing.Context) (routing.Route, error) { + if r.script != nil { + return r.script.pickRoute(ctx) + } + originalCtx := ctx rule, ctx, err := r.pickRouteInternal(ctx) if err != nil { @@ -221,6 +228,13 @@ func (r *Router) pickRouteInternal(ctx routing.Context) (*Rule, routing.Context, // Start implements common.Runnable. func (r *Router) Start() error { + if r.scriptPath != "" { + engine, err := newScriptEngine(r.scriptPath, r) + if err != nil { + return errors.New("failed to initialize routing script").Base(err) + } + r.script = engine + } return nil } @@ -235,6 +249,9 @@ func closeWebhooks(rules []*Rule) { // Close implements common.Closable. func (r *Router) Close() error { + if r.script != nil { + r.script.close() + } r.mu.Lock() defer r.mu.Unlock() closeWebhooks(*r.rules.Load()) diff --git a/app/router/script.go b/app/router/script.go new file mode 100644 index 000000000..b4a720b56 --- /dev/null +++ b/app/router/script.go @@ -0,0 +1,76 @@ +package router + +import ( + "time" + + "github.com/xtls/xray-core/app/dns" + "github.com/xtls/xray-core/common" + "github.com/xtls/xray-core/common/errors" + "github.com/xtls/xray-core/common/geodata" + "github.com/xtls/xray-core/common/log" + xlua "github.com/xtls/xray-core/common/lua" + "github.com/xtls/xray-core/features/routing" + lua "github.com/yuin/gopher-lua" +) + +const scriptExecutionTimeout = 6 * time.Second + +type scriptEngine struct { + pool *xlua.Pool +} + +func newScriptEngine(path string, router *Router) (*scriptEngine, error) { + program, err := xlua.CompileFile(path) + if err != nil { + return nil, err + } + + pool, err := xlua.NewPool(router.ctx, scriptExecutionTimeout, program.NewStateFactory( + scriptExecutionTimeout*20, + func(L *lua.LState) { + geodata.RegisterLua(L) + log.RegisterLua(L) + router.RegisterLua(L) + dns.RegisterLua(L, router.dns) + }, + func(L *lua.LState) error { + if L.GetGlobal("HandleRoute").Type() != lua.LTFunction { + return errors.New("routing script must define HandleRoute(...)") + } + return nil + })) + if err != nil { + return nil, err + } + + errors.LogInfo(router.ctx, "routing script initialized from ", path) + return &scriptEngine{pool: pool}, nil +} + +func (e *scriptEngine) close() { + e.pool.Close() +} + +func (e *scriptEngine) pickRoute(ctx routing.Context) (routing.Route, error) { + var outboundTag, ruleTag string + var routeErr error + + if err := e.pool.WithState(nil, 0, func(L *lua.LState) error { + if err := callLuaRoute(L, ctx); err != nil { + return err + } + outboundTag, ruleTag, routeErr = readLuaRouteResult(L) + return nil + }); err != nil { + return nil, err + } + + if routeErr != nil { + return nil, routeErr + } + if outboundTag == "" { + return nil, common.ErrNoClue + } + + return &Route{Context: ctx, outboundTag: outboundTag, ruleTag: ruleTag}, nil +} diff --git a/app/router/script_test.go b/app/router/script_test.go new file mode 100644 index 000000000..7daeb868a --- /dev/null +++ b/app/router/script_test.go @@ -0,0 +1,382 @@ +package router + +import ( + "context" + stdnet "net" + "os" + "path/filepath" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + wireDNS "github.com/miekg/dns" + "github.com/xtls/xray-core/app/dispatcher" + appdns "github.com/xtls/xray-core/app/dns" + "github.com/xtls/xray-core/app/proxyman" + _ "github.com/xtls/xray-core/app/proxyman/outbound" + "github.com/xtls/xray-core/common" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/common/serial" + "github.com/xtls/xray-core/core" + featureDNS "github.com/xtls/xray-core/features/dns" + "github.com/xtls/xray-core/features/outbound" + "github.com/xtls/xray-core/features/routing" + routing_session "github.com/xtls/xray-core/features/routing/session" + "github.com/xtls/xray-core/proxy/blackhole" + "github.com/xtls/xray-core/proxy/freedom" +) + +type luaRouteDNSClient struct { + featureDNS.Client + lookup func(string, featureDNS.IPOption) ([]net.IP, uint32, error) +} + +func (d *luaRouteDNSClient) LookupIP(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) { + return d.lookup(domain, option) +} + +type luaRouteOutboundManager struct{ outbound.Manager } + +func (*luaRouteOutboundManager) Select(selectors []string) []string { return selectors } + +func writeRouteScript(t *testing.T, script string) string { + t.Helper() + path := filepath.Join(t.TempDir(), "route.lua") + if err := os.WriteFile(path, []byte(script), 0o600); err != nil { + t.Fatal(err) + } + return path +} + +func startLuaRouter(t *testing.T, script string, d featureDNS.Client, config *Config) *Router { + t.Helper() + if config == nil { + config = &Config{} + } + config.Script = writeRouteScript(t, script) + r := new(Router) + if err := r.Init(context.Background(), config, d, &luaRouteOutboundManager{}, nil); err != nil { + t.Fatal(err) + } + if err := r.Start(); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + if err := r.Close(); err != nil { + t.Error(err) + } + }) + return r +} + +func TestRouterScriptStartup(t *testing.T) { + for _, tc := range []struct{ name, script string }{ + {"syntax error", "function HandleRoute("}, + {"missing hook", "value = 1"}, + {"initialization error", `error("setup failed")`}, + } { + t.Run(tc.name, func(t *testing.T) { + r := new(Router) + if err := r.Init(context.Background(), &Config{Script: writeRouteScript(t, tc.script)}, nil, nil, nil); err != nil { + t.Fatal(err) + } + defer r.Close() + if err := r.Start(); err == nil { + t.Fatal("Start accepted an invalid routing script") + } + }) + } +} + +func TestRouterScriptRouting(t *testing.T) { + for _, tc := range []struct { + name, body string + wantTag, wantRule string + wantErr error + wantMessage string + wantCalls string + }{ + {name: "route", body: `return "lua-out", "lua-rule"`, wantTag: "lua-out", wantRule: "lua-rule", wantCalls: "2"}, + {name: "no match", body: `return nil`, wantErr: common.ErrNoClue, wantCalls: "2"}, + {name: "empty tag", body: `return ""`, wantErr: common.ErrNoClue, wantCalls: "2"}, + {name: "balancer error", body: `local tag, err = router:PickOutbound("missing"); return tag, nil, err`, wantMessage: "not found", wantCalls: "2"}, + {name: "string error", body: `return nil, nil, "blocked"`, wantMessage: "blocked", wantCalls: "2"}, + {name: "invalid tag", body: `return false`, wantMessage: "outboundTag", wantCalls: "2"}, + {name: "invalid rule", body: `return "lua-out", false`, wantMessage: "ruleTag", wantCalls: "2"}, + {name: "execution error", body: `error("execution failed")`, wantMessage: "execution failed", wantCalls: "1"}, + } { + t.Run(tc.name, func(t *testing.T) { + var dnsCalls atomic.Int32 + d := &luaRouteDNSClient{lookup: func(string, featureDNS.IPOption) ([]net.IP, uint32, error) { + dnsCalls.Add(1) + return []net.IP{{1, 2, 3, 4}}, 60, nil + }} + script := ` +local router = require("xray.router") +local calls = 0 +function HandleRoute(ctx, inbound) + calls = calls + 1 + if inbound == "count" then return "lua-out", tostring(calls) end + ` + tc.body + ` +end +` + r := startLuaRouter(t, script, d, &Config{ + DomainStrategy: Config_IpOnDemand, + Rule: []*RoutingRule{{ + TargetTag: &RoutingRule_Tag{Tag: "json-out"}, + Networks: []net.Network{net.Network_TCP}, + }}, + }) + ctx := newLuaRouteTestContext() + ctx.Content.SkipDNSResolve = false + route, err := r.PickRoute(ctx) + switch { + case tc.wantErr != nil: + if err != tc.wantErr { + t.Fatalf("route error = %v, want %v", err, tc.wantErr) + } + case tc.wantMessage != "": + if err == nil || !strings.Contains(err.Error(), tc.wantMessage) { + t.Fatalf("route error = %v, want %q", err, tc.wantMessage) + } + case err != nil: + t.Fatal(err) + } + if tc.wantTag == "" { + if route != nil { + t.Fatalf("route = %v, want nil", route) + } + } else if route == nil || route.GetOutboundTag() != tc.wantTag || route.GetRuleTag() != tc.wantRule || route.(*Route).Context != ctx { + t.Fatalf("route = %v; want %q, %q and original context", route, tc.wantTag, tc.wantRule) + } + + ctx.Inbound.Tag = "count" + route, err = r.PickRoute(ctx) + if err != nil || route == nil || route.GetOutboundTag() != "lua-out" || route.GetRuleTag() != tc.wantCalls { + t.Fatalf("next route = %v, %v; want lua-out, calls %s", route, err, tc.wantCalls) + } + if dnsCalls.Load() != 0 { + t.Fatal("script routing implicitly resolved DNS") + } + }) + } +} + +func TestRouterScriptModules(t *testing.T) { + ips := []net.IP{{127, 0, 0, 7}} + calls := 0 + d := &luaRouteDNSClient{lookup: func(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) { + calls++ + if domain != "mixed.example." || !option.IPv4Enable || option.IPv6Enable || !option.FakeEnable { + t.Fatalf("dns.Query arguments = %q, %+v", domain, option) + } + return ips, 17, nil + }} + r := startLuaRouter(t, ` +local dns = require("xray.dns") +local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8") +assert(dns.Servers == nil and type(dns.Query) == "function") +assert(type(require("xray.log").Info) == "function") +function HandleRoute(ctx, inbound, sourcePort, targetPort, localPort, domain) + local ips, ttl, err = dns.Query(domain, true, false, true) + assert(not err and ttl == 17) + assert(matcher:AnyMatch(ips) and matcher:AnyMatch(ctx:GetTargetIPs())) + return "out" +end`, d, nil) + if _, err := r.PickRoute(newLuaRouteTestContext()); err != nil { + t.Fatal(err) + } + if calls != 1 { + t.Fatalf("DNS calls = %d, want 1", calls) + } +} + +func TestRouterScriptBalancerReload(t *testing.T) { + config := func(tag string) *Config { + return &Config{BalancingRule: []*BalancingRule{{ + Tag: "balance", Strategy: "roundrobin", OutboundSelector: []string{tag}, + }}} + } + r := startLuaRouter(t, ` +local router = require("xray.router") +function HandleRoute() + local tag, err = router:PickOutbound("balance") + return tag, "balanced", err +end`, nil, config("old")) + pick := func(want string) { + t.Helper() + route, err := r.PickRoute(&routing_session.Context{}) + if err != nil || route.GetOutboundTag() != want || route.GetRuleTag() != "balanced" { + t.Fatalf("route = %v, %v, want %q", route, err, want) + } + } + + pick("old") + if err := r.SetOverrideTarget("balance", "override"); err != nil { + t.Fatal(err) + } + pick("override") + if err := r.SetOverrideTarget("balance", ""); err != nil { + t.Fatal(err) + } + if err := r.ReloadRules(config("new"), false); err != nil { + t.Fatal(err) + } + pick("new") +} + +func TestRouterScriptConcurrentBalancerReload(t *testing.T) { + config := func(tag string) *Config { + return &Config{BalancingRule: []*BalancingRule{{ + Tag: "balance", Strategy: "roundrobin", OutboundSelector: []string{tag}, + }}} + } + r := startLuaRouter(t, ` +local router = require("xray.router") +function HandleRoute() + local tag, err = router:PickOutbound("balance") + return tag, nil, err +end`, nil, config("a")) + + var wg sync.WaitGroup + for range 4 { + wg.Go(func() { + for range 20 { + route, err := r.PickRoute(&routing_session.Context{}) + if err != nil { + t.Errorf("PickRoute: %v", err) + return + } + if tag := route.GetOutboundTag(); tag != "a" && tag != "b" { + t.Errorf("unexpected tag %q", tag) + } + } + }) + } + wg.Go(func() { + for range 20 { + for _, tag := range []string{"a", "b"} { + if err := r.ReloadRules(config(tag), false); err != nil { + t.Error(err) + return + } + } + } + }) + wg.Wait() +} + +func TestRouterScriptDNSDispatcherReentry(t *testing.T) { + conn, err := stdnet.ListenPacket("udp4", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + port := conn.LocalAddr().(*stdnet.UDPAddr).Port + ready, stopped := make(chan struct{}), make(chan error, 1) + var queries atomic.Int32 + server := &wireDNS.Server{ + PacketConn: conn, + NotifyStartedFunc: func() { + close(ready) + }, + Handler: wireDNS.HandlerFunc(func(w wireDNS.ResponseWriter, query *wireDNS.Msg) { + queries.Add(1) + response := new(wireDNS.Msg).SetReply(query) + for _, question := range query.Question { + if question.Name == "nested.example." && question.Qtype == wireDNS.TypeA { + response.Answer = append(response.Answer, &wireDNS.A{ + Hdr: wireDNS.RR_Header{Name: question.Name, Rrtype: wireDNS.TypeA, Class: wireDNS.ClassINET, Ttl: 60}, + A: stdnet.IP{127, 0, 0, 7}, + }) + } + } + if err := w.WriteMsg(response); err != nil { + t.Error(err) + } + }), + } + go func() { stopped <- server.ActivateAndServe() }() + defer func() { + server.Shutdown() + select { + case err := <-stopped: + if err != nil { + t.Error(err) + } + case <-time.After(3 * time.Second): + t.Error("DNS server did not stop") + } + }() + select { + case <-ready: + case err := <-stopped: + t.Fatalf("DNS server startup: %v", err) + case <-time.After(3 * time.Second): + t.Fatal("DNS server did not start") + } + + dnsScript := writeRouteScript(t, ` +local server = require("xray.dns").Servers[1] +function HandleDNSQuery(domain, ipv4, ipv6, fake) + return server:Query(domain, ipv4, ipv6, fake) +end`) + routerScript := writeRouteScript(t, ` +local router = require("xray.router") +local dns = require("xray.dns") +local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.7") +local active = false +function HandleRoute(ctx, inbound, sourcePort, targetPort, localPort, domain, network, + protocol, user, vlessRoute, skipDNSResolve) + assert(not active, "borrowed Router VM reentered") + if inbound == "dns" then + assert(network == router.NetworkUDP and skipDNSResolve == false) + return "direct", "dns-route" + end + active = true + local ips, ttl, err = dns.Query("nested.example", true, false, false) + assert(not err and matcher:AnyMatch(ips) and active) + active = false + return "direct", "outer-route" +end`) + instance, err := core.New(&core.Config{ + App: []*serial.TypedMessage{ + serial.ToTypedMessage(&appdns.Config{ + Tag: "dns", Script: dnsScript, DisableCache: true, + NameServer: []*appdns.NameServer{{ + Id: "upstream", TimeoutMs: 1000, + Address: &net.Endpoint{ + Network: net.Network_UDP, + Address: &net.IPOrDomain{Address: &net.IPOrDomain_Ip{Ip: []byte{127, 0, 0, 1}}}, + Port: uint32(port), + }, + }}, + }), + serial.ToTypedMessage(&Config{Script: routerScript}), + serial.ToTypedMessage(&dispatcher.Config{}), + serial.ToTypedMessage(&proxyman.OutboundConfig{}), + }, + Outbound: []*core.OutboundHandlerConfig{ + {Tag: "default", ProxySettings: serial.ToTypedMessage(&blackhole.Config{})}, + {Tag: "direct", ProxySettings: serial.ToTypedMessage(&freedom.Config{ + FinalRules: []*freedom.FinalRuleConfig{{Action: freedom.RuleAction_Allow}}, + })}, + }, + }) + if err != nil { + t.Fatal(err) + } + defer instance.Close() + if err := instance.Start(); err != nil { + t.Fatal(err) + } + r := instance.GetFeature(routing.RouterType()).(*Router) + route, err := r.PickRoute(newLuaRouteTestContext()) + if err != nil || route.GetOutboundTag() != "direct" || route.GetRuleTag() != "outer-route" { + t.Fatalf("nested DNS routing = %v, %v", route, err) + } + if queries.Load() == 0 { + t.Fatal("DNS query did not pass through the dispatcher") + } +} diff --git a/common/geodata/lua.go b/common/geodata/lua.go new file mode 100644 index 000000000..43f7363ff --- /dev/null +++ b/common/geodata/lua.go @@ -0,0 +1,163 @@ +package geodata + +import ( + xlua "github.com/xtls/xray-core/common/lua" + "github.com/xtls/xray-core/common/net" + lua "github.com/yuin/gopher-lua" +) + +// RegisterLua makes xray.geodata available to require in an LState. +func RegisterLua(L *lua.LState) { + L.PreloadModule("xray.geodata", func(L *lua.LState) int { + module := L.CreateTable(0, 2) + + module.RawSetString("BuildDomainMatcher", L.NewFunction(func(L *lua.LState) int { + parsed, err := ParseDomainRules(luaRules(L), 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 + } + xlua.PushWithDirectMethods(L, matcher, map[string]xlua.DirectMethod{ + "Match": newLuaDomainMatch(xlua.NewSlicePusher[uint32](L)), + "MatchAny": luaDomainMatchAny, + }) + return 1 + })) + + module.RawSetString("BuildIPMatcher", L.NewFunction(func(L *lua.LState) int { + parsed, err := ParseIPRules(luaRules(L)) + if err != nil { + L.RaiseError("%v", err) + return 0 + } + matcher, err := IPReg.BuildIPMatcher(parsed) + if err != nil { + L.RaiseError("%v", err) + return 0 + } + xlua.PushWithDirectMethods(L, matcher, map[string]xlua.DirectMethod{ + "Match": luaIPMatch, + "AnyMatch": luaIPAnyMatch, + "Matches": luaIPMatches, + "FilterIPs": newLuaIPFilterIPs(xlua.NewSlicePusher[net.IP](L)), + }) + return 1 + })) + + L.Push(module) + return 1 + }) +} + +// Read native Go values by type assertion; slices keep their original storage. +func readLuaIPMatcherArgs[T any](L *lua.LState) (IPMatcher, T, bool) { + var input T + if L.GetTop() != 2 { + return nil, input, false + } + value, ok := L.Get(1).(*lua.LUserData) + if !ok { + return nil, input, false + } + matcher, ok := value.Value.(IPMatcher) + if !ok { + return nil, input, false + } + if L.Get(2) == lua.LNil { + return matcher, input, true + } + value, ok = L.Get(2).(*lua.LUserData) + if !ok { + return nil, input, false + } + input, ok = value.Value.(T) + return matcher, input, ok +} + +func luaIPMatch(L *lua.LState) (int, bool) { + matcher, ip, ok := readLuaIPMatcherArgs[net.IP](L) + if !ok { + return 0, false + } + L.Push(lua.LBool(matcher.Match(ip))) + return 1, true +} + +func luaIPAnyMatch(L *lua.LState) (int, bool) { + matcher, ips, ok := readLuaIPMatcherArgs[[]net.IP](L) + if !ok { + return 0, false + } + L.Push(lua.LBool(matcher.AnyMatch(ips))) + return 1, true +} + +func luaIPMatches(L *lua.LState) (int, bool) { + matcher, ips, ok := readLuaIPMatcherArgs[[]net.IP](L) + if !ok { + return 0, false + } + L.Push(lua.LBool(matcher.Matches(ips))) + return 1, true +} + +func newLuaIPFilterIPs(pushIPs func(*lua.LState, []net.IP)) xlua.DirectMethod { + return func(L *lua.LState) (int, bool) { + matcher, ips, ok := readLuaIPMatcherArgs[[]net.IP](L) + if !ok { + return 0, false + } + matched, unmatched := matcher.FilterIPs(ips) + pushIPs(L, matched) + pushIPs(L, unmatched) + return 2, true + } +} + +func newLuaDomainMatch(pushMatches func(*lua.LState, []uint32)) xlua.DirectMethod { + return func(L *lua.LState) (int, bool) { + if L.GetTop() == 2 { + if value, ok := L.Get(1).(*lua.LUserData); ok { + matcher, validMatcher := value.Value.(DomainMatcher) + domain, validDomain := L.Get(2).(lua.LString) + if validMatcher && validDomain { + pushMatches(L, matcher.Match(string(domain))) + return 1, true + } + } + } + return 0, false + } +} + +func luaDomainMatchAny(L *lua.LState) (int, bool) { + if L.GetTop() == 2 { + if value, ok := L.Get(1).(*lua.LUserData); ok { + matcher, validMatcher := value.Value.(DomainMatcher) + domain, validDomain := L.Get(2).(lua.LString) + if validMatcher && validDomain { + L.Push(lua.LBool(matcher.MatchAny(string(domain)))) + return 1, true + } + } + } + return 0, false +} + +func luaRules(L *lua.LState) []string { + rules := make([]string, L.GetTop()) + for i := range rules { + value, ok := L.Get(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..f66a57bc4 --- /dev/null +++ b/common/geodata/lua_test.go @@ -0,0 +1,172 @@ +package geodata + +import ( + "fmt" + "testing" + + "github.com/xtls/xray-core/common/net" + lua "github.com/yuin/gopher-lua" +) + +func TestLuaIPMatcher(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) + ips := L.NewUserData() + ips.Value = []net.IP{ip.Value.(net.IP), net.ParseIP("8.8.8.8")} + L.SetGlobal("ips", ips) + if err := L.DoString(` + local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8", "::1") + assert(matcher:Match(ip)) + assert(matcher:AnyMatch(ips)) + assert(not matcher:Matches(ips)) + local matched, unmatched = matcher:FilterIPs(ips) + assert(type(matched) == "userdata" and type(unmatched) == "userdata") + assert(#matched == 1 and #unmatched == 1) + `); err != nil { + t.Fatal(err) + } +} + +func TestLuaDomainMatcher(t *testing.T) { + L := lua.NewState() + defer L.Close() + RegisterLua(L) + if err := L.DoString(` + local matcher = require("xray.geodata").BuildDomainMatcher("example.com", "full:other.com") + assert(matcher:MatchAny("example.com")) + assert(matcher:MatchAny("www.example.com")) + assert(matcher:MatchAny("other.com")) + assert(not matcher:MatchAny("www.other.com")) + assert(#(matcher:Match("www.example.com")) == 1) + `); err != nil { + t.Fatal(err) + } +} + +func TestLuaMatchersRejectInvalidRules(t *testing.T) { + for _, tc := range []struct { + name string + script string + }{ + {"IP rule", `require("xray.geodata").BuildIPMatcher("not-an-ip")`}, + {"non-string domain rule", `require("xray.geodata").BuildDomainMatcher("example.com", 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") + } + }) + } +} + +func TestLuaMatcherArgumentsAndAliases(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 geodata = require("xray.geodata") +local matcher = geodata.BuildIPMatcher("127.0.0.0/8") +assert(matcher.Match == matcher.match and matcher.AnyMatch == matcher.anyMatch) +assert(matcher.Matches == matcher.matches and matcher.FilterIPs == matcher.filterIPs) +assert(matcher:match(ip)) +assert(matcher:anyMatch({ip}) and matcher:matches({ip})) +assert(not matcher:AnyMatch(nil)) +assert(matcher:Matches(nil) == matcher:Matches({})) +local matched, unmatched = matcher:FilterIPs({ip}) +assert(#matched == 1 and matched[1]:Equal(ip)) +assert(matcher:AnyMatch(matched) and matcher:Matches(matched)) +local filtered, excluded = matcher:filterIPs(matched) +assert(#filtered == 1 and #excluded == 0 and filtered[1]:Equal(ip)) +local emptyMatched, emptyUnmatched = matcher:FilterIPs(nil) +assert(#emptyMatched == 0 and #emptyUnmatched == 0) +matcher:SetReverse(true) +assert(not matcher:Match(ip) and not matcher:AnyMatch(matched)) +matcher:ToggleReverse() +assert(matcher:Match(ip) and matcher:AnyMatch(matched)) +assert(matcher.missing == nil) + +local domain = geodata.BuildDomainMatcher("full:example.com") +assert(domain.Match == domain.match and domain.MatchAny == domain.matchAny) +assert(domain:matchAny("example.com")) +assert(#domain:Match("example.com") == 1) +assert(domain:match("example.com")[1] == 0) +assert(not pcall(function() matcher:AnyMatch() end)) +assert(not pcall(function() matcher:AnyMatch(matched, true) end)) +assert(not pcall(function() matcher.AnyMatch(ip, matched) end)) +assert(not pcall(function() matcher:Match(true) end)) +assert(not pcall(function() domain:MatchAny(123) end)) +assert(not pcall(function() domain:MatchAny("example.com", true) end)) +assert(not pcall(function() matcher:FilterIPs(true) end)) +assert(not pcall(function() matcher:FilterIPs(matched, true) end)) +assert(not pcall(function() domain:Match(123) end)) +assert(not pcall(function() domain:Match("example.com", true) end)) +`); err != nil { + t.Fatal(err) + } +} + +// BenchmarkLuaMatcherCall measures repeated calls with prebuilt matchers and inputs. +func BenchmarkLuaMatcherCall(b *testing.B) { + L := lua.NewState() + defer L.Close() + RegisterLua(L) + ip := net.ParseIP("127.0.0.1") + for name, value := range map[string]any{"ip": ip, "ips": []net.IP{ip}} { + ud := L.NewUserData() + ud.Value = value + L.SetGlobal(name, ud) + } + if err := L.DoString(` +local geodata = require("xray.geodata") +ipMatcher = geodata.BuildIPMatcher("127.0.0.0/8") +domainMatcher = geodata.BuildDomainMatcher("full:example.com") +`); err != nil { + b.Fatal(err) + } + for _, benchmark := range []struct { + name, expression string + }{ + {"ip_match", "ipMatcher:Match(ip)"}, + {"ip_match_lower", "ipMatcher:match(ip)"}, + {"ip_any_match", "ipMatcher:AnyMatch(ips)"}, + {"ip_any_match_lower", "ipMatcher:anyMatch(ips)"}, + {"ip_matches", "ipMatcher:Matches(ips)"}, + {"ip_matches_lower", "ipMatcher:matches(ips)"}, + {"domain_match_any", `domainMatcher:MatchAny("example.com")`}, + {"domain_match_any_lower", `domainMatcher:matchAny("example.com")`}, + {"ip_filter", "select(1, ipMatcher:FilterIPs(ips)) ~= nil"}, + {"ip_filter_lower", "select(1, ipMatcher:filterIPs(ips)) ~= nil"}, + {"domain_match", `#domainMatcher:Match("example.com") == 1`}, + {"domain_match_lower", `#domainMatcher:match("example.com") == 1`}, + {"ip_lua_table", "ipMatcher:AnyMatch({ip})"}, + {"ip_lua_table_lower", "ipMatcher:anyMatch({ip})"}, + } { + b.Run(benchmark.name, func(b *testing.B) { + if err := L.DoString(fmt.Sprintf("function benchmarkMatch() return %s end", benchmark.expression)); err != nil { + b.Fatal(err) + } + fn := L.GetGlobal("benchmarkMatch") + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + if err := L.CallByParam(lua.P{Fn: fn, NRet: 1, Protect: true}); err != nil { + b.Fatal(err) + } + if L.Get(-1) != lua.LTrue { + b.Fatal("matcher returned false") + } + L.Pop(1) + } + }) + } +} diff --git a/common/log/lua.go b/common/log/lua.go new file mode 100644 index 000000000..99db47baf --- /dev/null +++ b/common/log/lua.go @@ -0,0 +1,61 @@ +package log + +import ( + "path/filepath" + "strings" + + lua "github.com/yuin/gopher-lua" +) + +// RegisterLua makes xray.log available to require in an LState. +func RegisterLua(L *lua.LState) { + L.PreloadModule("xray.log", func(L *lua.LState) int { + module := L.CreateTable(0, 4) + var source, prefix string // cache + for name, severity := range map[string]Severity{ + "Debug": Severity_Debug, + "Info": Severity_Info, + "Warning": Severity_Warning, + "Error": Severity_Error, + } { + module.RawSetString(name, L.NewFunction(func(L *lua.LState) int { + if GetSeverity() < severity { + return 0 + } + var content strings.Builder + // Prefix with the calling script's filename. + if caller, ok := L.GetStack(1); ok { + if _, err := L.GetInfo("S", caller, lua.LNil); err == nil && caller.Source != "" { + if caller.Source != source { + source = caller.Source + prefix = filepath.Base(strings.TrimPrefix(source, "@")) + ": " + } + content.WriteString(prefix) + } + } + for i := 1; i <= L.GetTop(); i++ { + content.WriteString(luaLogString(L, L.Get(i))) + } + Record(&GeneralMessage{ + Severity: severity, + Content: content.String(), + }) + return 0 + })) + } + L.Push(module) + return 1 + }) +} + +func luaLogString(L *lua.LState, value lua.LValue) string { + if ud, ok := value.(*lua.LUserData); ok { + if err, ok := ud.Value.(error); ok { + return err.Error() + } + } + if _, ok := L.GetMetaField(value, "__tostring").(*lua.LFunction); ok { + return L.ToStringMeta(value).String() + } + return value.String() +} diff --git a/common/log/lua_test.go b/common/log/lua_test.go new file mode 100644 index 000000000..759bae537 --- /dev/null +++ b/common/log/lua_test.go @@ -0,0 +1,213 @@ +package log + +import ( + "errors" + "fmt" + "os" + "path/filepath" + "testing" + + lua "github.com/yuin/gopher-lua" +) + +type luaLogHandler struct { + messages []Message +} + +func (h *luaLogHandler) Handle(msg Message) { + h.messages = append(h.messages, msg) +} + +func TestLuaLog(t *testing.T) { + previous := logHandler.Load() + t.Cleanup(func() { logHandler.Store(previous) }) + handler := &luaLogHandler{} + RegisterHandler(handler) + + L := lua.NewState() + defer L.Close() + RegisterLua(L) + nativeError := L.NewUserData() + nativeError.Value = fmt.Errorf("lookup failed: %w", errors.New("upstream timeout")) + L.SetGlobal("nativeError", nativeError) + path := filepath.Join(t.TempDir(), "logging.lua") + if err := os.WriteFile(path, []byte(` + local log = require("xray.log") + assert(log == require("xray.log")) + log.Debug("query: ", "example.com") + log.Info("count=", 42, ", enabled=", true, ", value=", nil) + log.Warning(setmetatable({}, { + __tostring = function() return "fallback" end + })) + assert(select("#", log.Error("failed")) == 0) + log.Error("DNS failed: ", nativeError) + log.Warning(nativeError) + local ok, err = pcall(function() error("Lua failure", 0) end) + assert(not ok) + log.Error(err) + local calls = 0 + local custom = setmetatable({}, { + __tostring = function() calls = calls + 1; return "custom" end + }) + log.Info(custom, custom) + assert(calls == 2) + log.Info("a", "b", "c", "d", "e", "f", "g", "h", "i", "j", "k", "l") + log.Info() + function logHook() + log.Info("hook") + end + `), 0o600); err != nil { + t.Fatal(err) + } + if err := L.DoFile(path); err != nil { + t.Fatal(err) + } + if err := L.DoString(` + logHook() + require("xray.log").Info("anonymous") + `); err != nil { + t.Fatal(err) + } + other := filepath.Join(t.TempDir(), "other.lua") + if err := os.WriteFile(other, []byte(` + local log = require("xray.log") + log.Info("other") + logHook() + log.Info("other again") + `), 0o600); err != nil { + t.Fatal(err) + } + if err := L.DoFile(other); err != nil { + t.Fatal(err) + } + + want := []struct { + severity Severity + message string + }{ + {Severity_Debug, "[Debug] logging.lua: query: example.com"}, + {Severity_Info, "[Info] logging.lua: count=42, enabled=true, value=nil"}, + {Severity_Warning, "[Warning] logging.lua: fallback"}, + {Severity_Error, "[Error] logging.lua: failed"}, + {Severity_Error, "[Error] logging.lua: DNS failed: lookup failed: upstream timeout"}, + {Severity_Warning, "[Warning] logging.lua: lookup failed: upstream timeout"}, + {Severity_Error, "[Error] logging.lua: Lua failure"}, + {Severity_Info, "[Info] logging.lua: customcustom"}, + {Severity_Info, "[Info] logging.lua: abcdefghijkl"}, + {Severity_Info, "[Info] logging.lua: "}, + {Severity_Info, "[Info] logging.lua: hook"}, + {Severity_Info, "[Info] : anonymous"}, + {Severity_Info, "[Info] other.lua: other"}, + {Severity_Info, "[Info] logging.lua: hook"}, + {Severity_Info, "[Info] other.lua: other again"}, + } + if len(handler.messages) != len(want) { + t.Fatalf("logged %d messages, want %d", len(handler.messages), len(want)) + } + for i, expected := range want { + msg, ok := handler.messages[i].(*GeneralMessage) + if !ok { + t.Fatalf("message %d has type %T, want *GeneralMessage", i, handler.messages[i]) + } + if msg.Severity != expected.severity || msg.String() != expected.message { + t.Errorf("message %d = %q with severity %v, want %q with severity %v", i, msg.String(), msg.Severity, expected.message, expected.severity) + } + } +} + +type luaSeverityLogHandler struct { + luaLogHandler + level Severity +} + +func (h *luaSeverityLogHandler) Severity() Severity { return h.level } + +func TestLuaLogSeverity(t *testing.T) { + previous := logHandler.Load() + t.Cleanup(func() { logHandler.Store(previous) }) + L := lua.NewState() + defer L.Close() + RegisterLua(L) + for _, level := range []Severity{Severity_Unknown, Severity_Error, Severity_Warning, Severity_Info, Severity_Debug, Severity_Warning} { + t.Run(level.String(), func(t *testing.T) { + handler := &luaSeverityLogHandler{level: level} + RegisterHandler(handler) + want := []Severity{} + for _, severity := range []Severity{Severity_Error, Severity_Warning, Severity_Info, Severity_Debug} { + if severity <= level { + want = append(want, severity) + } + } + if err := L.DoString(fmt.Sprintf(` + local log = require("xray.log") + local calls = 0 + local value = setmetatable({}, { + __tostring = function() calls = calls + 1; return "message" end + }) + for _, write in ipairs({log.Error, log.Warning, log.Info, log.Debug}) do + assert(select("#", write(value)) == 0) + end + assert(calls == %d) + `, len(want))); err != nil { + t.Fatal(err) + } + if len(handler.messages) != len(want) { + t.Fatalf("logged %d messages, want %d", len(handler.messages), len(want)) + } + for i, severity := range want { + msg := handler.messages[i].(*GeneralMessage) + if msg.Severity != severity || msg.Content != ": message" { + t.Errorf("message %d = %v, want severity %v and content %q", i, msg, severity, ": message") + } + } + }) + } +} + +type luaDiscardLogHandler struct{ level Severity } + +func (luaDiscardLogHandler) Handle(Message) {} +func (h luaDiscardLogHandler) Severity() Severity { return h.level } + +func BenchmarkLuaLog(b *testing.B) { + benchmarkLuaLog(b, Severity_Debug) +} + +func BenchmarkLuaLogFiltered(b *testing.B) { + benchmarkLuaLog(b, Severity_Warning) +} + +func benchmarkLuaLog(b *testing.B, level Severity) { + previous := logHandler.Load() + b.Cleanup(func() { logHandler.Store(previous) }) + RegisterHandler(luaDiscardLogHandler{level: level}) + L := lua.NewState() + defer L.Close() + RegisterLua(L) + if err := L.DoString(`custom = setmetatable({}, {__tostring = function() return "custom" end})`); err != nil { + b.Fatal(err) + } + for _, benchmark := range []struct { + name, arguments string + }{ + {"strings", `"query: ", "example.com"`}, + {"mixed", `"count=", 42, ", enabled=", true, ", value=", nil`}, + {"many_arguments", `"a", "b", "c", "d", "e", "f", "g", "h", "i", "j", "k", "l"`}, + {"tostring", "custom"}, + } { + b.Run(benchmark.name, func(b *testing.B) { + if err := L.DoString(fmt.Sprintf(`local log = require("xray.log") +function benchmarkLog() log.Info(%s) end`, benchmark.arguments)); err != nil { + b.Fatal(err) + } + fn := L.GetGlobal("benchmarkLog") + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + if err := L.CallByParam(lua.P{Fn: fn, NRet: 0, Protect: true}); err != nil { + b.Fatal(err) + } + } + }) + } +} diff --git a/common/lua/lua.go b/common/lua/lua.go new file mode 100644 index 000000000..88fd75f14 --- /dev/null +++ b/common/lua/lua.go @@ -0,0 +1,3 @@ +// Package lua provides shared GopherLua programs, state management, and value +// conversion and validation helpers for Xray scripts. +package lua diff --git a/common/lua/luar.go b/common/lua/luar.go new file mode 100644 index 000000000..072ee663f --- /dev/null +++ b/common/lua/luar.go @@ -0,0 +1,65 @@ +package lua + +import ( + glua "github.com/yuin/gopher-lua" + luar "layeh.com/gopher-luar" +) + +// NewSlicePusher captures luar's slice metatable during state initialization. +// The returned function wraps slices without reflection or metatable lookup, +// and pushes nil for nil slices. Use it with this state or its coroutines. +func NewSlicePusher[T any](L *glua.LState) func(*glua.LState, []T) { + metatable := luar.New(L, []T{}).(*glua.LUserData).Metatable + return func(L *glua.LState, values []T) { + if values == nil { + L.Push(glua.LNil) + return + } + userdata := L.NewUserData() + userdata.Value = values + userdata.Metatable = metatable + L.Push(userdata) + } +} + +// DirectMethod handles a Lua call without luar's reflected method invocation. +// It returns the result count and whether it handled the arguments. On false, +// it must leave the stack unchanged for the original luar wrapper. +type DirectMethod func(L *glua.LState) (nresults int, handled bool) + +// PushWithDirectMethods pushes a luar userdata with typed Go method bindings. +// Handled calls bypass luar's argument conversion and reflect.Call; method lookup +// uses the methods table directly instead of luar's reflected __index handler. +// value must expose methods only. Bindings and their closures are installed once +// per Go type per LState, outside the method-call hot path. +func PushWithDirectMethods(L *glua.LState, value any, directMethods map[string]DirectMethod) { + userdata := luar.New(L, value).(*glua.LUserData) + metatable := userdata.Metatable.(*glua.LTable) + methods := metatable.RawGetString("methods").(*glua.LTable) + if metatable.RawGetString("__index") != methods { + for name, direct := range directMethods { + original := methods.RawGetString(name) + fn := L.NewFunction(func(L *glua.LState) int { + if nresults, handled := direct(L); handled { + return nresults + } + return callLuarMethod(L, original) + }) + // Keep luar's method aliases on the same direct binding. + for key, method := methods.Next(glua.LNil); key != glua.LNil; key, method = methods.Next(key) { + if method == original { + methods.RawSet(key, fn) + } + } + } + metatable.RawSetString("__index", methods) + } + L.Push(userdata) +} + +func callLuarMethod(L *glua.LState, method glua.LValue) int { + nargs := L.GetTop() + L.Insert(method, 1) + L.Call(nargs, glua.MultRet) + return L.GetTop() +} diff --git a/common/lua/luar_test.go b/common/lua/luar_test.go new file mode 100644 index 000000000..cad82b6c6 --- /dev/null +++ b/common/lua/luar_test.go @@ -0,0 +1,84 @@ +package lua + +import ( + "net" + "testing" + + glua "github.com/yuin/gopher-lua" + luar "layeh.com/gopher-luar" +) + +func TestSlicePusher(t *testing.T) { + L := glua.NewState() + defer L.Close() + push := NewSlicePusher[int](L) + values := []int{3, 5} + L.SetGlobal("getValues", L.NewFunction(func(L *glua.LState) int { + push(L, values) + return 1 + })) + if err := L.DoString(` +local values = getValues() +assert(#values == 2 and values[1] == 3 and values[2] == 5) +values[2] = 7 +local co = coroutine.create(function() + local values = getValues() + assert(#values == 2 and values[1] == 3 and values[2] == 7) + return true +end) +local ok, result = coroutine.resume(co) +assert(ok and result == true) +`); err != nil { + t.Fatal(err) + } + if values[1] != 7 { + t.Fatal("slice storage was copied") + } + push(L, nil) + if L.Get(-1) != glua.LNil { + t.Fatal("nil slice must push Lua nil") + } + L.Pop(1) + push(L, []int{}) + L.SetGlobal("empty", L.Get(-1)) + L.Pop(1) + if err := L.DoString(`assert(type(empty) == "userdata" and #empty == 0)`); err != nil { + t.Fatal(err) + } +} + +func TestSlicePusherMetatablePerState(t *testing.T) { + first := glua.NewState() + defer first.Close() + second := glua.NewState() + defer second.Close() + NewSlicePusher[int](first)(first, []int{1}) + NewSlicePusher[int](second)(second, []int{1}) + if first.Get(-1).(*glua.LUserData).Metatable == second.Get(-1).(*glua.LUserData).Metatable { + t.Fatal("independent states share a slice metatable") + } +} + +func BenchmarkSlicePusher(b *testing.B) { + L := glua.NewState() + defer L.Close() + ips := []net.IP{net.ParseIP("127.0.0.1")} + pushIPs := NewSlicePusher[net.IP](L) + for _, benchmark := range []struct { + name string + push func(*glua.LState, []net.IP) + }{ + {"bare", func(L *glua.LState, ips []net.IP) { PushUserData(L, ips) }}, + {"luar", func(L *glua.LState, ips []net.IP) { L.Push(luar.New(L, ips)) }}, + {"cached", pushIPs}, + } { + b.Run(benchmark.name, func(b *testing.B) { + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + benchmark.push(L, ips) + L.Pop(1) + } + }) + } +} diff --git a/common/lua/pool.go b/common/lua/pool.go new file mode 100644 index 000000000..ec9a58906 --- /dev/null +++ b/common/lua/pool.go @@ -0,0 +1,151 @@ +package lua + +import ( + "context" + "errors" + "sync" + "time" + + glua "github.com/yuin/gopher-lua" +) + +const maxIdleStates = 16 + +// Pool lends each state to one caller at a time. It grows on contention and +// keeps up to maxIdleStates idle states until Close. Acquire/Release callers +// decide reusability; WithState uses its callback's error. +type Pool struct { + ctx context.Context + cancel context.CancelFunc + timeout time.Duration + + factory LStateFactory + idle []*glua.LState + top int + + mu sync.Mutex + active sync.WaitGroup + closed bool +} + +// NewPool tests the factory by creating one state during initialization. +func NewPool(ctx context.Context, timeout time.Duration, factory LStateFactory) (*Pool, error) { + if timeout <= 0 { + return nil, errors.New("Lua pool timeout must be positive") + } + + poolCtx, cancel := context.WithCancel(ctx) + + state, err := factory(poolCtx) + if err != nil { + cancel() + return nil, err + } + + return &Pool{ctx: poolCtx, cancel: cancel, timeout: timeout, factory: factory, idle: []*glua.LState{state}, top: state.GetTop()}, nil +} + +// Acquire returns an initialized exclusive state, growing the pool if necessary. +// ctx is passed to the factory for state creation; nil uses the pool context. +func (p *Pool) Acquire(ctx context.Context) (*glua.LState, error) { + p.mu.Lock() + if p.closed { + p.mu.Unlock() + return nil, errors.New("Lua pool is closed") + } + if err := p.ctx.Err(); err != nil { + p.mu.Unlock() + return nil, err + } + if ctx == nil { + ctx = p.ctx + } else if err := ctx.Err(); err != nil { + p.mu.Unlock() + return nil, err + } + + p.active.Add(1) + + n := len(p.idle) + if n != 0 { + state := p.idle[n-1] + p.idle[n-1] = nil + 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(ctx) + if err != nil { + p.active.Done() + return nil, err + } + + return state, nil +} + +// WithState runs work on an exclusive state and releases it afterward. +// Nil ctx and zero timeout use pool defaults. The timeout starts after acquisition. +func (p *Pool) WithState(ctx context.Context, timeout time.Duration, work func(*glua.LState) error) error { + state, err := p.Acquire(ctx) + if err != nil { + return err + } + if ctx == nil { + ctx = p.ctx + } + if timeout == 0 { + timeout = p.timeout + } + ctx, cancel := context.WithTimeout(ctx, timeout) + state.SetContext(ctx) + reusable := false + defer func() { + cancel() + p.Release(state, reusable) + }() + err = work(state) + reusable = err == nil + return err +} + +// Release resets a state for reuse or closes it. +func (p *Pool) Release(state *glua.LState, reusable bool) { + if reusable { + state.RemoveContext() + state.SetTop(p.top) + 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 the pool context, 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/pool_test.go b/common/lua/pool_test.go new file mode 100644 index 000000000..49bb25bbd --- /dev/null +++ b/common/lua/pool_test.go @@ -0,0 +1,466 @@ +package lua + +import ( + "context" + "errors" + "testing" + "time" + + glua "github.com/yuin/gopher-lua" +) + +func newTestPool(t testing.TB, ctx context.Context, timeout time.Duration, factory LStateFactory) *Pool { + t.Helper() + pool, err := NewPool(ctx, timeout, factory) + if err != nil { + t.Fatal(err) + } + t.Cleanup(pool.Close) + return pool +} + +func assertPoolCloseBlocked(t *testing.T, done <-chan struct{}) { + t.Helper() + select { + case <-done: + t.Fatal("Close returned while work was still active") + case <-time.After(20 * time.Millisecond): + } +} + +func TestPoolTimeoutValidation(t *testing.T) { + for _, tc := range []struct { + name string + timeout time.Duration + wantErr bool + }{ + {"zero", 0, true}, + {"negative", -time.Nanosecond, true}, + {"positive", time.Nanosecond, false}, + } { + t.Run(tc.name, func(t *testing.T) { + called := false + pool, err := NewPool(context.Background(), tc.timeout, func(context.Context) (*glua.LState, error) { + called = true + return glua.NewState(), nil + }) + if pool != nil { + t.Cleanup(pool.Close) + } + if (err != nil) != tc.wantErr { + t.Fatalf("NewPool error = %v, want error %t", err, tc.wantErr) + } + if tc.wantErr && (pool != nil || called) { + t.Fatal("invalid timeout created a pool or called the factory") + } + }) + } +} + +func TestPoolFactoryFailure(t *testing.T) { + failure := errors.New("factory failed") + _, err := NewPool(context.Background(), time.Second, func(context.Context) (*glua.LState, error) { + return nil, failure + }) + if !errors.Is(err, failure) { + t.Fatalf("NewPool error = %v, want original factory error", err) + } + + calls := 0 + pool := newTestPool(t, context.Background(), time.Second, func(context.Context) (*glua.LState, error) { + calls++ + if calls == 1 { + return glua.NewState(), nil + } + return nil, failure + }) + state, err := pool.Acquire(nil) + if err != nil { + t.Fatal(err) + } + defer pool.Release(state, true) + err = pool.WithState(nil, 0, func(*glua.LState) error { + t.Error("work ran after factory failure") + return nil + }) + if !errors.Is(err, failure) { + t.Fatalf("WithState error = %v, want original factory error", err) + } +} + +func TestPoolReusesStatesAndLimitsIdle(t *testing.T) { + created := 0 + pool := newTestPool(t, context.Background(), time.Second, func(context.Context) (*glua.LState, error) { + created++ + return glua.NewState(), nil + }) + var borrowed []*glua.LState + defer func() { + for _, state := range borrowed { + pool.Release(state, false) + } + }() + for range maxIdleStates + 3 { + state, err := pool.Acquire(nil) + if err != nil { + t.Fatal(err) + } + borrowed = append(borrowed, state) + state.SetContext(context.Background()) + } + states := borrowed + for _, state := range states { + pool.Release(state, true) + } + borrowed = nil + open := 0 + for _, state := range states { + if !state.IsClosed() { + if state.Context() != nil { + t.Fatal("Release left a context on a reusable state") + } + open++ + } + } + if open != maxIdleStates { + t.Fatalf("retained %d states, want %d", open, maxIdleStates) + } + if err := pool.WithState(nil, 0, func(*glua.LState) error { return nil }); err != nil { + t.Fatal(err) + } + if created != len(states) { + t.Fatalf("created %d states, want %d", created, len(states)) + } + pool.Close() + for _, state := range states { + if !state.IsClosed() { + t.Fatal("Close left an idle state open") + } + } +} + +func TestPoolWithStateOptions(t *testing.T) { + key := struct{}{} + parent := context.WithValue(context.Background(), key, "pool") + caller := context.WithValue(context.Background(), key, "caller") + pool := newTestPool(t, parent, time.Second, func(context.Context) (*glua.LState, error) { + return glua.NewState(), nil + }) + for _, tc := range []struct { + name string + ctx context.Context + timeout time.Duration + wantValue string + wantTimeout time.Duration + }{ + {"defaults", nil, 0, "pool", time.Second}, + {"context", caller, 0, "caller", time.Second}, + {"timeout", nil, 2 * time.Second, "pool", 2 * time.Second}, + {"both", caller, 2 * time.Second, "caller", 2 * time.Second}, + } { + t.Run(tc.name, func(t *testing.T) { + started := time.Now() + err := pool.WithState(tc.ctx, tc.timeout, func(L *glua.LState) error { + ctx := L.Context() + if ctx.Value(key) != tc.wantValue { + t.Errorf("context value = %v, want %q", ctx.Value(key), tc.wantValue) + } + deadline, ok := ctx.Deadline() + if !ok || deadline.Before(started.Add(tc.wantTimeout)) || deadline.After(time.Now().Add(tc.wantTimeout)) { + t.Errorf("deadline = %v, want timeout %v", deadline, tc.wantTimeout) + } + return nil + }) + if err != nil { + t.Fatal(err) + } + }) + } +} + +func TestPoolFactoryContext(t *testing.T) { + caller, cancel := context.WithTimeout(context.Background(), time.Minute) + defer cancel() + for _, tc := range []struct { + name string + ctx context.Context + }{ + {"default", nil}, + {"caller", caller}, + } { + t.Run(tc.name, func(t *testing.T) { + var contexts []context.Context + pool := newTestPool(t, context.Background(), time.Second, func(ctx context.Context) (*glua.LState, error) { + contexts = append(contexts, ctx) + return glua.NewState(), nil + }) + state, err := pool.Acquire(nil) + if err != nil { + t.Fatal(err) + } + defer pool.Release(state, true) + if err := pool.WithState(tc.ctx, 2*time.Second, func(*glua.LState) error { return nil }); err != nil { + t.Fatal(err) + } + want := tc.ctx + if want == nil { + want = pool.ctx + } + if len(contexts) != 2 || contexts[0] != pool.ctx || contexts[1] != want { + t.Fatal("factory did not receive the initialization and acquisition contexts unchanged") + } + }) + } +} + +func TestPoolWithStateLifecycle(t *testing.T) { + failure := errors.New("work failed") + for _, tc := range []struct { + name string + work func(*glua.LState, context.CancelFunc) error + reusable bool + wantPanic bool + wantErr error + }{ + {"success", func(*glua.LState, context.CancelFunc) error { return nil }, true, false, nil}, + {"canceled success", func(_ *glua.LState, cancel context.CancelFunc) error { + cancel() + return nil + }, true, false, nil}, + {"error", func(*glua.LState, context.CancelFunc) error { return failure }, false, false, failure}, + {"timeout", func(L *glua.LState, _ context.CancelFunc) error { return L.DoString("while true do end") }, false, false, nil}, + {"panic", func(*glua.LState, context.CancelFunc) error { panic(failure) }, false, true, nil}, + } { + t.Run(tc.name, func(t *testing.T) { + pool := newTestPool(t, context.Background(), 10*time.Millisecond, func(context.Context) (*glua.LState, error) { + state := glua.NewState() + state.Push(glua.LTrue) + return state, nil + }) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + var state *glua.LState + var workCtx context.Context + var recovered any + err := func() (err error) { + defer func() { recovered = recover() }() + return pool.WithState(ctx, 0, func(L *glua.LState) error { + state, workCtx = L, L.Context() + L.Push(glua.LFalse) + return tc.work(L, cancel) + }) + }() + if tc.wantPanic { + if recovered != failure { + t.Fatalf("panic = %v, want original panic", recovered) + } + } else { + if recovered != nil || (err == nil) != tc.reusable { + t.Fatalf("WithState error = %v, panic = %v", err, recovered) + } + if tc.wantErr != nil && !errors.Is(err, tc.wantErr) { + t.Fatalf("WithState error = %v, want %v", err, tc.wantErr) + } + } + if workCtx.Err() == nil { + t.Fatal("WithState did not cancel the execution context") + } + if closed := state.IsClosed(); closed == tc.reusable { + t.Fatalf("state closed = %t, want %t", closed, !tc.reusable) + } + if tc.reusable && (state.Context() != nil || state.GetTop() != 1 || state.Get(1) != glua.LTrue) { + t.Fatal("WithState did not reset the state for reuse") + } + if err := pool.WithState(nil, 0, func(L *glua.LState) error { + if (L == state) != tc.reusable { + t.Error("unexpected state reuse") + } + return nil + }); err != nil { + t.Fatal(err) + } + }) + } +} + +func TestPoolClose(t *testing.T) { + pool := newTestPool(t, context.Background(), time.Second, func(context.Context) (*glua.LState, error) { + return glua.NewState(), nil + }) + finishCtx, finish := context.WithCancel(context.Background()) + t.Cleanup(finish) + started, done := make(chan *glua.LState, 1), make(chan error, 1) + var workCtx context.Context + go func() { + done <- pool.WithState(nil, 0, func(L *glua.LState) error { + workCtx = L.Context() + started <- L + <-finishCtx.Done() + return nil + }) + }() + var state *glua.LState + select { + case state = <-started: + case <-time.After(time.Second): + t.Fatal("WithState did not start") + } + closed := make(chan struct{}) + go func() { + pool.Close() + close(closed) + }() + select { + case <-workCtx.Done(): + case <-time.After(time.Second): + t.Fatal("Close did not cancel work using the pool context") + } + if !errors.Is(workCtx.Err(), context.Canceled) { + t.Fatalf("work context error = %v, want context.Canceled", workCtx.Err()) + } + assertPoolCloseBlocked(t, closed) + finish() + select { + case err := <-done: + if err != nil { + t.Fatalf("successful work returned an error: %v", err) + } + case <-time.After(time.Second): + t.Fatal("WithState did not finish") + } + select { + case <-closed: + case <-time.After(time.Second): + t.Fatal("Close did not finish after WithState") + } + if !state.IsClosed() { + t.Fatal("Release returned a state to a closed pool") + } + if state, err := pool.Acquire(nil); state != nil || err == nil || errors.Is(err, context.Canceled) { + t.Fatalf("Acquire after Close = %v, %v; want closed pool error", state, err) + } + pool.Close() +} + +func TestPoolCloseWaitsForFactory(t *testing.T) { + finishCtx, finish := context.WithCancel(context.Background()) + started, canceled := make(chan struct{}), make(chan struct{}) + first := true + pool := newTestPool(t, context.Background(), time.Second, func(ctx context.Context) (*glua.LState, error) { + if first { + first = false + return glua.NewState(), nil + } + close(started) + <-ctx.Done() + close(canceled) + <-finishCtx.Done() + return nil, ctx.Err() + }) + t.Cleanup(finish) + state, err := pool.Acquire(nil) + if err != nil { + t.Fatal(err) + } + pool.Release(state, false) + acquireDone := make(chan error, 1) + go func() { + _, err := pool.Acquire(nil) + acquireDone <- err + }() + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("state creation did not start") + } + closed := make(chan struct{}) + go func() { + pool.Close() + close(closed) + }() + select { + case <-canceled: + case <-time.After(time.Second): + t.Fatal("Close did not cancel state creation") + } + assertPoolCloseBlocked(t, closed) + finish() + select { + case err := <-acquireDone: + if !errors.Is(err, context.Canceled) { + t.Fatalf("Acquire error = %v, want context.Canceled", err) + } + case <-time.After(time.Second): + t.Fatal("state creation did not finish") + } + select { + case <-closed: + case <-time.After(time.Second): + t.Fatal("Close did not finish after state creation") + } +} + +func TestPoolCloseWaitsForCallerContext(t *testing.T) { + pool := newTestPool(t, context.Background(), time.Minute, func(context.Context) (*glua.LState, error) { + return glua.NewState(), nil + }) + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + started, done := make(chan context.Context, 1), make(chan error, 1) + go func() { + done <- pool.WithState(ctx, 0, func(L *glua.LState) error { + started <- L.Context() + <-L.Context().Done() + return L.Context().Err() + }) + }() + var workCtx context.Context + select { + case workCtx = <-started: + case <-time.After(time.Second): + t.Fatal("WithState did not start") + } + closed := make(chan struct{}) + go func() { + pool.Close() + close(closed) + }() + select { + case <-pool.ctx.Done(): + case <-time.After(time.Second): + t.Fatal("Close did not cancel the pool context") + } + assertPoolCloseBlocked(t, closed) + if workCtx.Err() != nil || ctx.Err() != nil { + t.Fatal("Close canceled the caller's execution context") + } + cancel() + select { + case err := <-done: + if !errors.Is(err, context.Canceled) { + t.Fatalf("WithState error = %v, want context.Canceled", err) + } + case <-time.After(time.Second): + t.Fatal("WithState did not stop after caller cancellation") + } + select { + case <-closed: + case <-time.After(time.Second): + t.Fatal("Close did not finish after WithState") + } +} + +func BenchmarkPoolAcquireRelease(b *testing.B) { + pool := newTestPool(b, context.Background(), time.Second, func(context.Context) (*glua.LState, error) { + return glua.NewState(), nil + }) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + state, err := pool.Acquire(nil) + if err != nil { + b.Fatal(err) + } + pool.Release(state, true) + } +} diff --git a/common/lua/program.go b/common/lua/program.go new file mode 100644 index 000000000..38eb225fc --- /dev/null +++ b/common/lua/program.go @@ -0,0 +1,77 @@ +package lua + +import ( + "bufio" + "context" + "os" + "time" + + 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 +} + +// LStateFactory returns a fully initialized state or nil and an error. +// Implementations must close partial states on failure; callers own successful states. +type LStateFactory func(context.Context) (*glua.LState, error) + +// 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 state, runs register, executes the program under ctx, and +// runs validate. It removes the initialization context before returning a state +// owned by the caller. +func (p *Program) NewState(ctx context.Context, register func(*glua.LState), validate func(*glua.LState) error) (*glua.LState, error) { + L := glua.NewState() + valid := false + defer func() { + if !valid { + L.Close() + } + }() + L.SetContext(ctx) + defer L.RemoveContext() + if register != nil { + register(L) + } + L.Push(L.NewFunctionFromProto(p.proto)) + // Execute the Lua script's top level. + if err := L.PCall(0, 0, nil); err != nil { + return nil, err + } + if validate != nil { + if err := validate(L); err != nil { + return nil, err + } + } + valid = true + return L, nil +} + +// NewStateFactory returns a factory that gives each state an initialization timeout. +func (p *Program) NewStateFactory(initTimeout time.Duration, register func(*glua.LState), validate func(*glua.LState) error) LStateFactory { + return func(ctx context.Context) (*glua.LState, error) { + initCtx, cancel := context.WithTimeout(ctx, initTimeout) + defer cancel() + return p.NewState(initCtx, register, validate) + } +} diff --git a/common/lua/program_test.go b/common/lua/program_test.go new file mode 100644 index 000000000..f4593f7e6 --- /dev/null +++ b/common/lua/program_test.go @@ -0,0 +1,76 @@ +package lua + +import ( + "context" + "errors" + "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, nil) + if err != nil { + t.Fatal(err) + } + defer first.Close() + first.SetGlobal("value", glua.LNumber(42)) + second, err := program.NewState(context.Background(), nil, 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, 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) + } +} + +func TestNewStateClosesFailedValidation(t *testing.T) { + path := filepath.Join(t.TempDir(), "state.lua") + if err := os.WriteFile(path, []byte("value = 1"), 0o600); err != nil { + t.Fatal(err) + } + program, err := CompileFile(path) + if err != nil { + t.Fatal(err) + } + wantErr := errors.New("invalid script") + var checked *glua.LState + L, err := program.NewState(context.Background(), nil, func(L *glua.LState) error { + checked = L + return wantErr + }) + if L != nil || !errors.Is(err, wantErr) || checked == nil || !checked.IsClosed() { + t.Fatalf("state = %v, error = %v, checked state closed = %t", L, err, checked != nil && checked.IsClosed()) + } +} diff --git a/common/lua/utils.go b/common/lua/utils.go new file mode 100644 index 000000000..f321e6572 --- /dev/null +++ b/common/lua/utils.go @@ -0,0 +1,95 @@ +package lua + +import ( + "math" + + "github.com/xtls/xray-core/common/errors" + glua "github.com/yuin/gopher-lua" +) + +type number interface { + ~int | ~int8 | ~int16 | ~int32 | ~int64 | + ~uint | ~uint8 | ~uint16 | ~uint32 | ~uint64 | ~uintptr | + ~float32 | ~float64 +} + +// PushNumber converts a Go number to a Lua number and pushes it. +func PushNumber[T number](L *glua.LState, value T) { + L.Push(glua.LNumber(value)) +} + +// PushString converts a Go string to a Lua string and pushes it. +func PushString(L *glua.LState, value string) { + L.Push(glua.LString(value)) +} + +// PushNil pushes Lua nil. +func PushNil(L *glua.LState) { + L.Push(glua.LNil) +} + +// PushUserData pushes a native Go value without copying it. +func PushUserData(L *glua.LState, value any) { + ud := L.NewUserData() + ud.Value = value + L.Push(ud) +} + +// PushError pushes nil or the original Go error as userdata. +func PushError(L *glua.LState, err error) { + if err == nil { + L.Push(glua.LNil) + return + } + PushUserData(L, err) +} + +// ReadUserData reads a native Go value of type T without copying it. +// Other Lua values or userdata containing a different type return invalidMessage. +func ReadUserData[T any](value glua.LValue, invalidMessage string) (T, error) { + if ud, ok := value.(*glua.LUserData); ok { + if result, ok := ud.Value.(T); ok { + return result, nil + } + } + var zero T + return zero, errors.New(invalidMessage) +} + +// ReadError accepts nil, a native Go error, or a Lua string. +// Native errors retain their identity; other values return invalidMessage. +func ReadError(value glua.LValue, invalidMessage string) error { + if value == glua.LNil { + return nil + } + if ud, ok := value.(*glua.LUserData); ok { + if err, ok := ud.Value.(error); ok { + return err + } + } + if message, ok := value.(glua.LString); ok { + return errors.New(string(message)) + } + return errors.New(invalidMessage) +} + +// ReadUint32 accepts only integral Lua numbers in the uint32 range. +func ReadUint32(value glua.LValue, invalidMessage string) (uint32, error) { + number, ok := value.(glua.LNumber) + if !ok || number < 0 || number > math.MaxUint32 || math.Trunc(float64(number)) != float64(number) { + return 0, errors.New(invalidMessage) + } + return uint32(number), nil +} + +// ReadOptionalString accepts a Lua string or nil, which becomes an empty string. +// It does not coerce other values to strings. +func ReadOptionalString(value glua.LValue, invalidMessage string) (string, error) { + if value == glua.LNil { + return "", nil + } + if result, ok := value.(glua.LString); ok { + return string(result), nil + } + return "", errors.New(invalidMessage) +} diff --git a/common/lua/utils_test.go b/common/lua/utils_test.go new file mode 100644 index 000000000..3a6f6a97a --- /dev/null +++ b/common/lua/utils_test.go @@ -0,0 +1,121 @@ +package lua + +import ( + "errors" + "math" + "strings" + "testing" + + glua "github.com/yuin/gopher-lua" +) + +func TestReadUint32(t *testing.T) { + for _, tc := range []struct { + name string + value glua.LValue + want uint32 + wantErr bool + }{ + {name: "zero", value: glua.LNumber(0)}, + {name: "integer", value: glua.LNumber(45), want: 45}, + {name: "maximum", value: glua.LNumber(math.MaxUint32), want: math.MaxUint32}, + {name: "fraction", value: glua.LNumber(1.5), wantErr: true}, + {name: "negative", value: glua.LNumber(-1), wantErr: true}, + {name: "overflow", value: glua.LNumber(math.MaxUint32 + 1), wantErr: true}, + {name: "NaN", value: glua.LNumber(math.NaN()), wantErr: true}, + {name: "positive infinity", value: glua.LNumber(math.Inf(1)), wantErr: true}, + {name: "negative infinity", value: glua.LNumber(math.Inf(-1)), wantErr: true}, + {name: "nil", value: glua.LNil, wantErr: true}, + {name: "numeric string", value: glua.LString("45"), wantErr: true}, + {name: "boolean", value: glua.LTrue, wantErr: true}, + } { + t.Run(tc.name, func(t *testing.T) { + got, err := ReadUint32(tc.value, "invalid number") + if got != tc.want || (err != nil) != tc.wantErr { + t.Fatalf("ReadUint32() = %d, %v; want %d, error %t", got, err, tc.want, tc.wantErr) + } + if err != nil && !strings.Contains(err.Error(), "invalid number") { + t.Fatalf("error = %v, want invalid number", err) + } + }) + } +} + +func TestReadOptionalString(t *testing.T) { + for _, tc := range []struct { + name string + value glua.LValue + want string + wantErr bool + }{ + {name: "nil", value: glua.LNil}, + {name: "empty", value: glua.LString("")}, + {name: "string", value: glua.LString("out"), want: "out"}, + {name: "number", value: glua.LNumber(1), wantErr: true}, + {name: "boolean", value: glua.LFalse, wantErr: true}, + } { + t.Run(tc.name, func(t *testing.T) { + got, err := ReadOptionalString(tc.value, "invalid string") + if got != tc.want || (err != nil) != tc.wantErr { + t.Fatalf("ReadOptionalString() = %q, %v; want %q, error %t", got, err, tc.want, tc.wantErr) + } + if err != nil && !strings.Contains(err.Error(), "invalid string") { + t.Fatalf("error = %v, want invalid string", err) + } + }) + } +} + +func TestUserDataRoundTrip(t *testing.T) { + L := glua.NewState() + defer L.Close() + want := []int{1, 2} + PushUserData(L, want) + if L.GetTop() != 1 { + t.Fatalf("stack top = %d, want 1", L.GetTop()) + } + got, err := ReadUserData[[]int](L.Get(-1), "invalid userdata") + if err != nil || len(got) != len(want) || &got[0] != &want[0] { + t.Fatalf("userdata = %v, %v; want original slice", got, err) + } + PushUserData(L, []int(nil)) + if got, err := ReadUserData[[]int](L.Get(-1), "invalid userdata"); err != nil || got != nil { + t.Fatalf("nil slice userdata = %v, %v", got, err) + } + for _, value := range []glua.LValue{glua.LNil, glua.LString("1"), L.NewTable(), L.Get(1)} { + if got, err := ReadUserData[int](value, "invalid userdata"); got != 0 || err == nil || !strings.Contains(err.Error(), "invalid userdata") { + t.Fatalf("ReadUserData(%v) = %d, %v; want invalid userdata", value, got, err) + } + } +} + +func TestErrorRoundTrip(t *testing.T) { + L := glua.NewState() + defer L.Close() + want := errors.New("upstream failed") + for _, err := range []error{nil, want} { + PushError(L, err) + if L.GetTop() != 1 { + t.Fatalf("stack top = %d, want 1", L.GetTop()) + } + if err == nil && L.Get(-1) != glua.LNil { + t.Fatalf("nil error pushed as %v", L.Get(-1)) + } + if got := ReadError(L.Get(-1), "invalid error"); got != err { + t.Fatalf("ReadError() = %v, want original error %v", got, err) + } + L.Pop(1) + } + for _, message := range []string{"script failed", ""} { + if err := ReadError(glua.LString(message), "invalid error"); err == nil || !strings.Contains(err.Error(), message) { + t.Fatalf("string error = %v, want %q", err, message) + } + } + wrong := L.NewUserData() + wrong.Value = "not a native error" + for _, value := range []glua.LValue{glua.LTrue, glua.LNumber(1), L.NewTable(), wrong, L.NewUserData()} { + if err := ReadError(value, "invalid error"); err == nil || !strings.Contains(err.Error(), "invalid error") { + t.Fatalf("ReadError(%v) = %v, want invalid error", value, err) + } + } +} diff --git a/common/platform/platform.go b/common/platform/platform.go index b60b8bd22..4c5cf2a0f 100644 --- a/common/platform/platform.go +++ b/common/platform/platform.go @@ -1,6 +1,8 @@ package platform // import "github.com/xtls/xray-core/common/platform" import ( + "errors" + "fmt" "os" "path/filepath" "strconv" @@ -90,3 +92,49 @@ func GetConfDirPath() string { configPath := NewEnvFlag(ConfdirLocation).GetValue(func() string { return "" }) return configPath } + +// ResolveLuaFile finds a local Lua script and returns its absolute path. +// Relative paths: XRAY_LOCATION_CONFDIR > XRAY_LOCATION_CONFIG > working dir > executable dir. +func ResolveLuaFile(path string) (string, error) { + if path == "" { + return "", errors.New("Lua file path is empty") + } + paths := []string{path} + if !filepath.IsAbs(path) { + paths = nil + for _, dir := range []string{ + GetConfDirPath(), + NewEnvFlag(ConfigLocation).GetValue(func() string { return "" }), + ".", + getExecutableDir(), + } { + if dir != "" { + paths = append(paths, filepath.Join(dir, path)) + } + } + } + return resolveFile(paths) +} + +func resolveFile(paths []string) (string, error) { + var tried []string + for _, path := range paths { + path, err := filepath.Abs(path) + if err != nil { + return "", fmt.Errorf("failed to resolve file path: %w", err) + } + tried = append(tried, path) + info, err := os.Stat(path) + if errors.Is(err, os.ErrNotExist) { + continue + } + if err != nil { + return "", fmt.Errorf("failed to inspect file %q: %w", path, err) + } + if !info.Mode().IsRegular() { + return "", fmt.Errorf("file is not a regular file: %s", path) + } + return path, nil + } + return "", fmt.Errorf("file not found; tried %q: %w", tried, os.ErrNotExist) +} diff --git a/common/platform/platform_test.go b/common/platform/platform_test.go index 854c397f6..95d5b6545 100644 --- a/common/platform/platform_test.go +++ b/common/platform/platform_test.go @@ -1,6 +1,7 @@ package platform_test import ( + "errors" "os" "path/filepath" "runtime" @@ -64,3 +65,53 @@ func TestGetAssetLocation(t *testing.T) { } } } + +func TestResolveLuaFile(t *testing.T) { + workingDir := t.TempDir() + t.Chdir(workingDir) + executable, err := os.Executable() + common.Must(err) + file, err := os.CreateTemp(filepath.Dir(executable), "lua-*.lua") + common.Must(err) + common.Must(file.Close()) + defer os.Remove(file.Name()) + + name := filepath.Base(file.Name()) + paths := []string{ + filepath.Join(t.TempDir(), name), + filepath.Join(t.TempDir(), name), + filepath.Join(workingDir, name), + file.Name(), + } + t.Setenv(ConfdirLocation, filepath.Dir(paths[0])) + t.Setenv(ConfigLocation, filepath.Dir(paths[1])) + for _, path := range paths[:3] { + common.Must(os.WriteFile(path, nil, 0o600)) + } + if got, err := ResolveLuaFile(paths[2]); err != nil || got != paths[2] { + t.Fatalf("absolute path = %q, %v; want %q", got, err, paths[2]) + } + for i, want := range paths { + if i == 2 { + t.Setenv(ConfdirLocation, "") + t.Setenv(ConfigLocation, "") + } + if got, err := ResolveLuaFile(name); err != nil || got != want { + t.Fatalf("resolved path = %q, %v; want %q", got, err, want) + } + common.Must(os.Remove(want)) + } + if _, err := ResolveLuaFile(name); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("missing file error = %v", err) + } + + t.Setenv(ConfdirLocation, filepath.Dir(paths[0])) + t.Setenv(ConfigLocation, filepath.Dir(paths[1])) + common.Must(os.Mkdir(paths[0], 0o700)) + common.Must(os.WriteFile(paths[1], nil, 0o600)) + for _, path := range []string{"", name, filepath.Join(t.TempDir(), name)} { + if _, err := ResolveLuaFile(path); err == nil { + t.Fatalf("accepted invalid path %q", path) + } + } +} diff --git a/go.mod b/go.mod index 6ddfa871c..b9abcee36 100644 --- a/go.mod +++ b/go.mod @@ -21,6 +21,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 @@ -34,6 +35,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 f72773596..c7f69432d 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= @@ -81,6 +84,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= @@ -107,6 +113,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= @@ -157,6 +164,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..07b482742 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,14 @@ func (c *DNSConfig) Build() (*dns.Config, error) { QueryStrategy: resolveQueryStrategy(c.QueryStrategy), } + if c.Script != "" { + path, err := platform.ResolveLuaFile(c.Script) + if err != nil { + return nil, errors.New("failed to resolve DNS script: ", c.Script).Base(err) + } + config.Script = path + } + if c.ClientIP != nil { if !c.ClientIP.Family().IsIP() { return nil, errors.New("not an IP address:", c.ClientIP.String()) diff --git a/infra/conf/dns_test.go b/infra/conf/dns_test.go index 278f34c38..5657f0c70 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(domain, ipv4, ipv6, fake) 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) + } +} diff --git a/infra/conf/router.go b/infra/conf/router.go index 3f85ce8fa..ba5cfeef2 100644 --- a/infra/conf/router.go +++ b/infra/conf/router.go @@ -7,6 +7,7 @@ import ( "github.com/xtls/xray-core/app/router" "github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/geodata" + "github.com/xtls/xray-core/common/platform" "github.com/xtls/xray-core/common/serial" "google.golang.org/protobuf/proto" @@ -72,6 +73,7 @@ type RouterConfig struct { RuleList []json.RawMessage `json:"rules"` DomainStrategy *string `json:"domainStrategy"` Balancers []*BalancingRule `json:"balancers"` + Script string `json:"script"` } func (c *RouterConfig) getDomainStrategy() router.Config_DomainStrategy { @@ -92,6 +94,15 @@ func (c *RouterConfig) getDomainStrategy() router.Config_DomainStrategy { func (c *RouterConfig) Build() (*router.Config, error) { config := new(router.Config) + + if c.Script != "" { + path, err := platform.ResolveLuaFile(c.Script) + if err != nil { + return nil, errors.New("failed to resolve routing script").Base(err) + } + config.Script = path + } + config.DomainStrategy = c.getDomainStrategy() var rawRuleList []json.RawMessage diff --git a/infra/conf/router_test.go b/infra/conf/router_test.go index 130cf4f78..04e4bcd83 100644 --- a/infra/conf/router_test.go +++ b/infra/conf/router_test.go @@ -2,6 +2,8 @@ package conf_test import ( "encoding/json" + "os" + "path/filepath" "testing" "time" _ "unsafe" @@ -236,3 +238,39 @@ func TestRouterConfig(t *testing.T) { }, }) } + +func TestRouterScriptConfig(t *testing.T) { + dir := t.TempDir() + t.Setenv("xray.location.confdir", dir) + path := filepath.Join(dir, "route.lua") + if err := os.WriteFile(path, []byte("function HandleRoute() end"), 0o600); err != nil { + t.Fatal(err) + } + + for _, tc := range []struct { + name string + script string + wantError bool + }{ + {"relative", "route.lua", false}, + {"absolute", path, false}, + {"missing", "missing.lua", true}, + {"directory", dir, true}, + } { + t.Run(tc.name, func(t *testing.T) { + built, err := (&RouterConfig{Script: tc.script}).Build() + if tc.wantError { + if err == nil { + t.Fatal("Build accepted invalid script path") + } + return + } + if err != nil { + t.Fatal(err) + } + if built.Script != path { + t.Fatalf("script path = %q, want %q", built.Script, path) + } + }) + } +}