Compare commits

..
34 Commits
Author SHA1 Message Date
Meo597 d86ad7c501 test: rework and calibrate Lua DNS and routing benchmarks 2026-10-05 15:16:01 +08:00
Meo597 8d80e9924f lua: unify IP slice access and cache slice metatables 2026-10-05 09:06:20 +08:00
Meo597 103f5cbfdb pool: clear idle state references on acquire
Prevent the idle slice's backing array from retaining discarded states.
2026-10-05 08:37:12 +08:00
Meo597 52be8f8170 router: preserve Lua states after recoverable route errors & refactor 2026-10-05 06:38:24 +08:00
Meo597 248cb666f3 dns: preserve Lua states after recoverable query errors & refactor 2026-10-05 06:15:49 +08:00
Meo597 54610a39a5 lua: preallocate tables with appropriate capacities 2026-10-05 04:29:01 +08:00
Meo597 37d9894dde geodata: reduce Lua matcher build allocations 2026-10-05 04:11:03 +08:00
Meo597 e73fb0d500 geodata: skip reflection for matcher calls 2026-10-05 04:01:08 +08:00
Meo597 d74b8b1ba5 log: skip filtered Lua messages early 2026-10-05 01:10:06 +08:00
Meo597 edb8e7477e test(log): fix Lua test compilation after merging main 2026-10-05 00:51:56 +08:00
Meo597 6cd6c61578 Merge commit 'b26a91de4f3294e26a0ad0a970b81a386a41f789' into lua 2026-10-05 00:46:32 +08:00
Meo597 db2dc8840a log: cut log allocations and cache filename prefixes 2026-10-05 00:45:25 +08:00
Meo597 7ab6930f27 refactor: move push helpers to common and use xlua as the import alias 2026-10-04 07:12:54 +08:00
Meo597 73fb3e8f4a refactor: extract shared result validation and userdata helpers 2026-10-04 06:29:00 +08:00
Meo597 745526f14c refactor(dns): make server-specific Lua registration private 2026-10-04 05:41:34 +08:00
Meo597 399563b6d9 refactor: move execution context and timeout management into pool 2026-10-04 05:22:47 +08:00
Meo597 e38794ed88 lua: reset pooled state stack to prevent memory leaks 2026-10-04 01:36:48 +08:00
Meo597 2610e57ecf lua: refactor to simplify state creation and pooled script execution 2026-10-03 21:57:54 +08:00
Meo597 5afe260f10 router: lowercase target domain passed to Lua hook 2026-10-01 22:05:09 +08:00
Meo597 5d1d8200d9 dns: unify the Lua API for local and configured DNS
Expose localdns.Client in xray.dns.Servers with the ID "localhost",
so Lua scripts can use the same server API for both clients.
2026-10-01 21:58:09 +08:00
Meo597 2440f53cdd feat(router): add Lua scripting for routing script 2026-10-01 18:25:53 +08:00
RPRX b26a91de4f Xray-core v26.9.30
Sponsor & Donation & NFTs: https://github.com/XTLS/Xray-core/issues/3668
Project X Channel: https://t.me/projectXtls

Announcement of NFTs by Project X: https://github.com/XTLS/Xray-core/discussions/3633
Project X NFT: https://opensea.io/assets/ethereum/0x5ee362866001613093361eb8569d59c4141b76d1/1

VLESS Post-Quantum Encryption: https://github.com/XTLS/Xray-core/pull/5067
VLESS NFT: https://opensea.io/collection/vless

XHTTP: Beyond REALITY: https://github.com/XTLS/Xray-core/discussions/4113
REALITY NFT: https://opensea.io/assets/ethereum/0x5ee362866001613093361eb8569d59c4141b76d1/2
2026-09-30 07:40:04 +00:00
patternihaandClaude Opus 5.5 1f304916bd TUN inbound: Add autoSystemWfpBlockLeak on Windows (blocks "dns" and "misconfigtun" IPv4/IPv6 traffic leaks outside the TUN); Rename autoSystemDNS to autoSystemDnsToGateway on Linux (and change some behaviors) (#6853)
https://github.com/XTLS/Xray-core/pull/6853#issuecomment-5899791359
https://github.com/XTLS/Xray-core/pull/6853#issuecomment-5901287980
https://github.com/XTLS/Xray-core/pull/6853#issuecomment-5903680113
https://github.com/XTLS/Xray-core/pull/6853#issuecomment-5904123488
https://github.com/XTLS/Xray-core/pull/6853#issuecomment-5904647772
https://github.com/XTLS/Xray-core/pull/6853#issuecomment-5905047424

Fixes https://github.com/XTLS/Xray-core/issues/6454#issuecomment-5863800676

---------

Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 06:26:16 +00:00
Meo597 1c52c65872 geodata: rename Lua matcher constructors to BuildDomainMatcher and BuildIPMatcher 2026-09-28 22:52:45 +08:00
Meo597 5724db08f4 lua: standardize hooks and host APIs on PascalCase 2026-09-28 22:52:14 +08:00
Meo597 3d3306503d geodata: flatten Lua matcher rule arguments 2026-09-28 22:23:04 +08:00
Meo597 5e1bb92b98 dns: reduce Lua allocations with flat arguments and returns and native IP slice userdata 2026-09-28 22:12:47 +08:00
Meo597 459301d42e refine Lua script path fallback 2026-09-27 23:15:51 +08:00
Meo597 3982028a9c lua: lower camel case 2026-09-26 17:58:25 +08:00
Meo597 70b8e9a61d log caller 2026-09-26 15:42:57 +08:00
Meo597 219f758060 add log module 2026-09-26 14:40:15 +08:00
Meo597 9628003594 Reduce memory usage 2026-09-26 06:20:34 +08:00
Meo597 72d9ab50b9 add tests 2026-09-26 06:16:54 +08:00
Meo597 235843c5d2 feat(dns): add Lua scripting for DNS queries 2026-09-26 05:37:23 +08:00
59 changed files with 5290 additions and 105 deletions
+3
View File
@@ -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)
}
}
+25 -6
View File
@@ -93,6 +93,7 @@ type NameServer struct {
UnexpectedIp []*geodata.IPRule `protobuf:"bytes,13,rep,name=unexpected_ip,json=unexpectedIp,proto3" json:"unexpected_ip,omitempty"`
ActUnprior bool `protobuf:"varint,14,opt,name=actUnprior,proto3" json:"actUnprior,omitempty"`
PolicyID uint32 `protobuf:"varint,17,opt,name=policyID,proto3" json:"policyID,omitempty"`
Id string `protobuf:"bytes,18,opt,name=id,proto3" json:"id,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
@@ -239,6 +240,13 @@ func (x *NameServer) GetPolicyID() uint32 {
return 0
}
func (x *NameServer) GetId() string {
if x != nil {
return x.Id
}
return ""
}
type Config struct {
state protoimpl.MessageState `protogen:"open.v1"`
// NameServer list used by this DNS client.
@@ -258,8 +266,10 @@ type Config struct {
DisableFallback bool `protobuf:"varint,10,opt,name=disableFallback,proto3" json:"disableFallback,omitempty"`
DisableFallbackIfMatch bool `protobuf:"varint,11,opt,name=disableFallbackIfMatch,proto3" json:"disableFallbackIfMatch,omitempty"`
EnableParallelQuery bool `protobuf:"varint,14,opt,name=enableParallelQuery,proto3" json:"enableParallelQuery,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
// Absolute path to the Lua DNS query script.
Script string `protobuf:"bytes,15,opt,name=script,proto3" json:"script,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *Config) Reset() {
@@ -369,6 +379,13 @@ func (x *Config) GetEnableParallelQuery() bool {
return false
}
func (x *Config) GetScript() string {
if x != nil {
return x.Script
}
return ""
}
type Config_HostMapping struct {
state protoimpl.MessageState `protogen:"open.v1"`
Domain *geodata.DomainRule `protobuf:"bytes,2,opt,name=domain,proto3" json:"domain,omitempty"`
@@ -435,7 +452,7 @@ var File_app_dns_config_proto protoreflect.FileDescriptor
const file_app_dns_config_proto_rawDesc = "" +
"\n" +
"\x14app/dns/config.proto\x12\fxray.app.dns\x1a\x1ccommon/net/destination.proto\x1a\x1bcommon/geodata/geodat.proto\"\xde\x05\n" +
"\x14app/dns/config.proto\x12\fxray.app.dns\x1a\x1ccommon/net/destination.proto\x1a\x1bcommon/geodata/geodat.proto\"\xee\x05\n" +
"\n" +
"NameServer\x123\n" +
"\aaddress\x18\x01 \x01(\v2\x19.xray.common.net.EndpointR\aaddress\x12\x1b\n" +
@@ -461,10 +478,11 @@ const file_app_dns_config_proto_rawDesc = "" +
"\n" +
"actUnprior\x18\x0e \x01(\bR\n" +
"actUnprior\x12\x1a\n" +
"\bpolicyID\x18\x11 \x01(\rR\bpolicyIDB\x0f\n" +
"\bpolicyID\x18\x11 \x01(\rR\bpolicyID\x12\x0e\n" +
"\x02id\x18\x12 \x01(\tR\x02idB\x0f\n" +
"\r_disableCacheB\r\n" +
"\v_serveStaleB\x12\n" +
"\x10_serveExpiredTTLJ\x04\b\x04\x10\x05\"\x82\x05\n" +
"\x10_serveExpiredTTLJ\x04\b\x04\x10\x05\"\x9a\x05\n" +
"\x06Config\x129\n" +
"\vname_server\x18\x05 \x03(\v2\x18.xray.app.dns.NameServerR\n" +
"nameServer\x12\x1b\n" +
@@ -480,7 +498,8 @@ const file_app_dns_config_proto_rawDesc = "" +
"\x0fdisableFallback\x18\n" +
" \x01(\bR\x0fdisableFallback\x126\n" +
"\x16disableFallbackIfMatch\x18\v \x01(\bR\x16disableFallbackIfMatch\x120\n" +
"\x13enableParallelQuery\x18\x0e \x01(\bR\x13enableParallelQuery\x1a}\n" +
"\x13enableParallelQuery\x18\x0e \x01(\bR\x13enableParallelQuery\x12\x16\n" +
"\x06script\x18\x0f \x01(\tR\x06script\x1a}\n" +
"\vHostMapping\x127\n" +
"\x06domain\x18\x02 \x01(\v2\x1f.xray.common.geodata.DomainRuleR\x06domain\x12\x0e\n" +
"\x02ip\x18\x03 \x03(\fR\x02ip\x12%\n" +
+4
View File
@@ -27,6 +27,7 @@ message NameServer {
repeated xray.common.geodata.IPRule unexpected_ip = 13;
bool actUnprior = 14;
uint32 policyID = 17;
string id = 18;
}
enum QueryStrategy {
@@ -73,4 +74,7 @@ message Config {
bool disableFallbackIfMatch = 11;
bool enableParallelQuery = 14;
// Absolute path to the Lua DNS query script.
string script = 15;
}
+16
View File
@@ -31,6 +31,8 @@ type DNS struct {
domainMatcher geodata.DomainMatcher
matcherInfos []*DomainMatcherInfo
checkSystem bool
script *scriptEngine
scriptPath string
}
// DomainMatcherInfo contains information attached to index returned by Server.domainMatcher.
@@ -180,6 +182,7 @@ func New(ctx context.Context, config *Config) (*DNS, error) {
disableFallbackIfMatch: config.DisableFallbackIfMatch,
enableParallelQuery: config.EnableParallelQuery,
checkSystem: checkSystem,
scriptPath: config.Script,
}, nil
}
@@ -190,11 +193,21 @@ func (*DNS) Type() interface{} {
// Start implements common.Runnable.
func (s *DNS) Start() error {
if s.scriptPath != "" {
engine, err := newScriptEngine(s.scriptPath, s)
if err != nil {
return errors.New("failed to initialize DNS script").Base(err)
}
s.script = engine
}
return nil
}
// Close implements common.Closable.
func (s *DNS) Close() error {
if s.script != nil {
s.script.close()
}
return nil
}
@@ -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 {
+169
View File
@@ -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
}
+118
View File
@@ -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)
}
})
}
}
+272
View File
@@ -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
}
+2 -1
View File
@@ -29,6 +29,7 @@ type Server interface {
// Client is the interface for DNS client.
type Client struct {
id string
server Server
skipFallback bool
expectedIPs geodata.IPMatcher
@@ -97,7 +98,7 @@ func NewClient(
ipOption dns.IPOption,
updateRules func(bool),
) (*Client, error) {
client := &Client{}
client := &Client{id: ns.Id}
err := core.RequireFeatures(ctx, func(dispatcher routing.Dispatcher) error {
// Create a new server for each client for now
server, err := NewServer(ctx, ns.Address.AsDestination(), dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP)
+1 -1
View File
@@ -49,5 +49,5 @@ func NewLocalNameServer() *LocalNameServer {
// NewLocalDNSClient creates localdns client object for directly lookup in system DNS.
func NewLocalDNSClient(ipOption dns.IPOption) *Client {
return &Client{server: NewLocalNameServer(), ipOption: &ipOption}
return &Client{id: "localhost", server: NewLocalNameServer(), ipOption: &ipOption}
}
+63
View File
@@ -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
}
+280
View File
@@ -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)
}
}
+14 -4
View File
@@ -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" +
+2
View File
@@ -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;
}
+175
View File
@@ -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)
}
+223
View File
@@ -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
}
+274
View File
@@ -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)
+17
View File
@@ -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())
+76
View File
@@ -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
}
+382
View File
@@ -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")
}
}
+163
View File
@@ -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
}
+172
View File
@@ -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)
}
})
}
}
+61
View File
@@ -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()
}
+213
View File
@@ -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] <string>: 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 != "<string>: message" {
t.Errorf("message %d = %v, want severity %v and content %q", i, msg, severity, "<string>: 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)
}
}
})
}
}
+3
View File
@@ -0,0 +1,3 @@
// Package lua provides shared GopherLua programs, state management, and value
// conversion and validation helpers for Xray scripts.
package lua
+65
View File
@@ -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()
}
+84
View File
@@ -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)
}
})
}
}
+151
View File
@@ -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()
}
+466
View File
@@ -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)
}
}
+77
View File
@@ -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)
}
}
+76
View File
@@ -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())
}
}
+95
View File
@@ -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)
}
+121
View File
@@ -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)
}
}
}
+48
View File
@@ -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)
}
+51
View File
@@ -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)
}
}
}
+1 -1
View File
@@ -20,7 +20,7 @@ import (
var (
Version_x byte = 26
Version_y byte = 9
Version_z byte = 9
Version_z byte = 30
)
var (
+3
View File
@@ -97,6 +97,9 @@ func New() *Client {
r := &net.Resolver{
PreferGo: true,
Dial: func(ctx context.Context, network, address string) (net.Conn, error) {
if internet.IsSkippedDNSServer(address) {
return nil, errors.New("skipped DNS server ", address)
}
return d.DialContext(ctx, network, address)
},
}
+23
View File
@@ -0,0 +1,23 @@
package localdns
import (
"context"
"net/netip"
"testing"
"github.com/xtls/xray-core/transport/internet"
)
func TestSkippedDNSServers(t *testing.T) {
internet.SkipDNSServers([]netip.Addr{netip.MustParseAddr("203.0.113.53")})
t.Cleanup(func() { internet.SkipDNSServers(nil) })
c := New()
if _, err := c.r.Dial(context.Background(), "udp", "203.0.113.53:53"); err == nil {
t.Error("a skipped DNS server was dialed")
}
conn, err := c.r.Dial(context.Background(), "udp", "127.0.0.1:53")
if err != nil {
t.Fatal(err)
}
conn.Close()
}
+2
View File
@@ -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
)
+9
View File
@@ -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=
@@ -105,6 +111,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=
@@ -155,6 +162,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=
+14
View File
@@ -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())
+50
View File
@@ -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)
}
}
+11
View File
@@ -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
+38
View File
@@ -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)
}
})
}
}
+32 -2
View File
@@ -5,8 +5,12 @@ import (
"fmt"
"math/big"
"net"
"runtime"
"slices"
"strconv"
"strings"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/proxy/tun"
"google.golang.org/protobuf/proto"
)
@@ -20,7 +24,8 @@ type TunConfig struct {
UserLevel uint32 `json:"userLevel"`
AutoSystemRoutingTable []string `json:"autoSystemRoutingTable"`
AutoOutboundsInterface *string `json:"autoOutboundsInterface"`
AutoSystemDNS bool `json:"autoSystemDNS"`
AutoSystemDnsToGateway bool `json:"autoSystemDnsToGateway"`
AutoSystemWfpBlockLeak []string `json:"autoSystemWfpBlockLeak"`
}
func (v *TunConfig) Build() (proto.Message, error) {
@@ -32,7 +37,32 @@ func (v *TunConfig) Build() (proto.Message, error) {
DNS: v.DNS,
UserLevel: v.UserLevel,
AutoSystemRoutingTable: v.AutoSystemRoutingTable,
AutoSystemDns: v.AutoSystemDNS,
AutoSystemDnsToGateway: v.AutoSystemDnsToGateway,
}
for _, leak := range v.AutoSystemWfpBlockLeak {
switch leak := strings.ToLower(leak); leak {
case "dns", "misconfigtun":
config.AutoSystemWfpBlockLeak = append(config.AutoSystemWfpBlockLeak, leak)
default:
return nil, errors.New("unknown autoSystemWfpBlockLeak value: ", leak)
}
}
// Each option needs other settings on the system it takes effect on: the
// filters go along with the routes of autoSystemRoutingTable, "dns" lets
// DNS through the TUN only, and autoSystemDnsToGateway points the system
// DNS at the gateway.
switch runtime.GOOS {
case "windows":
if len(config.AutoSystemWfpBlockLeak) > 0 && len(v.AutoSystemRoutingTable) == 0 {
return nil, errors.New("autoSystemWfpBlockLeak needs autoSystemRoutingTable to be set")
}
if slices.Contains(config.AutoSystemWfpBlockLeak, "dns") && len(v.DNS) == 0 {
return nil, errors.New(`autoSystemWfpBlockLeak "dns" needs dns to be set`)
}
case "linux":
if v.AutoSystemDnsToGateway && len(v.Gateway) == 0 {
return nil, errors.New("autoSystemDnsToGateway needs gateway to be set")
}
}
if v.AutoOutboundsInterface != nil {
config.AutoOutboundsInterface = *v.AutoOutboundsInterface
+71
View File
@@ -0,0 +1,71 @@
package conf_test
import (
"encoding/json"
"runtime"
"testing"
. "github.com/xtls/xray-core/infra/conf"
"github.com/xtls/xray-core/proxy/tun"
)
func TestTunConfigAutoSystem(t *testing.T) {
creator := func() Buildable {
return new(TunConfig)
}
runMultiTestCase(t, []TestCase{
{
Input: `{"name": "xray0"}`,
Parser: loadJSON(creator),
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500},
},
{
Input: `{"name": "xray0", "gateway": ["10.0.0.1/24"], "autoSystemDnsToGateway": true}`,
Parser: loadJSON(creator),
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500, Gateway: []string{"10.0.0.1/24"}, AutoSystemDnsToGateway: true},
},
{
Input: `{"name": "xray0", "dns": ["1.1.1.1"], "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["dns", "misconfigtun"]}`,
Parser: loadJSON(creator),
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500, DNS: []string{"1.1.1.1"}, AutoSystemRoutingTable: []string{"0.0.0.0/0"}, AutoOutboundsInterface: "auto", AutoSystemWfpBlockLeak: []string{"dns", "misconfigtun"}},
},
{
Input: `{"name": "xray0", "dns": ["1.1.1.1"], "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["DNS"]}`,
Parser: loadJSON(creator),
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500, DNS: []string{"1.1.1.1"}, AutoSystemRoutingTable: []string{"0.0.0.0/0"}, AutoOutboundsInterface: "auto", AutoSystemWfpBlockLeak: []string{"dns"}},
},
})
}
// TestTunConfigAutoSystemNeeds checks that an option is rejected without the
// setting it needs, only on the system it takes effect on.
func TestTunConfigAutoSystemNeeds(t *testing.T) {
for _, c := range []struct {
input string
goos string // where it is rejected
}{
{`{"name": "xray0", "autoSystemWfpBlockLeak": ["misconfigtun"]}`, "windows"},
{`{"name": "xray0", "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["misconfigtun"]}`, ""},
{`{"name": "xray0", "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["dns"]}`, "windows"},
{`{"name": "xray0", "autoSystemDnsToGateway": true}`, "linux"},
} {
config := new(TunConfig)
if err := json.Unmarshal([]byte(c.input), config); err != nil {
t.Fatal(err)
}
if _, err := config.Build(); (err != nil) != (runtime.GOOS == c.goos) {
t.Errorf("%s on %s: error = %v", c.input, runtime.GOOS, err)
}
}
}
func TestTunConfigAutoSystemWfpBlockLeakUnknown(t *testing.T) {
config := new(TunConfig)
if err := json.Unmarshal([]byte(`{"name": "xray0", "autoSystemWfpBlockLeak": ["dns", "ip"]}`), config); err != nil {
t.Fatal(err)
}
if _, err := config.Build(); err == nil {
t.Error("an unknown autoSystemWfpBlockLeak value was accepted")
}
}
+28 -13
View File
@@ -15,27 +15,28 @@ Plainly enabling it in the config probably will result nothing, or lock your rou
## DETAILS
By default, enabling the feature will only bring the tun interface up. \
When configured explicitly, Windows and Linux can apply interface addresses from `gateway`, while macOS uses the first IPv4 prefix from `gateway` to configure the utun point-to-point address. \
When configured explicitly, Windows and Linux can apply interface addresses from `gateway`, while macOS and FreeBSD use the first IPv4 prefix from `gateway` for the point-to-point address. \
Without `gateway`, the systems differ: Xray assigns no address on Linux, Windows gives the interface link-local addresses itself (an IPv6 one at once, an IPv4 one from `169.254.0.0/16` after a few seconds), and macOS and FreeBSD use `169.254.10.1/30`. \
Windows, Linux and macOS can also apply system routes from `autoSystemRoutingTable`.
macOS does not configure system DNS from the `dns` field, and neither does Linux by default; system DNS remains managed by the OS or distribution-specific network services. \
For more advanced routing policies or rules, OS level configuration can still manage the named interface (e.g. xray0) when it appears.
This keeps complex system level routing and rules in a single place of responsibility - the OS itself. \
Examples of how to achieve this on a simple Linux system (Ubuntu with systemd-networkd) can be found at the end of this README.
### SYSTEM DNS ON LINUX (`autoSystemDNS`)
### SYSTEM DNS ON LINUX (`autoSystemDnsToGateway`)
On Linux, setting `autoSystemDNS` to `true` lets the inbound point the system resolver at the tun interface, so name lookups resolve through Xray instead of going out over the physical link. It is off by default, and it is Linux-only.
On Linux, setting `autoSystemDnsToGateway` to `true` lets the inbound point the system resolver at the tun interface, so name lookups resolve through Xray instead of going out over the physical link. It is off by default, and it is Linux-only.
It uses `resolvectl`, which means it applies only when all of these hold:
It uses `resolvectl`, which means it only works when all of these hold. Where Xray can tell that one does not, it does not start:
- the system runs systemd and `resolvectl` is on `PATH`
- `systemd-resolved` is enabled and actually managing DNS (installed but not running has no effect)
- `systemd-resolved` is enabled and actually managing DNS (installed but not running is not enough)
- systemd-resolved is version 240 or newer, where `default-route` exists
- no `dns` upstream resolves through the system resolver, directly or through its own bootstrap (see below)
The address handed over is the first IPv4 `gateway` incremented by one (e.g. `192.168.100.1/30` -> `192.168.100.2`). It is not taken from `dns`: handing `1.1.1.1` to `resolvectl dns` would make systemd-resolved query that server directly over the physical link, which is the leak this option exists to close.
The address handed over is the first IPv4 `gateway`, or without one the first IPv6 `gateway`, incremented by one (e.g. `192.168.100.1/30` -> `192.168.100.2`, `fc00::1/64` -> `fc00::2`). Without any `gateway`, the config is rejected. It is not taken from `dns`: handing `1.1.1.1` to `resolvectl dns` would make systemd-resolved query that server directly over the physical link, which is the leak this option exists to close.
Because that address has to actually answer, the takeover is checked before it happens. A query from the interface address to that address is routed through the configured rules, and host-wide DNS is only changed when the result is a DNS-capable outbound. Otherwise the option does nothing and DNS is left to the OS. In practice this means you also need a routing rule sending the interface's port 53 to a `dns` outbound, for example:
Because that address has to actually answer, the takeover is checked before it happens. A query from the interface address to that address is routed through the configured rules, and host-wide DNS is only changed when the result is a DNS-capable outbound. Otherwise DNS is left alone and Xray does not start. In practice this means you also need a routing rule sending the interface's port 53 to a `dns` outbound, for example:
```json
"routing": {
@@ -49,19 +50,19 @@ The check is a preflight, not a proof for arbitrary rules. It sends its query fr
It is also a check for the dependencies it knows about, not a proof that no indirect one exists. A hostname-based upstream that bootstraps through system DNS is the case in point: `https+local://dns.google/dns-query` resolves its own hostname with `DialSystem`, so once the takeover is in place that bootstrap goes `resolved -> TUN -> DNS outbound -> bootstrap -> resolved` and the query times out. The preflight does not see it, because the dependency sits in the upstream's bootstrap rather than in the clients it inspects. Upstream resolution, bootstrap included, therefore has to stay independent of the resolver path being redirected; configuring the address instead of the hostname, or resolving the hostname beforehand, avoids it.
The upstream requirement in the list above matters as much as the routing rule. With no name servers configured, Core resolves through a client that forwards to the system resolver; pointing the system resolver at the TUN would then close a loop through the DNS outbound, `resolved -> TUN -> DNS outbound -> system resolver -> resolved`, and resolution stops. The takeover is refused in that case.
The upstream requirement in the list above matters as much as the routing rule. With no name servers configured, Core resolves through a client that forwards to the system resolver; pointing the system resolver at the TUN would then close a loop through the DNS outbound, `resolved -> TUN -> DNS outbound -> system resolver -> resolved`, and resolution stops. The takeover is refused in that case, and Xray does not start.
The same applies to a name server pointed at `localhost`, and to a `dns` section that is present but lists no name servers. One such upstream is enough to refuse the takeover even when independent upstreams are configured alongside it: name servers are selected per domain, so a domain-specific rule can still choose the local one, and the loop then affects whichever domains reach it. The check is deliberately broader than the loop it observed, because the alternative would be to drop a name server the user configured.
Where it does not apply, DNS is left alone and the leak described in XTLS/Xray-core#6454 remains:
Where it cannot apply, Xray does not start, rather than run with the leak described in XTLS/Xray-core#6454, so leave the option off there:
| Environment | Behaviour |
|---|---|
| systemd distribution with systemd-resolved enabled | applies |
| Alpine, Void, Devuan, OpenRC-based, OpenWrt | no `resolvectl`, skipped |
| DNS managed by dnsmasq / unbound / BIND / static `resolv.conf` | unreachable by `resolvectl`, skipped |
| Containers without a systemd-resolved daemon | skipped |
| systemd older than 240 | `default-route` unavailable, skipped |
| Alpine, Void, Devuan, OpenRC-based, OpenWrt | no `resolvectl`, does not start |
| DNS managed by dnsmasq / unbound / BIND / static `resolv.conf` | unreachable by `resolvectl`, does not start |
| Containers without a systemd-resolved daemon | does not start |
| systemd older than 240 | `default-route` unavailable, does not start |
On `Close()` the setting is reverted. It is **not** reverted if the process is killed with `SIGKILL`, since a process cannot handle that signal; run `resolvectl revert <iface>` to clean up by hand. An application that brings its own DNS endpoint is unaffected either way — this only covers the system resolver.
@@ -198,6 +199,20 @@ To make it start, wintun.dll specific for your Windows/arch must be present next
After the start network adapter with the name you chose in the config will be created in the system, and exist while Xray is running.
When `dns` is set, those servers are applied to the adapter. Windows is kept from registering the TUN's addresses in DNS, and its DNS cache is flushed when the TUN starts and stops.
With `autoSystemWfpBlockLeak`, which needs `autoSystemRoutingTable` (the config is rejected otherwise), Xray also adds Windows Filtering Platform filters that keep two kinds of traffic of every program but Xray itself from leaving outside the TUN, each chosen by a value in the list, e.g. `"autoSystemWfpBlockLeak": ["dns", "misconfigtun"]`:
- `"dns"` (needs `dns`, the config is rejected otherwise): DNS (port 53) only goes through the TUN. Windows keeps sending name queries to the DNS servers of the other interfaces as well, out through those interfaces whatever the routes say, and other programs reach a resolver on the local network (e.g. `192.168.1.1` handed out by DHCP) through its more specific LAN route instead of the TUN. On Windows 11 and Server 2022 and later, where those queries may also go over HTTPS or TLS, Windows' DNS Client service cannot connect outside the TUN at all, except for name resolution on the local network (LLMNR, mDNS). The `dns` servers therefore have to lie within `gateway` or `autoSystemRoutingTable` (a warning is logged otherwise), and DNS servers that should be reached directly belong in Xray's own `dns` settings.
- `"misconfigtun"`: an IP version without routes in `autoSystemRoutingTable`, IPv4 or IPv6, is blocked entirely, in both directions, as it would bypass the TUN. Only loopback and what Windows itself needs on the local link (DHCP, and for IPv6 neighbor and multicast listener discovery) remain allowed. An address of that version in `gateway` is not needed: without one, Windows gives the TUN link-local addresses itself, an IPv6 one at once and an IPv4 one from `169.254.0.0/16` after some seconds (until then, IPv4 routed to the TUN is unreachable), and what is routed to the TUN goes through it with those.
With the filters in place, Xray's own connections out also get past Windows Firewall's block rules (other firewalls may still block them), while connections to Xray's inbounds stay subject to them.
Names that Xray resolves through the system resolver, such as an outbound's server address given as a domain with the default `AsIs` domain strategy, would be looked up by Windows on Xray's behalf, and those queries would then go into the TUN too. While DNS is restricted this way and `autoOutboundsInterface` is in use (the default with `autoSystemRoutingTable`), Xray therefore resolves them itself, with its own queries to the DNS servers of the other interfaces. That bypasses Windows' DNS cache, and its name resolution on the local network (LLMNR, mDNS): a server address given as a domain is looked up again for every connection, and a DNS server that does not answer delays each lookup. Having Xray's own `dns` resolve it, through the outbound's `sockopt.domainStrategy`, avoids that. The `localhost` DNS server queries the same servers whenever `autoOutboundsInterface` is in use. Both skip the TUN's own DNS servers, unless another interface uses them as well: queried from Xray itself, they would lead back into it, or nowhere.
If the filters cannot be added, Xray does not start. They are removed when Xray exits. Not covered is name resolution on the local network (LLMNR, mDNS, NetBIOS), except over an IP version that is blocked.
`autoSystemWfpBlockLeak` (Windows only) is empty by default, as the filters break some setups: with `"dns"`, a local DNS resolver other programs use (e.g. on `127.0.0.1:53`), the DNS of another VPN on its own interface, virtual machines whose NAT resolves names on the host, or signing in to a captive portal; with `"misconfigtun"`, IPv4 or IPv6 on the local network while no route of that version leads to the TUN. Without the filters, DNS may leak as described above. To keep an IP version out of the TUN on purpose while still blocking DNS leaks, use only `["dns"]`.
You can give the adapter ip address manually, you can live Windows to give it autogenerated ip address (which take few seconds), it doesn't matter, the traffic going _through_ the interface will be forwarded into the app for proxying. \
Minimal configuration that will work for local machine is routing passing the traffic on-link through the interface.
You will need the interface id for that, unfortunately it is going to change with every Xray start due to implementation ambiguity between Xray and wintun driver.
+16 -6
View File
@@ -32,7 +32,8 @@ type Config struct {
AutoSystemRoutingTable []string `protobuf:"bytes,6,rep,name=auto_system_routing_table,json=autoSystemRoutingTable,proto3" json:"auto_system_routing_table,omitempty"`
AutoOutboundsInterface string `protobuf:"bytes,7,opt,name=auto_outbounds_interface,json=autoOutboundsInterface,proto3" json:"auto_outbounds_interface,omitempty"`
Desc string `protobuf:"bytes,8,opt,name=desc,proto3" json:"desc,omitempty"`
AutoSystemDns bool `protobuf:"varint,9,opt,name=auto_system_dns,json=autoSystemDns,proto3" json:"auto_system_dns,omitempty"`
AutoSystemDnsToGateway bool `protobuf:"varint,9,opt,name=auto_system_dns_to_gateway,json=autoSystemDnsToGateway,proto3" json:"auto_system_dns_to_gateway,omitempty"`
AutoSystemWfpBlockLeak []string `protobuf:"bytes,10,rep,name=auto_system_wfp_block_leak,json=autoSystemWfpBlockLeak,proto3" json:"auto_system_wfp_block_leak,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
@@ -123,18 +124,25 @@ func (x *Config) GetDesc() string {
return ""
}
func (x *Config) GetAutoSystemDns() bool {
func (x *Config) GetAutoSystemDnsToGateway() bool {
if x != nil {
return x.AutoSystemDns
return x.AutoSystemDnsToGateway
}
return false
}
func (x *Config) GetAutoSystemWfpBlockLeak() []string {
if x != nil {
return x.AutoSystemWfpBlockLeak
}
return nil
}
var File_proxy_tun_config_proto protoreflect.FileDescriptor
const file_proxy_tun_config_proto_rawDesc = "" +
"\n" +
"\x16proxy/tun/config.proto\x12\x0exray.proxy.tun\"\xaa\x02\n" +
"\x16proxy/tun/config.proto\x12\x0exray.proxy.tun\"\xfa\x02\n" +
"\x06Config\x12\x12\n" +
"\x04name\x18\x01 \x01(\tR\x04name\x12\x10\n" +
"\x03MTU\x18\x02 \x01(\rR\x03MTU\x12\x18\n" +
@@ -144,8 +152,10 @@ const file_proxy_tun_config_proto_rawDesc = "" +
"user_level\x18\x05 \x01(\rR\tuserLevel\x129\n" +
"\x19auto_system_routing_table\x18\x06 \x03(\tR\x16autoSystemRoutingTable\x128\n" +
"\x18auto_outbounds_interface\x18\a \x01(\tR\x16autoOutboundsInterface\x12\x12\n" +
"\x04desc\x18\b \x01(\tR\x04desc\x12&\n" +
"\x0fauto_system_dns\x18\t \x01(\bR\rautoSystemDnsBL\n" +
"\x04desc\x18\b \x01(\tR\x04desc\x12:\n" +
"\x1aauto_system_dns_to_gateway\x18\t \x01(\bR\x16autoSystemDnsToGateway\x12:\n" +
"\x1aauto_system_wfp_block_leak\x18\n" +
" \x03(\tR\x16autoSystemWfpBlockLeakBL\n" +
"\x12com.xray.proxy.tunP\x01Z#github.com/xtls/xray-core/proxy/tun\xaa\x02\x0eXray.Proxy.Tunb\x06proto3"
var (
+2 -1
View File
@@ -15,5 +15,6 @@ message Config {
repeated string auto_system_routing_table = 6;
string auto_outbounds_interface = 7;
string desc = 8;
bool auto_system_dns = 9;
bool auto_system_dns_to_gateway = 9;
repeated string auto_system_wfp_block_leak = 10;
}
+4 -2
View File
@@ -166,12 +166,14 @@ func (t *Handler) Start() error {
}
// Platform-specific system DNS takeover, where the platform implements it.
// Non-fatal: a failure leaves DNS management with the OS.
// Rather no TUN than one that the system DNS bypasses.
if c, ok := tunInterface.(interface {
ConfigureSystemDNS(context.Context, string) error
}); ok {
if err := c.ConfigureSystemDNS(t.ctx, t.tag); err != nil {
errors.LogInfoInner(t.ctx, err, "[tun] system DNS not configured")
_ = tunStack.Close()
_ = tunInterface.Close()
return errors.New("unable to set the system DNS (remove autoSystemDnsToGateway to run without)").Base(err)
}
}
+21 -15
View File
@@ -53,23 +53,29 @@ var resolvectlRunner = func(name string, args ...string) ([]byte, error) {
}
// systemDNSAddrs derives the addresses used for the system DNS takeover from the
// first IPv4 gateway: the gateway address itself is what a query from this
// interface appears to come from, and the next address is what the resolver is
// pointed at. The latter belongs to the TUN and is answered inside Xray;
// handing the configured public resolvers to resolvectl instead would leave the
// system querying them directly over the physical link, defeating the point of
// the TUN.
// first IPv4 gateway, or without one, the first IPv6 gateway: the gateway
// address itself is what a query from this interface appears to come from, and
// the next address is what the resolver is pointed at. The latter belongs to
// the TUN and is answered inside Xray; handing the configured public resolvers
// to resolvectl instead would leave the system querying them directly over the
// physical link, defeating the point of the TUN.
func systemDNSAddrs(gateway []string) (source, dns netip.Addr, ok bool) {
var first6 netip.Addr
for _, address := range gateway {
prefix, err := netip.ParsePrefix(address)
if err != nil {
continue
}
addr := prefix.Addr()
if !addr.Is4() {
continue
if addr.Is4() {
return addr, addr.Next(), true
}
return addr, addr.Next(), true
if !first6.IsValid() {
first6 = addr
}
}
if first6.IsValid() {
return first6, first6.Next(), true
}
return netip.Addr{}, netip.Addr{}, false
}
@@ -115,11 +121,11 @@ const probeSourcePort = 49152
// Overridable for tests.
var verifyDNSRouting = func(ctx context.Context, inboundTag, source, address string) error {
ip, err := netip.ParseAddr(address)
if err != nil || !ip.Is4() {
if err != nil {
return errors.New("invalid DNS address ", address).Base(err)
}
src, err := netip.ParseAddr(source)
if err != nil || !src.Is4() {
if err != nil || src.Is4() != ip.Is4() {
return errors.New("invalid source address ", source).Base(err)
}
@@ -182,10 +188,10 @@ var verifyDNSRouting = func(ctx context.Context, inboundTag, source, address str
//
// It acts only when the config opts in, and it verifies the data path first:
// unless a query to the advertised address would actually be handled, host-wide
// resolution is left to the OS, which is the documented default. Errors are
// returned to the caller, which treats them as non-fatal.
// resolution is left to the OS and an error returned. The caller does not start
// the TUN on an error, as the system DNS would bypass it.
func (t *LinuxTun) ConfigureSystemDNS(ctx context.Context, inboundTag string) error {
if !t.options.AutoSystemDns {
if !t.options.AutoSystemDnsToGateway {
return nil
}
if t.systemDNSSet {
@@ -202,7 +208,7 @@ func (t *LinuxTun) ConfigureSystemDNS(ctx context.Context, inboundTag string) er
source, address, ok := systemDNSAddrs(t.options.Gateway)
if !ok {
return errors.New("no IPv4 gateway, cannot derive a system DNS address")
return errors.New("no gateway, cannot derive a system DNS address")
}
iface := t.ifaceName()
+12
View File
@@ -191,3 +191,15 @@ func TestVerifyDNSRoutingDecisions(t *testing.T) {
})
}
}
// Without an IPv4 gateway, the takeover uses the first IPv6 one, and the probe
// carries IPv6 addresses.
func TestVerifyDNSRoutingIPv6(t *testing.T) {
ctx := newRouteTestContext(t, true, udpNameServer([]byte{9, 9, 9, 9}), []*router.RoutingRule{port53Rule()})
if err := verifyDNSRouting(ctx, routeTestInboundTag, "fc00::1", "fc00::2"); err != nil {
t.Fatalf("expected the takeover to be accepted, got: %v", err)
}
if err := verifyDNSRouting(ctx, routeTestInboundTag, routeTestSource, "fc00::2"); err == nil {
t.Fatal("expected mixed IPv4 and IPv6 addresses to be refused")
}
}
+17 -8
View File
@@ -58,9 +58,9 @@ func recorder(t *testing.T, failOn string) *[][]string {
func optedInTun() *LinuxTun {
return &LinuxTun{
options: &Config{
Name: "xray_tun",
Gateway: []string{"192.168.100.1/30"},
AutoSystemDns: true,
Name: "xray_tun",
Gateway: []string{"192.168.100.1/30"},
AutoSystemDnsToGateway: true,
},
tunLink: testLink("xray_tun"),
}
@@ -79,7 +79,7 @@ func TestConfigureSystemDNSDisabledByDefault(t *testing.T) {
calls := recorder(t, "")
t1 := optedInTun()
t1.options.AutoSystemDns = false
t1.options.AutoSystemDnsToGateway = false
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
t.Fatalf("unexpected error: %v", err)
@@ -103,7 +103,7 @@ func TestConfigureSystemDNSNoGateway(t *testing.T) {
t1.options.Gateway = nil
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err == nil {
t.Fatal("expected an error when no IPv4 gateway is configured")
t.Fatal("expected an error when no gateway is configured")
}
if len(*probes) != 0 {
t.Errorf("routing probe must not run without a gateway, got %d calls", len(*probes))
@@ -351,9 +351,18 @@ func TestSystemDNSAddrs(t *testing.T) {
wantOK: false,
},
{
name: "ipv6 only",
gateway: []string{"fc00::1/64"},
wantOK: false,
name: "ipv6 only",
gateway: []string{"fc00::1/64"},
wantSource: "fc00::1",
wantDNS: "fc00::2",
wantOK: true,
},
{
name: "first ipv6 without ipv4",
gateway: []string{"fc00::1/64", "fd00::1/64"},
wantSource: "fc00::1",
wantDNS: "fc00::2",
wantOK: true,
},
}
+229 -2
View File
@@ -3,17 +3,25 @@
package tun
import (
"bytes"
"context"
"crypto/md5"
"encoding/binary"
go_errors "errors"
"net"
"net/netip"
"os/exec"
"path/filepath"
"slices"
"strconv"
"strings"
"sync"
"syscall"
"time"
"unsafe"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/transport/internet"
"golang.org/x/sys/windows"
"golang.zx2c4.com/wintun"
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
@@ -38,6 +46,10 @@ type WindowsTun struct {
luid winipcfg.LUID
cbr winipcfg.ChangeCallback
cbi winipcfg.ChangeCallback
wfp windows.Handle
resolver *savedResolver
skipStop chan struct{}
skipDone chan struct{}
closed bool
}
@@ -197,19 +209,105 @@ startOver:
}
}
// Windows lists the TUN's DNS servers among the system's ones, which Go's
// resolver queries for Xray's own lookups past the TUN, where they lead
// nowhere or back into Xray. Not skipped are those another interface uses
// as well, as that could leave no server at all. As those can change at
// any time, they are looked at again as often as Go rereads its servers.
if len(dns) > 0 {
skipped, err := tunOnlyDNS(t.luid, dns)
if err != nil {
skipped = dns
}
internet.SkipDNSServers(skipped)
t.skipStop, t.skipDone = make(chan struct{}), make(chan struct{})
go func() {
defer close(t.skipDone)
ticker := time.NewTicker(5 * time.Second)
defer ticker.Stop()
for {
select {
case <-ticker.C:
if skipped, err := tunOnlyDNS(t.luid, dns); err == nil {
internet.SkipDNSServers(skipped)
}
case <-t.skipStop:
return
}
}
}()
}
// Keep Windows from registering the TUN's addresses, and the host name
// with them, through dynamic DNS updates. Best effort.
if address4 || address6 {
if err := disableDNSRegistration(t.luid, dns); err != nil {
errors.LogDebugInner(context.Background(), err, "[tun] unable to disable DNS registration")
}
}
// With autoSystemWfpBlockLeak, once the system routes lead to the TUN,
// keep DNS ("dns", if dns is set), and an IP version no route of which
// leads to the TUN ("misconfigtun"), from leaving through the other
// interfaces. Addresses do not matter: without one of a version in
// gateway, Windows gives the TUN a link-local one.
leaks := t.options.AutoSystemWfpBlockLeak
blockDNS := slices.Contains(leaks, "dns") && len(dns) > 0
blockIPv4 := slices.Contains(leaks, "misconfigtun") && !route4
blockIPv6 := slices.Contains(leaks, "misconfigtun") && !route6
if (route4 || route6) && (blockDNS || blockIPv4 || blockIPv6) {
if t.wfp, err = blockLeaks(t.luid, blockDNS, blockIPv4, blockIPv6); err != nil {
var blocked []string
for _, b := range []struct {
on bool
what string
}{{blockDNS, "DNS"}, {blockIPv4, "IPv4"}, {blockIPv6, "IPv6"}} {
if b.on {
blocked = append(blocked, b.what)
}
}
// Rather no TUN than a leaking one.
return errors.New("unable to block ", strings.Join(blocked, " and "), " outside the TUN (remove autoSystemWfpBlockLeak to run without)").Base(err)
}
errors.LogInfo(context.Background(), "[tun] outside the TUN, blocked DNS: ", blockDNS, ", blocked IPv4: ", blockIPv4, ", blocked IPv6: ", blockIPv6)
if blockDNS {
covered := slices.Clone(addresses)
for _, route := range routesData {
covered = append(covered, route.Destination)
}
for _, server := range dnsOutsideTUN(dns, covered) {
errors.LogWarning(context.Background(), "[tun] DNS server ", server, " is in neither gateway nor autoSystemRoutingTable, so queries to it cannot go through the TUN and are blocked")
}
// With updater, the dialer controllers bind Xray's own sockets
// to the physical interface.
if updater != nil {
t.resolver = resolveOnOwn()
}
}
}
if len(dns) > 0 || route4 || route6 {
if err := flushDNSCache(); err != nil {
errors.LogInfoInner(context.Background(), err, "[tun] unable to flush DNS cache")
}
}
if updater != nil {
t.cbr, err = winipcfg.RegisterRouteChangeCallback(func(notificationType winipcfg.MibNotificationType, route *winipcfg.MibIPforwardRow2) {
// Only a registered callback goes into the fields: a nil pointer in
// them would not compare equal to nil in Close.
cbr, err := winipcfg.RegisterRouteChangeCallback(func(notificationType winipcfg.MibNotificationType, route *winipcfg.MibIPforwardRow2) {
updater.Update()
})
if err != nil {
return err
}
t.cbi, err = winipcfg.RegisterInterfaceChangeCallback(func(notificationType winipcfg.MibNotificationType, iface *winipcfg.MibIPInterfaceRow) {
t.cbr = cbr
cbi, err := winipcfg.RegisterInterfaceChangeCallback(func(notificationType winipcfg.MibNotificationType, iface *winipcfg.MibIPInterfaceRow) {
updater.Update()
})
if err != nil {
return err
}
t.cbi = cbi
}
return nil
}
@@ -236,6 +334,20 @@ func (t *WindowsTun) Close() error {
t.luid.FlushIPAddresses(windows.AF_INET6)
t.luid.FlushDNS(windows.AF_INET6)
}
if t.wfp != 0 {
closeWFPEngine(t.wfp)
}
if t.resolver != nil {
t.resolver.restore()
}
if t.skipStop != nil {
close(t.skipStop)
<-t.skipDone
}
internet.SkipDNSServers(nil)
if len(t.options.DNS) > 0 || len(t.options.AutoSystemRoutingTable) > 0 {
flushDNSCache()
}
if t.session != (wintun.Session{}) {
t.session.End()
}
@@ -245,6 +357,121 @@ func (t *WindowsTun) Close() error {
return nil
}
type savedResolver struct {
preferGo bool
dial func(ctx context.Context, network, address string) (net.Conn, error)
}
// resolveOnOwn has Go resolve the names Xray would otherwise ask Windows for,
// on Xray's own sockets, which the dialer controllers bind to the physical
// interface, and skipping the TUN's DNS servers, as localdns does. Windows'
// resolver runs in the DNS Client service, whose queries the DNS filter lets
// through the TUN only, so Xray's own lookups, like of an outbound's server
// domain, would go into Xray again and could end up waiting on themselves.
//
// It changes net.DefaultResolver for the whole process, which covers every
// lookup that would reach Windows' resolver; restore undoes it.
func resolveOnOwn() *savedResolver {
saved := &savedResolver{net.DefaultResolver.PreferGo, net.DefaultResolver.Dial}
dialer := &net.Dialer{Control: func(network, address string, c syscall.RawConn) error {
for _, ctl := range internet.Controllers {
if err := ctl(network, address, c); err != nil {
return err
}
}
return nil
}}
// Go's resolver moves on to the next server right away when a dial fails.
net.DefaultResolver.Dial = func(ctx context.Context, network, address string) (net.Conn, error) {
if internet.IsSkippedDNSServer(address) {
return nil, errors.New("skipped DNS server ", address)
}
return dialer.DialContext(ctx, network, address)
}
net.DefaultResolver.PreferGo = true
return saved
}
func (s *savedResolver) restore() {
net.DefaultResolver.PreferGo = s.preferGo
net.DefaultResolver.Dial = s.dial
}
// tunOnlyDNS returns those of servers, the TUN's DNS servers, that Go's
// resolver does not also get from another interface: one that is up and has
// a gateway, as it reads them.
func tunOnlyDNS(tun winipcfg.LUID, servers []netip.Addr) ([]netip.Addr, error) {
adapters, err := winipcfg.GetAdaptersAddresses(windows.AF_UNSPEC, winipcfg.GAAFlagIncludeGateways)
if err != nil {
return nil, err
}
var others []netip.Addr
for _, adapter := range adapters {
if adapter.LUID == tun || adapter.OperStatus != winipcfg.IfOperStatusUp || adapter.FirstGatewayAddress == nil {
continue
}
for server := adapter.FirstDNSServerAddress; server != nil; server = server.Next {
if addr, ok := netip.AddrFromSlice(server.Address.IP()); ok {
others = append(others, addr.Unmap())
}
}
}
return slices.DeleteFunc(slices.Clone(servers), func(server netip.Addr) bool {
return slices.Contains(others, server.Unmap())
}), nil
}
// disableDNSRegistration turns off the dynamic DNS registration of the
// interface's addresses. dns are its DNS servers.
func disableDNSRegistration(luid winipcfg.LUID, dns []netip.Addr) error {
guid, err := luid.GUID()
if err != nil {
return err
}
err = winipcfg.SetInterfaceDnsSettings(*guid, &winipcfg.DnsInterfaceSettings{
Version: winipcfg.DnsInterfaceSettingsVersion1,
Flags: winipcfg.DnsInterfaceSettingsFlagRegistrationEnabled,
})
if err == nil || !go_errors.Is(err, windows.ERROR_PROC_NOT_FOUND) {
return err
}
return disableDNSRegistrationByNetsh(luid, dns)
}
// disableDNSRegistrationByNetsh does it for Windows before 10 1809, which
// lacks SetInterfaceDnsSettings. The setting is the interface's, not the
// address family's, but netsh only applies it along with a DNS server, which
// replaces the IPv4 ones, so they are set again afterwards.
func disableDNSRegistrationByNetsh(luid winipcfg.LUID, dns []netip.Addr) error {
row, err := luid.Interface()
if err != nil {
return err
}
server := "127.0.0.1" // any will do when there is no IPv4 one
if i := slices.IndexFunc(dns, netip.Addr.Is4); i >= 0 {
server = dns[i].String()
}
err = runNetsh("interface", "ipv4", "set", "dnsservers", "name="+strconv.FormatUint(uint64(row.InterfaceIndex), 10), "source=static", "address="+server, "register=none", "validate=no")
return errors.Combine(err, luid.SetDNS(windows.AF_INET, dns, nil))
}
// runNetsh runs netsh.exe from the system directory. netsh reports some
// failures, like a syntax error, only in its output, even with exit code 0,
// so any output counts as a failure.
func runNetsh(args ...string) error {
system32, err := windows.GetSystemDirectory()
if err != nil {
return err
}
cmd := exec.Command(filepath.Join(system32, "netsh.exe"), args...)
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
output, err := cmd.CombinedOutput()
if output = bytes.TrimSpace(output); err != nil || len(output) > 0 {
return errors.New("netsh ", strings.Join(args, " "), ": ", string(output)).Base(err)
}
return nil
}
func (t *WindowsTun) Name() (string, error) {
row, err := t.luid.Interface()
if err != nil {
+471
View File
@@ -0,0 +1,471 @@
//go:build windows
package tun
import (
"net/netip"
"os"
"runtime"
"slices"
"unsafe"
"github.com/xtls/xray-core/common/errors"
"golang.org/x/sys/windows"
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
)
var (
modfwpuclnt = windows.NewLazySystemDLL("fwpuclnt.dll")
moddnsapi = windows.NewLazySystemDLL("dnsapi.dll")
procFwpmEngineOpen0 = modfwpuclnt.NewProc("FwpmEngineOpen0")
procFwpmEngineClose0 = modfwpuclnt.NewProc("FwpmEngineClose0")
procFwpmTransactionBegin0 = modfwpuclnt.NewProc("FwpmTransactionBegin0")
procFwpmTransactionCommit0 = modfwpuclnt.NewProc("FwpmTransactionCommit0")
procFwpmTransactionAbort0 = modfwpuclnt.NewProc("FwpmTransactionAbort0")
procFwpmSubLayerAdd0 = modfwpuclnt.NewProc("FwpmSubLayerAdd0")
procFwpmFilterAdd0 = modfwpuclnt.NewProc("FwpmFilterAdd0")
procFwpmGetAppIdFromFileName0 = modfwpuclnt.NewProc("FwpmGetAppIdFromFileName0")
procFwpmFreeMemory0 = modfwpuclnt.NewProc("FwpmFreeMemory0")
procDnsFlushResolverCache = moddnsapi.NewProc("DnsFlushResolverCache")
)
// fwptypes.h and fwpmtypes.h
const (
rpcCAuthnWinNT = 10 // RPC_C_AUTHN_WINNT
fwpmSessionFlagDynamic = 1 // FWPM_SESSION_FLAG_DYNAMIC
fwpmFilterFlagClearActionRight = 8 // FWPM_FILTER_FLAG_CLEAR_ACTION_RIGHT
fwpUint8 = 1 // FWP_UINT8
fwpUint16 = 2 // FWP_UINT16
fwpUint32 = 3 // FWP_UINT32
fwpUint64 = 4 // FWP_UINT64
fwpByteArray16Type = 11 // FWP_BYTE_ARRAY16_TYPE
fwpByteBlobType = 12 // FWP_BYTE_BLOB_TYPE
fwpSecurityDescriptorType = 14 // FWP_SECURITY_DESCRIPTOR_TYPE
fwpMatchEqual = 0 // FWP_MATCH_EQUAL
fwpMatchFlagsAllSet = 6 // FWP_MATCH_FLAGS_ALL_SET
fwpConditionFlagIsLoopback = 1 // FWP_CONDITION_FLAG_IS_LOOPBACK
fwpActionBlock = 0x1001 // FWP_ACTION_BLOCK
fwpActionPermit = 0x1002 // FWP_ACTION_PERMIT
)
// fwpmu.h
var (
fwpmLayerALEAuthConnectV4 = windows.GUID{Data1: 0xc38d57d1, Data2: 0x05a7, Data3: 0x4c33, Data4: [8]byte{0x90, 0x4f, 0x7f, 0xbc, 0xee, 0xe6, 0x0e, 0x82}}
fwpmLayerALEAuthConnectV6 = windows.GUID{Data1: 0x4a72393b, Data2: 0x319f, Data3: 0x44bc, Data4: [8]byte{0x84, 0xc3, 0xba, 0x54, 0xdc, 0xb3, 0xb6, 0xb4}}
fwpmLayerALEAuthRecvAcceptV4 = windows.GUID{Data1: 0xe1cd9fe7, Data2: 0xf4b5, Data3: 0x4273, Data4: [8]byte{0x96, 0xc0, 0x59, 0x2e, 0x48, 0x7b, 0x86, 0x50}}
fwpmLayerALEAuthRecvAcceptV6 = windows.GUID{Data1: 0xa3b42c97, Data2: 0x9f04, Data3: 0x4672, Data4: [8]byte{0xb8, 0x7e, 0xce, 0xe9, 0xc4, 0x83, 0x25, 0x7f}}
fwpmConditionFlags = windows.GUID{Data1: 0x632ce23b, Data2: 0x5167, Data3: 0x435c, Data4: [8]byte{0x86, 0xd7, 0xe9, 0x03, 0x68, 0x4a, 0xa8, 0x0c}}
fwpmConditionIPArrivalInterface = windows.GUID{Data1: 0x618a9b6d, Data2: 0x386b, Data3: 0x4136, Data4: [8]byte{0xad, 0x6e, 0xb5, 0x15, 0x87, 0xcf, 0xb1, 0xcd}}
fwpmConditionIPLocalInterface = windows.GUID{Data1: 0x4cd62a49, Data2: 0x59c3, Data3: 0x4969, Data4: [8]byte{0xb7, 0xf3, 0xbd, 0xa5, 0xd3, 0x28, 0x90, 0xa4}}
fwpmConditionIPLocalPort = windows.GUID{Data1: 0x0c1ba1af, Data2: 0x5765, Data3: 0x453f, Data4: [8]byte{0xaf, 0x22, 0xa8, 0xf7, 0x91, 0xac, 0x77, 0x5b}} // also FWPM_CONDITION_ICMP_TYPE
fwpmConditionIPNexthopInterface = windows.GUID{Data1: 0x93ae8f5b, Data2: 0x7f6f, Data3: 0x4719, Data4: [8]byte{0x98, 0xc8, 0x14, 0xe9, 0x74, 0x29, 0xef, 0x04}}
fwpmConditionIPProtocol = windows.GUID{Data1: 0x3971ef2b, Data2: 0x623e, Data3: 0x4f9a, Data4: [8]byte{0x8c, 0xb1, 0x6e, 0x79, 0xb8, 0x06, 0xb9, 0xa7}}
fwpmConditionIPRemoteAddress = windows.GUID{Data1: 0xb235ae9a, Data2: 0x1d64, Data3: 0x49b8, Data4: [8]byte{0xa4, 0x4c, 0x5f, 0xf3, 0xd9, 0x09, 0x50, 0x45}}
fwpmConditionIPRemotePort = windows.GUID{Data1: 0xc35a604d, Data2: 0xd22b, Data3: 0x4e1a, Data4: [8]byte{0x91, 0xb4, 0x68, 0xf6, 0x74, 0xee, 0x67, 0x4b}} // also FWPM_CONDITION_ICMP_CODE
fwpmConditionALEAppID = windows.GUID{Data1: 0xd78e1e87, Data2: 0x8644, Data3: 0x4ea5, Data4: [8]byte{0x94, 0x37, 0xd8, 0x09, 0xec, 0xef, 0xc9, 0x71}}
fwpmConditionALEUserID = windows.GUID{Data1: 0xaf043a0a, Data2: 0xb34d, Data3: 0x4f86, Data4: [8]byte{0x97, 0x9c, 0xc9, 0x03, 0x71, 0xaf, 0x6e, 0x66}}
)
// dnsClientSID is the SID of Windows' DNS Client service, NT SERVICE\Dnscache.
// Service SIDs derive from the service name, so it is the same everywhere (sc
// showsid dnscache).
const dnsClientSID = "S-1-5-80-859482183-879914841-863379149-1145462774-2388618682"
// ff02::1:2, where DHCPv6 clients send to. A package-level variable never
// moves, so conditions may refer to it through uintptr.
var ipv6AllDHCPv6Servers = [16]byte{0xff, 0x02, 13: 0x01, 15: 0x02}
type fwpByteBlob struct {
size uint32
data *byte
}
// fwpValue0 is FWP_VALUE0 as well as FWP_CONDITION_VALUE0. Their union holds
// a scalar of at most 32 bits, or a pointer for the larger types.
type fwpValue0 struct {
typ uint32
value uintptr
}
type fwpmDisplayData0 struct {
name *uint16
description *uint16
}
type fwpmSession0 struct {
sessionKey windows.GUID
displayData fwpmDisplayData0
flags uint32
txnWaitTimeoutInMSec uint32
processID uint32
sid *windows.SID
username *uint16
kernelMode int32
}
type fwpmSublayer0 struct {
subLayerKey windows.GUID
displayData fwpmDisplayData0
flags uint32
providerKey *windows.GUID
providerData fwpByteBlob
weight uint16
}
type fwpmFilterCondition0 struct {
fieldKey windows.GUID
matchType uint32
conditionValue fwpValue0
}
type fwpmAction0 struct {
typ uint32
filterType windows.GUID
}
type fwpmFilter0 struct {
filterKey windows.GUID
displayData fwpmDisplayData0
flags uint32
providerKey *windows.GUID
providerData fwpByteBlob
layerKey windows.GUID
subLayerKey windows.GUID
weight fwpValue0
numFilterConditions uint32
filterCondition *fwpmFilterCondition0
action fwpmAction0
_ uint32 // C aligns the following union to 8 bytes, as it holds a UINT64
providerContextKey windows.GUID
reserved *windows.GUID
_ [8 - unsafe.Sizeof(uintptr(0))]byte // and filterId as well, also on 32-bit
filterID uint64
effectiveWeight fwpValue0
}
// fwpmResult converts the DWORD status the Fwpm functions return.
func fwpmResult(r1, _ uintptr, _ error) error {
if r1 != 0 {
return windows.Errno(r1)
}
return nil
}
func utf16Ptr(s string) *uint16 {
p, _ := windows.UTF16PtrFromString(s)
return p
}
func condition(field *windows.GUID, typ uint32, value uintptr) fwpmFilterCondition0 {
return fwpmFilterCondition0{
fieldKey: *field,
matchType: fwpMatchEqual,
conditionValue: fwpValue0{typ: typ, value: value},
}
}
// blockLeaks keeps traffic from leaving through interfaces other than tun,
// for every program but Xray itself, whose outbounds (DNS included) use the
// other interfaces on purpose:
//
// - dns: DNS (port 53) may only go through the TUN. Windows sends a name
// query to the DNS servers of all interfaces, not only to those of the TUN:
// to the first server of each interface, then to all of them when no answer
// arrives within a second or two. It sends the queries for the servers of
// an interface out through that interface, whatever the routes say, and
// other programs reach an on-link resolver, like 192.168.1.1 from DHCP,
// through its LAN route, which is more specific than the TUN's default
// route. Since Windows 11 and Server 2022, Windows may also send its
// queries over HTTPS or TLS, so there its DNS Client service may not
// connect outside the TUN at all, except for name resolution on the local
// link (mDNS, LLMNR).
// - ipv4, ipv6: no IPv4, or no IPv6, at all, in either direction, for a TUN
// that no route of it leads to, except loopback and what Windows itself
// needs on the local link (DHCP, and for IPv6 neighbor and multicast
// listener discovery), none of which can leave it. The TUN carries what
// is routed to it even without an address of that IP version in gateway:
// Windows gives it link-local ones itself, an IPv6 one at once, an IPv4
// one from 169.254.0.0/16 after some seconds (until then, IPv4 routed to
// the TUN is unreachable).
//
// The filters live in a dynamic WFP session: closing the returned engine handle
// with closeWFPEngine deletes them, and so does Windows when the process dies.
func blockLeaks(tun winipcfg.LUID, dns, ipv4, ipv6 bool) (windows.Handle, error) {
engine, err := openWFPEngine()
if err != nil {
return 0, err
}
if err := fwpmResult(procFwpmTransactionBegin0.Call(uintptr(engine), 0)); err != nil {
closeWFPEngine(engine)
return 0, errors.New("FwpmTransactionBegin0 failed").Base(err)
}
err = addLeakFilters(engine, tun, dns, ipv4, ipv6)
if err == nil {
if err = fwpmResult(procFwpmTransactionCommit0.Call(uintptr(engine))); err != nil {
err = errors.New("FwpmTransactionCommit0 failed").Base(err)
}
}
if err != nil {
procFwpmTransactionAbort0.Call(uintptr(engine))
closeWFPEngine(engine)
return 0, err
}
return engine, nil
}
func openWFPEngine() (windows.Handle, error) {
if err := modfwpuclnt.Load(); err != nil {
return 0, err
}
// txnWaitTimeoutInMSec stays 0 for BFE's default, so that a transaction
// held by another program cannot hang the start forever.
session := fwpmSession0{
displayData: fwpmDisplayData0{name: utf16Ptr("Xray TUN")},
flags: fwpmSessionFlagDynamic,
}
var engine windows.Handle
if err := fwpmResult(procFwpmEngineOpen0.Call(0, rpcCAuthnWinNT, 0, uintptr(unsafe.Pointer(&session)), uintptr(unsafe.Pointer(&engine)))); err != nil {
return 0, errors.New("FwpmEngineOpen0 failed").Base(err)
}
return engine, nil
}
func closeWFPEngine(engine windows.Handle) {
procFwpmEngineClose0.Call(uintptr(engine))
}
// addLeakFilters adds the filters of blockLeaks in a sublayer of their own.
// blockLeaks runs it in a transaction, so that they take effect all at once.
func addLeakFilters(engine windows.Handle, tun winipcfg.LUID, dns, ipv4, ipv6 bool) error {
exe, err := os.Executable()
if err != nil {
return err
}
exePath, err := windows.UTF16PtrFromString(exe)
if err != nil {
return err
}
var appID *fwpByteBlob
if err := fwpmResult(procFwpmGetAppIdFromFileName0.Call(uintptr(unsafe.Pointer(exePath)), uintptr(unsafe.Pointer(&appID)))); err != nil {
return errors.New("FwpmGetAppIdFromFileName0 failed for ", exe).Base(err)
}
defer func() { procFwpmFreeMemory0.Call(uintptr(unsafe.Pointer(&appID))) }()
sublayer := fwpmSublayer0{
displayData: fwpmDisplayData0{name: utf16Ptr("Xray TUN")},
weight: 0xffff,
}
if sublayer.subLayerKey, err = windows.GenerateGUID(); err != nil {
return err
}
if err := fwpmResult(procFwpmSubLayerAdd0.Call(uintptr(engine), uintptr(unsafe.Pointer(&sublayer)), 0)); err != nil {
return errors.New("FwpmSubLayerAdd0 failed").Base(err)
}
add := func(layer *windows.GUID, name string, flags, action uint32, weight uint8, conditions ...fwpmFilterCondition0) error {
return addFilter(engine, &sublayer.subLayerKey, layer, "Xray TUN: "+name, flags, action, weight, conditions...)
}
var pinner runtime.Pinner
defer pinner.Unpin()
tunLUID := new(uint64)
*tunLUID = uint64(tun)
pinner.Pin(tunLUID) // the condition only holds it as uintptr
// The heaviest matching filter of a sublayer decides. All sublayers have
// their say, though, and a block in any of them beats a permit, unless
// the permit is hard: it clears the action right, and then the blocks of
// lower sublayers, Windows Firewall rules among them, no longer override
// it, only a callout's veto does. Xray's own connections out get such a
// hard permit. Connections from outside to Xray get an ordinary one, so
// that firewalls keep guarding its inbounds.
self := condition(&fwpmConditionALEAppID, fwpByteBlobType, uintptr(unsafe.Pointer(appID)))
dns53 := condition(&fwpmConditionIPRemotePort, fwpUint16, 53)
// DNS goes through the TUN when its local address is the TUN's, and it
// also leaves, or arrives, through the TUN. The local address alone
// decides by default, but with weak host sending or receiving enabled,
// packets of the TUN's address can use other interfaces. (The next hop,
// the interface replies would leave by, is not known for arriving ones.)
onTUN := func(field *windows.GUID) fwpmFilterCondition0 {
return condition(field, fwpUint64, uintptr(unsafe.Pointer(tunLUID)))
}
out := []fwpmFilterCondition0{dns53, onTUN(&fwpmConditionIPLocalInterface), onTUN(&fwpmConditionIPNexthopInterface)}
in := []fwpmFilterCondition0{dns53, onTUN(&fwpmConditionIPLocalInterface), onTUN(&fwpmConditionIPArrivalInterface)}
for _, layer := range []struct {
key *windows.GUID
selfFlags uint32
throughTUN []fwpmFilterCondition0
}{
{&fwpmLayerALEAuthConnectV4, fwpmFilterFlagClearActionRight, out},
{&fwpmLayerALEAuthRecvAcceptV4, 0, in},
{&fwpmLayerALEAuthConnectV6, fwpmFilterFlagClearActionRight, out},
{&fwpmLayerALEAuthRecvAcceptV6, 0, in},
} {
if err := add(layer.key, "permit Xray", layer.selfFlags, fwpActionPermit, 4, self); err != nil {
return err
}
if dns {
if err := add(layer.key, "permit DNS through the TUN", 0, fwpActionPermit, 3, layer.throughTUN...); err != nil {
return err
}
if err := add(layer.key, "block DNS", 0, fwpActionBlock, 2, dns53); err != nil {
return err
}
}
}
// Since Windows 11 and Server 2022 (build 20348), the DNS Client service
// may also send the queries for an interface's servers over HTTPS or TLS,
// out through that interface and to any port. So there it may only
// connect through the TUN, except for mDNS and LLMNR, which stay on the
// local link (over an IP version only while it is not blocked altogether).
// Earlier versions only query port 53, and may run the service in one
// process with others, which the filters would catch as well. Like
// Windows Firewall's rules for it, they recognize the service by its SID,
// which Windows puts in the token of its process: the security descriptor
// grants that SID the right to match (FWP_ACTRL_MATCH_FILTER, CC in SDDL).
if _, _, build := windows.RtlGetNtVersionNumbers(); dns && build >= 20348 {
sd, err := windows.SecurityDescriptorFromString("O:SYG:SYD:(A;;CCRC;;;" + dnsClientSID + ")")
if err != nil {
return err
}
sdBlob := &fwpByteBlob{size: sd.Length(), data: (*byte)(unsafe.Pointer(sd))}
pinner.Pin(sdBlob) // the condition only holds it as uintptr
dnsClient := condition(&fwpmConditionALEUserID, fwpSecurityDescriptorType, uintptr(unsafe.Pointer(sdBlob)))
// Conditions on the same field match when any of them does.
mdnsLLMNR := []fwpmFilterCondition0{dnsClient, condition(&fwpmConditionIPRemotePort, fwpUint16, 5353), condition(&fwpmConditionIPRemotePort, fwpUint16, 5355)}
for _, layer := range []struct {
key *windows.GUID
localLink bool
}{
{&fwpmLayerALEAuthConnectV4, !ipv4},
{&fwpmLayerALEAuthConnectV6, !ipv6},
} {
if err := add(layer.key, "permit the DNS Client service through the TUN", 0, fwpActionPermit, 3, dnsClient, onTUN(&fwpmConditionIPLocalInterface), onTUN(&fwpmConditionIPNexthopInterface)); err != nil {
return err
}
if layer.localLink {
if err := add(layer.key, "permit the DNS Client service's mDNS and LLMNR", 0, fwpActionPermit, 3, mdnsLLMNR...); err != nil {
return err
}
}
if err := add(layer.key, "block the DNS Client service", 0, fwpActionBlock, 2, dnsClient); err != nil {
return err
}
}
}
// Both directions: replies to a connection accepted from outside would
// leave through the physical link as well.
loopback := fwpmFilterCondition0{
fieldKey: fwpmConditionFlags,
matchType: fwpMatchFlagsAllSet,
conditionValue: fwpValue0{typ: fwpUint32, value: fwpConditionFlagIsLoopback},
}
if ipv4 {
// DHCP keeps the addresses of the other interfaces, which Xray's own
// connections use.
dhcp := []fwpmFilterCondition0{
condition(&fwpmConditionIPProtocol, fwpUint8, windows.IPPROTO_UDP),
condition(&fwpmConditionIPLocalPort, fwpUint16, 68),
condition(&fwpmConditionIPRemotePort, fwpUint16, 67),
}
for _, layer := range []*windows.GUID{&fwpmLayerALEAuthConnectV4, &fwpmLayerALEAuthRecvAcceptV4} {
if err := add(layer, "permit IPv4 loopback", 0, fwpActionPermit, 1, loopback); err != nil {
return err
}
if err := add(layer, "permit DHCP", 0, fwpActionPermit, 1, dhcp...); err != nil {
return err
}
if err := add(layer, "block IPv4", 0, fwpActionBlock, 0); err != nil {
return err
}
}
}
if ipv6 {
// Neighbor and multicast listener discovery, ICMPv6 130-137 and 143,
// whose type and code sit where the local and remote port are.
discovery := []fwpmFilterCondition0{condition(&fwpmConditionIPProtocol, fwpUint8, windows.IPPROTO_ICMPV6)}
for _, typ := range []uintptr{130, 131, 132, 133, 134, 135, 136, 137, 143} {
discovery = append(discovery, condition(&fwpmConditionIPLocalPort, fwpUint16, typ))
}
discovery = append(discovery, condition(&fwpmConditionIPRemotePort, fwpUint16, 0))
dhcpv6 := []fwpmFilterCondition0{
condition(&fwpmConditionIPProtocol, fwpUint8, windows.IPPROTO_UDP),
condition(&fwpmConditionIPLocalPort, fwpUint16, 546),
condition(&fwpmConditionIPRemotePort, fwpUint16, 547),
}
for _, direction := range []struct {
layer *windows.GUID
dhcpv6 []fwpmFilterCondition0
}{
// The client sends to the servers' multicast address, and they
// answer from their own.
{&fwpmLayerALEAuthConnectV6, slices.Concat(dhcpv6, []fwpmFilterCondition0{condition(&fwpmConditionIPRemoteAddress, fwpByteArray16Type, uintptr(unsafe.Pointer(&ipv6AllDHCPv6Servers)))})},
{&fwpmLayerALEAuthRecvAcceptV6, dhcpv6},
} {
if err := add(direction.layer, "permit IPv6 loopback", 0, fwpActionPermit, 1, loopback); err != nil {
return err
}
if err := add(direction.layer, "permit IPv6 neighbor and multicast listener discovery", 0, fwpActionPermit, 1, discovery...); err != nil {
return err
}
if err := add(direction.layer, "permit DHCPv6", 0, fwpActionPermit, 1, direction.dhcpv6...); err != nil {
return err
}
if err := add(direction.layer, "block IPv6", 0, fwpActionBlock, 0); err != nil {
return err
}
}
}
return nil
}
func addFilter(engine windows.Handle, sublayer, layer *windows.GUID, name string, flags, action uint32, weight uint8, conditions ...fwpmFilterCondition0) error {
filter := fwpmFilter0{
displayData: fwpmDisplayData0{name: utf16Ptr(name)},
flags: flags,
layerKey: *layer,
subLayerKey: *sublayer,
weight: fwpValue0{typ: fwpUint8, value: uintptr(weight)},
numFilterConditions: uint32(len(conditions)),
action: fwpmAction0{typ: action},
}
if len(conditions) > 0 {
filter.filterCondition = &conditions[0]
}
if err := fwpmResult(procFwpmFilterAdd0.Call(uintptr(engine), uintptr(unsafe.Pointer(&filter)), 0, 0)); err != nil {
return errors.New("FwpmFilterAdd0 failed for ", name).Base(err)
}
return nil
}
// dnsOutsideTUN returns the servers outside all of prefixes, the TUN's own
// subnets and routes: queries to them cannot go through the TUN.
func dnsOutsideTUN(servers []netip.Addr, prefixes []netip.Prefix) []netip.Addr {
var outside []netip.Addr
for _, server := range servers {
server = server.Unmap()
if !slices.ContainsFunc(prefixes, func(p netip.Prefix) bool { return p.Contains(server) }) {
outside = append(outside, server)
}
}
return outside
}
// flushDNSCache drops the answers Windows cached so far, like ipconfig
// /flushdns, so that names get resolved again with the current DNS setup.
func flushDNSCache() error {
if err := procDnsFlushResolverCache.Find(); err != nil {
return err
}
if r, _, err := procDnsFlushResolverCache.Call(); r == 0 {
return err
}
return nil
}
+206
View File
@@ -0,0 +1,206 @@
//go:build windows
package tun
import (
"context"
go_errors "errors"
"net"
"net/netip"
"slices"
"testing"
"unsafe"
"github.com/xtls/xray-core/transport/internet"
"golang.org/x/sys/windows"
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
)
// The WFP structures are handed to fwpuclnt.dll as they are, so their layout
// has to match what MSVC produces for 64-bit and for 32-bit Windows.
func TestWFPStructLayout(t *testing.T) {
check := func(name string, got, want64, want32 []uintptr) {
t.Helper()
want := want32
if unsafe.Sizeof(uintptr(0)) == 8 {
want = want64
}
if !slices.Equal(got, want) {
t.Errorf("%s: size and offsets are %v, want %v", name, got, want)
}
}
var blob fwpByteBlob
check("FWP_BYTE_BLOB",
[]uintptr{unsafe.Sizeof(blob), unsafe.Offsetof(blob.data)},
[]uintptr{16, 8}, []uintptr{8, 4})
var value fwpValue0
check("FWP_VALUE0",
[]uintptr{unsafe.Sizeof(value), unsafe.Offsetof(value.value)},
[]uintptr{16, 8}, []uintptr{8, 4})
var display fwpmDisplayData0
check("FWPM_DISPLAY_DATA0",
[]uintptr{unsafe.Sizeof(display), unsafe.Offsetof(display.description)},
[]uintptr{16, 8}, []uintptr{8, 4})
var action fwpmAction0
check("FWPM_ACTION0",
[]uintptr{unsafe.Sizeof(action), unsafe.Offsetof(action.filterType)},
[]uintptr{20, 4}, []uintptr{20, 4})
var cond fwpmFilterCondition0
check("FWPM_FILTER_CONDITION0",
[]uintptr{unsafe.Sizeof(cond), unsafe.Offsetof(cond.matchType), unsafe.Offsetof(cond.conditionValue)},
[]uintptr{40, 16, 24}, []uintptr{28, 16, 20})
var session fwpmSession0
check("FWPM_SESSION0",
[]uintptr{
unsafe.Sizeof(session), unsafe.Offsetof(session.displayData), unsafe.Offsetof(session.flags),
unsafe.Offsetof(session.txnWaitTimeoutInMSec), unsafe.Offsetof(session.processID), unsafe.Offsetof(session.sid),
unsafe.Offsetof(session.username), unsafe.Offsetof(session.kernelMode),
},
[]uintptr{72, 16, 32, 36, 40, 48, 56, 64},
[]uintptr{48, 16, 24, 28, 32, 36, 40, 44})
var sublayer fwpmSublayer0
check("FWPM_SUBLAYER0",
[]uintptr{
unsafe.Sizeof(sublayer), unsafe.Offsetof(sublayer.displayData), unsafe.Offsetof(sublayer.flags),
unsafe.Offsetof(sublayer.providerKey), unsafe.Offsetof(sublayer.providerData), unsafe.Offsetof(sublayer.weight),
},
[]uintptr{72, 16, 32, 40, 48, 64},
[]uintptr{44, 16, 24, 28, 32, 40})
var filter fwpmFilter0
check("FWPM_FILTER0",
[]uintptr{
unsafe.Sizeof(filter), unsafe.Offsetof(filter.displayData), unsafe.Offsetof(filter.flags),
unsafe.Offsetof(filter.providerKey), unsafe.Offsetof(filter.providerData), unsafe.Offsetof(filter.layerKey),
unsafe.Offsetof(filter.subLayerKey), unsafe.Offsetof(filter.weight), unsafe.Offsetof(filter.numFilterConditions),
unsafe.Offsetof(filter.filterCondition), unsafe.Offsetof(filter.action), unsafe.Offsetof(filter.providerContextKey),
unsafe.Offsetof(filter.reserved), unsafe.Offsetof(filter.filterID), unsafe.Offsetof(filter.effectiveWeight),
},
[]uintptr{200, 16, 32, 40, 48, 64, 80, 96, 112, 120, 128, 152, 168, 176, 184},
[]uintptr{152, 16, 24, 28, 32, 40, 56, 72, 80, 84, 88, 112, 128, 136, 144})
}
// TestLeakFiltersAccepted has WFP validate the filters by adding them inside a
// transaction that is then aborted, which leaves the system untouched. Adding
// filters requires an elevated process.
func TestLeakFiltersAccepted(t *testing.T) {
skipUnlessElevated := func(err error) {
t.Helper()
if go_errors.Is(err, windows.ERROR_ACCESS_DENIED) {
t.Skipf("WFP filters can only be added by an elevated process: %v", err)
}
t.Fatal(err)
}
engine, err := openWFPEngine()
if err != nil {
skipUnlessElevated(err)
}
defer closeWFPEngine(engine)
if err := fwpmResult(procFwpmTransactionBegin0.Call(uintptr(engine), 0)); err != nil {
skipUnlessElevated(err)
}
defer procFwpmTransactionAbort0.Call(uintptr(engine))
// Any interface stands in for the TUN; the loopback one always exists.
loopback, err := winipcfg.LUIDFromIndex(1)
if err != nil {
t.Fatal(err)
}
if err := addLeakFilters(engine, loopback, true, true, true); err != nil {
skipUnlessElevated(err)
}
}
func TestDNSClientSID(t *testing.T) {
sid, _, _, err := windows.LookupSID("", `NT SERVICE\Dnscache`)
if err != nil {
t.Fatal(err)
}
if sid.String() != dnsClientSID {
t.Errorf(`NT SERVICE\Dnscache is %v, not %v`, sid, dnsClientSID)
}
}
func TestDNSOutsideTUN(t *testing.T) {
prefixes := []netip.Prefix{
netip.MustParsePrefix("198.51.100.1/30"), // gateway, not masked
netip.MustParsePrefix("203.0.113.0/24"), // route
}
servers := []netip.Addr{
netip.MustParseAddr("198.51.100.2"),
netip.MustParseAddr("203.0.113.53"),
netip.MustParseAddr("::ffff:203.0.113.54"),
netip.MustParseAddr("8.8.8.8"),
netip.MustParseAddr("2001:db8::53"),
}
want := []netip.Addr{netip.MustParseAddr("8.8.8.8"), netip.MustParseAddr("2001:db8::53")}
if got := dnsOutsideTUN(servers, prefixes); !slices.Equal(got, want) {
t.Errorf("got %v, want %v", got, want)
}
}
func TestResolveOnOwn(t *testing.T) {
internet.SkipDNSServers([]netip.Addr{netip.MustParseAddr("::ffff:203.0.113.53")})
t.Cleanup(func() { internet.SkipDNSServers(nil) })
preferGo, dial := net.DefaultResolver.PreferGo, net.DefaultResolver.Dial
saved := resolveOnOwn()
t.Cleanup(saved.restore)
if !net.DefaultResolver.PreferGo || net.DefaultResolver.Dial == nil {
t.Fatal("net.DefaultResolver is unchanged")
}
if _, err := net.DefaultResolver.Dial(context.Background(), "udp", "203.0.113.53:53"); err == nil {
t.Error("the TUN's DNS server was not skipped")
}
conn, err := net.DefaultResolver.Dial(context.Background(), "udp", "127.0.0.1:53")
if err != nil {
t.Fatal(err)
}
conn.Close()
saved.restore()
if net.DefaultResolver.PreferGo != preferGo || (net.DefaultResolver.Dial == nil) != (dial == nil) {
t.Error("net.DefaultResolver is not restored")
}
}
// TestTunOnlyDNS checks that a DNS server another interface uses as well is
// not skipped, while one of the TUN alone is.
func TestTunOnlyDNS(t *testing.T) {
adapters, err := winipcfg.GetAdaptersAddresses(windows.AF_UNSPEC, winipcfg.GAAFlagIncludeGateways)
if err != nil {
t.Fatal(err)
}
var other netip.Addr
for _, adapter := range adapters {
if adapter.OperStatus == winipcfg.IfOperStatusUp && adapter.FirstGatewayAddress != nil && adapter.FirstDNSServerAddress != nil {
other, _ = netip.AddrFromSlice(adapter.FirstDNSServerAddress.Address.IP())
other = other.Unmap()
break
}
}
if !other.IsValid() {
t.Skip("no interface with a gateway and a DNS server")
}
tunOnly := netip.MustParseAddr("203.0.113.53")
// LUID 0 is no interface, so every one counts as another.
got, err := tunOnlyDNS(0, []netip.Addr{other, tunOnly})
if err != nil {
t.Fatal(err)
}
if !slices.Equal(got, []netip.Addr{tunOnly}) {
t.Errorf("got %v, want [%v]", got, tunOnly)
}
}
func TestFlushDNSCache(t *testing.T) {
if err := flushDNSCache(); err != nil {
t.Fatal(err)
}
}
+32
View File
@@ -0,0 +1,32 @@
package internet
import (
"net/netip"
"slices"
"sync/atomic"
)
var skippedDNSServers atomic.Pointer[[]netip.Addr]
// SkipDNSServers has the queries Xray sends to the system's DNS servers on its
// own, like those of localdns, skip servers until it is called again. The DNS
// servers of a TUN are only meant for what goes through it: queried by Xray
// itself they lead back into it, or nowhere.
func SkipDNSServers(servers []netip.Addr) {
skipped := make([]netip.Addr, len(servers))
for i, server := range servers {
skipped[i] = server.Unmap()
}
skippedDNSServers.Store(&skipped)
}
// IsSkippedDNSServer reports whether address, a DNS server as host:port, is to
// be skipped, see SkipDNSServers.
func IsSkippedDNSServer(address string) bool {
skipped := skippedDNSServers.Load()
if skipped == nil {
return false
}
server, err := netip.ParseAddrPort(address)
return err == nil && slices.Contains(*skipped, server.Addr().Unmap())
}
+27
View File
@@ -0,0 +1,27 @@
package internet_test
import (
"net/netip"
"testing"
"github.com/xtls/xray-core/transport/internet"
)
func TestSkipDNSServers(t *testing.T) {
internet.SkipDNSServers([]netip.Addr{netip.MustParseAddr("::ffff:203.0.113.53"), netip.MustParseAddr("2001:db8::53")})
t.Cleanup(func() { internet.SkipDNSServers(nil) })
for address, want := range map[string]bool{
"203.0.113.53:53": true,
"[2001:db8::53]:53": true,
"198.51.100.53:53": false,
"localhost:53": false,
} {
if got := internet.IsSkippedDNSServer(address); got != want {
t.Errorf("IsSkippedDNSServer(%q) = %v, want %v", address, got, want)
}
}
internet.SkipDNSServers(nil)
if internet.IsSkippedDNSServer("203.0.113.53:53") {
t.Error("still skipped after SkipDNSServers(nil)")
}
}
+2 -10
View File
@@ -5,7 +5,6 @@ import (
"encoding/binary"
"io"
"sync"
"sync/atomic"
"time"
"github.com/apernet/quic-go"
@@ -21,10 +20,8 @@ type interConn struct {
local net.Addr
remote net.Addr
client bool
user *protocol.MemoryUser
closeOnce sync.Once
aliveTCP *atomic.Int64
client bool
user *protocol.MemoryUser
}
func (c *interConn) User() *protocol.MemoryUser {
@@ -49,11 +46,6 @@ func (c *interConn) Write(b []byte) (int, error) {
func (c *interConn) Close() error {
c.stream.CancelRead(0)
if c.aliveTCP != nil {
c.closeOnce.Do(func() {
c.aliveTCP.Add(-1)
})
}
return c.stream.Close()
}
+7 -33
View File
@@ -9,7 +9,6 @@ import (
"runtime"
"strconv"
"sync"
"sync/atomic"
"time"
"github.com/apernet/quic-go"
@@ -36,12 +35,10 @@ type client struct {
finalMask *finalmask.FinalMask
quicParams *internet.QuicParams
conn *quic.Conn
tr *quic.Transport
pktConn net.PacketConn
udpSM *udpSessionManager
aliveTCP *atomic.Int64
lastAlive time.Time
conn *quic.Conn
tr *quic.Transport
pktConn net.PacketConn
udpSM *udpSessionManager
}
func (c *client) status() status {
@@ -238,15 +235,12 @@ func (c *client) tcp(ctx context.Context) (stat.Connection, error) {
return nil, err
}
c.lastAlive = time.Now()
c.aliveTCP.Add(1)
return &interConn{
stream: stream,
local: c.conn.LocalAddr(),
remote: c.conn.RemoteAddr(),
client: true,
aliveTCP: c.aliveTCP,
client: true,
}, nil
}
@@ -264,28 +258,10 @@ func (c *client) udp(ctx context.Context) (stat.Connection, error) {
func (c *client) clean() {
c.Lock()
defer c.Unlock()
switch c.status() {
case StatusInactive:
if c.status() == StatusInactive {
c.close()
return
case StatusNull:
return
}
var udpSessions int
if c.udpSM != nil {
c.udpSM.RLock()
udpSessions = len(c.udpSM.m)
c.udpSM.RUnlock()
}
if udpSessions == 0 && c.aliveTCP.Load() == 0 {
if c.lastAlive.Add(net.ConnIdleTimeout).Before(time.Now()) {
c.close()
return
}
} else {
c.lastAlive = time.Now()
}
c.Unlock()
}
type dialerConf struct {
@@ -345,8 +321,6 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me
socketConfig: streamSettings.SocketSettings,
finalMask: streamSettings.FinalMask,
quicParams: streamSettings.QuicParams,
aliveTCP: &atomic.Int64{},
lastAlive: time.Now(),
}
manager.m[dialerConf{dest, streamSettings}] = c
}