mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-11 00:38:24 +03:00
Xray-core: Add Lua script for dns and routing (#6823)
https://github.com/XTLS/Xray-core/pull/6823#issuecomment-5843754759 https://github.com/XTLS/Xray-core/pull/6823#issuecomment-5861456450 https://github.com/XTLS/Xray-core/pull/6823#issuecomment-6093596069
This commit is contained in:
@@ -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
@@ -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" +
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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" +
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
@@ -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())
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user