mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-02 13:56:39 +00:00
Compare commits
13
Commits
windows-tun-fix
...
lua
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5d1d8200d9 | ||
|
|
2440f53cdd | ||
|
|
1c52c65872 | ||
|
|
5724db08f4 | ||
|
|
3d3306503d | ||
|
|
5e1bb92b98 | ||
|
|
459301d42e | ||
|
|
3982028a9c | ||
|
|
70b8e9a61d | ||
|
|
219f758060 | ||
|
|
9628003594 | ||
|
|
72d9ab50b9 | ||
|
|
235843c5d2 |
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
+23
-4
@@ -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,6 +266,8 @@ 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"`
|
||||
// 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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -257,6 +270,9 @@ func (s *DNS) LookupIP(domain string, option dns.IPOption) ([]net.IP, uint32, er
|
||||
}
|
||||
|
||||
// Name servers lookup
|
||||
if s.script != nil {
|
||||
return s.script.query(domain, option)
|
||||
}
|
||||
if s.enableParallelQuery {
|
||||
return s.parallelQuery(domain, option)
|
||||
} else {
|
||||
|
||||
+201
@@ -0,0 +1,201 @@
|
||||
package dns
|
||||
|
||||
import (
|
||||
"context"
|
||||
"math"
|
||||
"strings"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
featureDNS "github.com/xtls/xray-core/features/dns"
|
||||
"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 {
|
||||
serverList := L.NewTable()
|
||||
for i, client := range servers {
|
||||
server := L.NewTable()
|
||||
|
||||
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)
|
||||
}
|
||||
addresses := L.NewUserData()
|
||||
addresses.Value = ips
|
||||
L.Push(addresses)
|
||||
L.Push(lua.LNumber(ttl))
|
||||
if err != nil {
|
||||
ud := L.NewUserData()
|
||||
ud.Value = err
|
||||
L.Push(ud)
|
||||
} else {
|
||||
L.Push(lua.LNil)
|
||||
}
|
||||
return 3
|
||||
}))
|
||||
serverList.RawSetInt(i+1, server)
|
||||
}
|
||||
|
||||
module := L.NewTable()
|
||||
if servers != nil {
|
||||
module.RawSetString("Servers", serverList)
|
||||
}
|
||||
if client != nil {
|
||||
module.RawSetString("Query", newLuaClientQuery(L, client))
|
||||
}
|
||||
L.Push(module)
|
||||
return 1
|
||||
})
|
||||
}
|
||||
|
||||
func newLuaClientQuery(L *lua.LState, client featureDNS.Client) *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)
|
||||
addresses := L.NewUserData()
|
||||
addresses.Value = ips
|
||||
L.Push(addresses)
|
||||
L.Push(lua.LNumber(ttl))
|
||||
if err != nil {
|
||||
ud := L.NewUserData()
|
||||
ud.Value = err
|
||||
L.Push(ud)
|
||||
} else {
|
||||
L.Push(lua.LNil)
|
||||
}
|
||||
return 3
|
||||
})
|
||||
}
|
||||
|
||||
// CallLuaHook invokes HandleDNSQuery in the supplied state.
|
||||
// Returned slices and IP bytes may share storage with DNS caches or matcher inputs.
|
||||
func (s *DNS) CallLuaHook(L *lua.LState, ctx context.Context, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||
previous, top := L.Context(), L.GetTop()
|
||||
L.SetContext(ctx)
|
||||
defer func() {
|
||||
L.SetTop(top)
|
||||
if previous == nil {
|
||||
L.RemoveContext()
|
||||
} else {
|
||||
L.SetContext(previous)
|
||||
}
|
||||
}()
|
||||
fn := L.GetGlobal("HandleDNSQuery")
|
||||
if fn.Type() != lua.LTFunction {
|
||||
return nil, 0, errors.New("DNS script must define HandleDNSQuery(domain, ipv4, ipv6, fake)")
|
||||
}
|
||||
if err := 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)); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
return readLuaDNSResult(L.Get(-3), L.Get(-2), L.Get(-1))
|
||||
}
|
||||
|
||||
func readLuaDNSResult(addresses, ttlValue, errorValue lua.LValue) ([]net.IP, uint32, error) {
|
||||
if errorValue != lua.LNil {
|
||||
if ud, ok := errorValue.(*lua.LUserData); ok {
|
||||
if err, ok := ud.Value.(error); ok {
|
||||
return nil, 0, err
|
||||
}
|
||||
}
|
||||
if s, ok := errorValue.(lua.LString); ok {
|
||||
return nil, 0, errors.New(string(s))
|
||||
}
|
||||
return nil, 0, errors.New("DNS script error must be an error or string")
|
||||
}
|
||||
ttl, ok := ttlValue.(lua.LNumber)
|
||||
if !ok || ttl < 0 || ttl > math.MaxUint32 || math.Trunc(float64(ttl)) != float64(ttl) {
|
||||
return nil, 0, errors.New("DNS script returned invalid TTL")
|
||||
}
|
||||
if addresses == lua.LNil {
|
||||
return nil, 0, featureDNS.ErrEmptyResponse
|
||||
}
|
||||
ud, ok := addresses.(*lua.LUserData)
|
||||
if !ok {
|
||||
return nil, 0, errors.New("DNS script IPs must be native IP slice userdata")
|
||||
}
|
||||
ips, ok := ud.Value.([]net.IP)
|
||||
if !ok {
|
||||
return nil, 0, errors.New("DNS script IPs must be native IP slice userdata")
|
||||
}
|
||||
if len(ips) == 0 {
|
||||
return nil, 0, featureDNS.ErrEmptyResponse
|
||||
}
|
||||
return ips, uint32(ttl), nil
|
||||
}
|
||||
@@ -0,0 +1,296 @@
|
||||
package dns
|
||||
|
||||
import (
|
||||
"context"
|
||||
go_errors "errors"
|
||||
"math"
|
||||
"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 TestReadLuaDNSResult(t *testing.T) {
|
||||
L := lua.NewState()
|
||||
defer L.Close()
|
||||
want := []net.IP{net.ParseIP("8.8.8.8"), {127, 0, 0, 1}, net.ParseIP("::1")}
|
||||
addresses := L.NewUserData()
|
||||
addresses.Value = want
|
||||
ips, ttl, err := readLuaDNSResult(addresses, lua.LNumber(45), lua.LNil)
|
||||
if err != nil || ttl != 45 || len(ips) != len(want) {
|
||||
t.Fatalf("readLuaDNSResult() = %v, TTL %d, %v", ips, ttl, err)
|
||||
}
|
||||
for i := range want {
|
||||
if !ips[i].Equal(want[i]) {
|
||||
t.Fatalf("IP %d = %v, want %v", i, ips[i], want[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadLuaDNSResultValidation(t *testing.T) {
|
||||
L := lua.NewState()
|
||||
defer L.Close()
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
change func(*[3]lua.LValue)
|
||||
want string
|
||||
}{
|
||||
{"fractional TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(1.5) }, "invalid TTL"},
|
||||
{"oversized TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(4294967296) }, "invalid TTL"},
|
||||
{"negative TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(-1) }, "invalid TTL"},
|
||||
{"NaN TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(math.NaN()) }, "invalid TTL"},
|
||||
{"missing TTL", func(v *[3]lua.LValue) { v[1] = lua.LNil }, "invalid TTL"},
|
||||
{"string IPs", func(v *[3]lua.LValue) { v[0] = lua.LString("127.0.0.1") }, "native IP slice"},
|
||||
{"wrong userdata", func(v *[3]lua.LValue) { v[0].(*lua.LUserData).Value = net.ParseIP("127.0.0.1") }, "native IP slice"},
|
||||
{"script error", func(v *[3]lua.LValue) { v[2] = lua.LString("blocked by script") }, "blocked by script"},
|
||||
{"invalid error", func(v *[3]lua.LValue) { v[2] = lua.LTrue }, "error or string"},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
addresses := L.NewUserData()
|
||||
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
|
||||
values := [3]lua.LValue{addresses, lua.LNumber(60), lua.LNil}
|
||||
tc.change(&values)
|
||||
_, _, err := readLuaDNSResult(values[0], values[1], values[2])
|
||||
if err == nil || !strings.Contains(err.Error(), tc.want) {
|
||||
t.Fatalf("readLuaDNSResult error = %v, want %q", err, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
addresses := L.NewUserData()
|
||||
addresses.Value = []net.IP(nil)
|
||||
for _, empty := range []lua.LValue{addresses, lua.LNil} {
|
||||
if _, _, err := readLuaDNSResult(empty, lua.LNumber(0), lua.LNil); !go_errors.Is(err, featureDNS.ErrEmptyResponse) {
|
||||
t.Fatalf("empty result error = %v, want ErrEmptyResponse", err)
|
||||
}
|
||||
}
|
||||
wantErr := go_errors.New("upstream failed")
|
||||
errorValue := L.NewUserData()
|
||||
errorValue.Value = wantErr
|
||||
if _, _, err := readLuaDNSResult(lua.LNil, lua.LNil, errorValue); err != wantErr {
|
||||
t.Fatalf("upstream error = %v, want original error %v", err, wantErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallLuaHookCancellation(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.WithTimeout(context.Background(), 50*time.Millisecond)
|
||||
defer cancel()
|
||||
_, _, err := (&DNS{}).CallLuaHook(L, ctx, "example.com", featureDNS.IPOption{IPv4Enable: true})
|
||||
if err == nil {
|
||||
t.Fatal("CallLuaHook did not stop after context cancellation")
|
||||
}
|
||||
if L.Context() != nil {
|
||||
t.Fatal("CallLuaHook left the canceled context on the Lua state")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallLuaHookNormalizesDomain(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)
|
||||
}
|
||||
s := &DNS{}
|
||||
if _, _, err := s.CallLuaHook(L, context.Background(), "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallLuaHookRestoresState(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
body string
|
||||
wantErr bool
|
||||
}{
|
||||
{"success", `return ips, 60`, false},
|
||||
{"error", `error("failed")`, true},
|
||||
} {
|
||||
t.Run(tc.name, func(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() " + tc.body + " end"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
previous, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
L.SetContext(previous)
|
||||
L.Push(lua.LTrue)
|
||||
_, _, err := (&DNS{}).CallLuaHook(L, context.Background(), "example.com", featureDNS.IPOption{IPv4Enable: true})
|
||||
if (err != nil) != tc.wantErr {
|
||||
t.Fatalf("hook error = %v, want error %t", err, tc.wantErr)
|
||||
}
|
||||
if L.Context() != previous || L.GetTop() != 1 || L.Get(1) != lua.LTrue {
|
||||
t.Fatal("hook did not restore the previous context and 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(matcher:AnyMatch(ips))
|
||||
local matched = matcher:FilterIPs(ips)
|
||||
return matched, ttl, err
|
||||
end
|
||||
`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got, ttl, err := server.CallLuaHook(L, context.Background(), "example.com", option)
|
||||
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())
|
||||
defer L.RemoveContext()
|
||||
geodata.RegisterLua(L)
|
||||
want := []net.IP{{127, 0, 0, 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))
|
||||
`); 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())
|
||||
defer L.RemoveContext()
|
||||
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)
|
||||
`); 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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
// BenchmarkLuaDNSHookCall isolates a preloaded Lua hook and its server:Query bridge.
|
||||
// The direct case measures the same DNS client without Lua.
|
||||
func BenchmarkLuaDNSHookCall(b *testing.B) {
|
||||
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{clients: []*Client{client}}
|
||||
L := lua.NewState()
|
||||
defer L.Close()
|
||||
server.RegisterLua(L)
|
||||
if err := L.DoString(`
|
||||
local server = require("xray.dns").Servers[1]
|
||||
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
||||
return server:Query(domain, ipv4, ipv6, fake)
|
||||
end
|
||||
`); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
for _, bench := range []struct {
|
||||
name string
|
||||
query func() ([]net.IP, uint32, error)
|
||||
}{
|
||||
{"direct", func() ([]net.IP, uint32, error) { return client.QueryIP(ctx, "example.com", option) }},
|
||||
{"lua_hook", func() ([]net.IP, uint32, error) { return server.CallLuaHook(L, ctx, "example.com", option) }},
|
||||
} {
|
||||
b.Run(bench.name, func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
var ips []net.IP
|
||||
var ttl uint32
|
||||
var err error
|
||||
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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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,73 @@
|
||||
package dns
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/geodata"
|
||||
"github.com/xtls/xray-core/common/log"
|
||||
luamgr "github.com/xtls/xray-core/common/lua"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/features/dns"
|
||||
lua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
const scriptExecutionTimeout = 10 * time.Second
|
||||
|
||||
type scriptEngine struct {
|
||||
dns *DNS
|
||||
pool *luamgr.Pool
|
||||
}
|
||||
|
||||
func newScriptEngine(path string, server *DNS) (*scriptEngine, error) {
|
||||
program, err := luamgr.CompileFile(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
e := &scriptEngine{dns: server}
|
||||
e.pool, err = luamgr.NewPool(server.ctx, func(poolCtx context.Context) (*lua.LState, error) {
|
||||
initCtx, cancel := context.WithTimeout(poolCtx, scriptExecutionTimeout)
|
||||
defer cancel()
|
||||
L, err := program.NewState(initCtx, func(L *lua.LState) {
|
||||
geodata.RegisterLua(L)
|
||||
log.RegisterLua(L)
|
||||
server.RegisterLua(L)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if L.GetGlobal("HandleDNSQuery").Type() != lua.LTFunction {
|
||||
L.Close()
|
||||
return nil, errors.New("DNS script must define HandleDNSQuery(domain, ipv4, ipv6, fake)")
|
||||
}
|
||||
return L, nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
errors.LogInfo(server.ctx, "DNS script initialized from ", path)
|
||||
return e, nil
|
||||
}
|
||||
|
||||
func (e *scriptEngine) close() {
|
||||
e.pool.Close()
|
||||
}
|
||||
|
||||
func (e *scriptEngine) query(domain string, option dns.IPOption) ([]net.IP, uint32, error) {
|
||||
L, err := e.pool.Acquire()
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
reusable := false
|
||||
defer func() {
|
||||
e.pool.Release(L, reusable)
|
||||
}()
|
||||
queryCtx, cancel := context.WithTimeout(e.pool.Context(), scriptExecutionTimeout)
|
||||
defer cancel()
|
||||
ips, ttl, err := e.dns.CallLuaHook(L, queryCtx, domain, option)
|
||||
if err == nil {
|
||||
reusable = true
|
||||
}
|
||||
return ips, ttl, err
|
||||
}
|
||||
@@ -0,0 +1,197 @@
|
||||
package dns
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
featureDNS "github.com/xtls/xray-core/features/dns"
|
||||
)
|
||||
|
||||
type geoIPScriptNameServer struct {
|
||||
name string
|
||||
answers map[string]net.IP
|
||||
ttl uint32
|
||||
calls int
|
||||
}
|
||||
|
||||
func (s *geoIPScriptNameServer) Name() string { return s.name }
|
||||
func (s *geoIPScriptNameServer) IsDisableCache() bool { return true }
|
||||
|
||||
func (s *geoIPScriptNameServer) QueryIP(ctx context.Context, domain string, _ featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
s.calls++
|
||||
ip, ok := s.answers[domain]
|
||||
if !ok {
|
||||
return nil, 0, featureDNS.ErrEmptyResponse
|
||||
}
|
||||
return []net.IP{ip}, s.ttl, nil
|
||||
}
|
||||
|
||||
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 := &geoIPScriptNameServer{
|
||||
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 := &geoIPScriptNameServer{
|
||||
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 TestDNSScriptHookErrorAndFakeDNSOption(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)
|
||||
if domain == "bad.example" then error("script failure") end
|
||||
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 := &geoIPScriptNameServer{
|
||||
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("bad.example", option); err == nil || !strings.Contains(err.Error(), "script failure") {
|
||||
t.Fatalf("hook failure = %v, want script failure", err)
|
||||
}
|
||||
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)
|
||||
}
|
||||
}
|
||||
+12
-2
@@ -587,6 +587,8 @@ 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"`
|
||||
// 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
|
||||
}
|
||||
@@ -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,206 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"context"
|
||||
"runtime"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"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.NewTable()
|
||||
|
||||
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 {
|
||||
L.Push(lua.LNil)
|
||||
pushLuaError(L, errors.New("balancer ", tag, " not found"))
|
||||
return 2
|
||||
}
|
||||
outboundTag, err := balancer.PickOutbound()
|
||||
L.Push(lua.LString(outboundTag))
|
||||
pushLuaError(L, err)
|
||||
return 2
|
||||
}))
|
||||
|
||||
module.RawSetString("FindProcess", L.NewFunction(func(L *lua.LState) int {
|
||||
pid, name, path, err := findProcess(checkLuaContext(L), net.FindProcess)
|
||||
L.Push(lua.LNumber(pid))
|
||||
L.Push(lua.LString(name))
|
||||
L.Push(lua.LString(path))
|
||||
pushLuaError(L, err)
|
||||
return 4
|
||||
}))
|
||||
|
||||
L.Push(module)
|
||||
return 1
|
||||
})
|
||||
}
|
||||
|
||||
func registerLuaContext(L *lua.LState) {
|
||||
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 {
|
||||
L.Push(lua.LString(value))
|
||||
} else {
|
||||
L.Push(lua.LNil)
|
||||
}
|
||||
return 1
|
||||
}))
|
||||
methods := L.NewTable()
|
||||
L.SetFuncs(methods, map[string]lua.LGFunction{
|
||||
"GetSourceIPs": func(L *lua.LState) int {
|
||||
return pushLuaIPs(L, checkLuaContext(L).GetSourceIPs())
|
||||
},
|
||||
"GetTargetIPs": func(L *lua.LState) int {
|
||||
return pushLuaIPs(L, checkLuaContext(L).GetTargetIPs())
|
||||
},
|
||||
"GetLocalIPs": func(L *lua.LState) int {
|
||||
return pushLuaIPs(L, checkLuaContext(L).GetLocalIPs())
|
||||
},
|
||||
"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
|
||||
}
|
||||
|
||||
func pushLuaIPs(L *lua.LState, ips []net.IP) int {
|
||||
addresses := L.NewUserData()
|
||||
addresses.Value = ips
|
||||
L.Push(addresses)
|
||||
return 1
|
||||
}
|
||||
|
||||
func pushLuaError(L *lua.LState, err error) {
|
||||
if err == nil {
|
||||
L.Push(lua.LNil)
|
||||
return
|
||||
}
|
||||
value := L.NewUserData()
|
||||
value.Value = err
|
||||
L.Push(value)
|
||||
}
|
||||
|
||||
// CallLuaHook invokes HandleRoute in the supplied state.
|
||||
func (r *Router) CallLuaHook(L *lua.LState, ctx context.Context, routeCtx routing.Context) (string, string, error) {
|
||||
previous, top := L.Context(), L.GetTop()
|
||||
L.SetContext(ctx)
|
||||
defer func() {
|
||||
L.SetTop(top)
|
||||
if previous == nil {
|
||||
L.RemoveContext()
|
||||
} else {
|
||||
L.SetContext(previous)
|
||||
}
|
||||
}()
|
||||
fn := L.GetGlobal("HandleRoute")
|
||||
if fn.Type() != lua.LTFunction {
|
||||
return "", "", errors.New("routing script must define HandleRoute(...)")
|
||||
}
|
||||
value := L.NewUserData()
|
||||
value.Value = routeCtx
|
||||
L.SetMetatable(value, L.GetTypeMetatable(luaContextType))
|
||||
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
|
||||
value, lua.LString(routeCtx.GetInboundTag()), lua.LNumber(routeCtx.GetSourcePort()),
|
||||
lua.LNumber(routeCtx.GetTargetPort()), lua.LNumber(routeCtx.GetLocalPort()),
|
||||
lua.LString(routeCtx.GetTargetDomain()), lua.LNumber(routeCtx.GetNetwork()),
|
||||
lua.LString(routeCtx.GetProtocol()), lua.LString(routeCtx.GetUser()),
|
||||
lua.LNumber(routeCtx.GetVlessRoute()), lua.LBool(routeCtx.GetSkipDNSResolve())); err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
return readLuaRouteResult(L.Get(-3), L.Get(-2), L.Get(-1))
|
||||
}
|
||||
|
||||
func readLuaRouteResult(tagValue, ruleValue, errorValue lua.LValue) (string, string, error) {
|
||||
if errorValue != lua.LNil {
|
||||
if value, ok := errorValue.(*lua.LUserData); ok {
|
||||
if err, ok := value.Value.(error); ok {
|
||||
return "", "", err
|
||||
}
|
||||
}
|
||||
if value, ok := errorValue.(lua.LString); ok {
|
||||
return "", "", errors.New(string(value))
|
||||
}
|
||||
return "", "", errors.New("routing script error must be an error or string")
|
||||
}
|
||||
if tagValue == lua.LNil {
|
||||
return "", "", nil
|
||||
}
|
||||
tag, ok := tagValue.(lua.LString)
|
||||
if !ok {
|
||||
return "", "", errors.New("routing script outboundTag must be a string or nil")
|
||||
}
|
||||
if tag == "" {
|
||||
return "", "", nil
|
||||
}
|
||||
var ruleTag string
|
||||
if ruleValue != lua.LNil {
|
||||
value, ok := ruleValue.(lua.LString)
|
||||
if !ok {
|
||||
return "", "", errors.New("routing script ruleTag must be a string")
|
||||
}
|
||||
ruleTag = string(value)
|
||||
}
|
||||
return string(tag), 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,292 @@
|
||||
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) (*Router, *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 r, L
|
||||
}
|
||||
|
||||
func TestLuaRouteBinding(t *testing.T) {
|
||||
r, 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(matcher:AnyMatch(sourceIPs) and matcher:AnyMatch(targetIPs) and matcher:AnyMatch(localIPs))
|
||||
assert(attributes.key == "value" and attributes.missing == nil)
|
||||
assert(not pcall(function() attributes.key = "changed" end))
|
||||
return "out", "rule"
|
||||
end`)
|
||||
|
||||
ctx := newLuaRouteTestContext()
|
||||
tag, rule, err := r.CallLuaHook(L, context.Background(), ctx)
|
||||
if err != nil || tag != "out" || rule != "rule" {
|
||||
t.Fatalf("hook = %q, %q, %v", tag, rule, err)
|
||||
}
|
||||
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 TestLuaRouteResult(t *testing.T) {
|
||||
nativeErr := go_errors.New("native failure")
|
||||
for _, tc := range []struct {
|
||||
name, body, tag, rule, wantErr string
|
||||
native bool
|
||||
}{
|
||||
{name: "route", body: `return "out", "rule"`, tag: "out", rule: "rule"},
|
||||
{name: "no match", body: `return nil`},
|
||||
{name: "empty tag", body: `return ""`},
|
||||
{name: "invalid tag", body: `return 1`, wantErr: "outboundTag"},
|
||||
{name: "invalid rule", body: `return "out", false`, wantErr: "ruleTag"},
|
||||
{name: "string error", body: `return nil, nil, "script failure"`, wantErr: "script failure"},
|
||||
{name: "native error", body: `return nil, nil, nativeError`, native: true},
|
||||
{name: "runtime error", body: `error("runtime failure")`, wantErr: "runtime failure"},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
r, L := newLuaRouterState(t, "function HandleRoute() "+tc.body+" end")
|
||||
value := L.NewUserData()
|
||||
value.Value = nativeErr
|
||||
L.SetGlobal("nativeError", value)
|
||||
previous := context.WithValue(context.Background(), struct{}{}, true)
|
||||
L.SetContext(previous)
|
||||
L.Push(lua.LTrue)
|
||||
|
||||
tag, rule, err := r.CallLuaHook(L, context.Background(), &routing_session.Context{})
|
||||
if tag != tc.tag || rule != tc.rule {
|
||||
t.Fatalf("result = %q, %q, %v", tag, rule, err)
|
||||
}
|
||||
switch {
|
||||
case tc.native:
|
||||
if err != nativeErr {
|
||||
t.Fatalf("error = %v, want original error", err)
|
||||
}
|
||||
case tc.wantErr != "":
|
||||
if err == nil || !strings.Contains(err.Error(), tc.wantErr) {
|
||||
t.Fatalf("error = %v, want %q", err, tc.wantErr)
|
||||
}
|
||||
case err != nil:
|
||||
t.Fatal(err)
|
||||
}
|
||||
if L.Context() != previous || L.GetTop() != 1 || L.Get(1) != lua.LTrue {
|
||||
t.Fatal("hook did not restore the previous context and stack")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLuaRouteCancellation(t *testing.T) {
|
||||
r, L := newLuaRouterState(t, `function HandleRoute() while true do end end`)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
if _, _, err := r.CallLuaHook(L, ctx, &routing_session.Context{}); err == nil {
|
||||
t.Fatal("CallLuaHook did not stop after context cancellation")
|
||||
}
|
||||
if L.Context() != nil || L.GetTop() != 0 {
|
||||
t.Fatal("CallLuaHook did not restore the Lua state")
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkLuaRouteHookCall isolates a preloaded Lua hook and its routing context bridge.
|
||||
// The direct case runs an equivalent native routing rule.
|
||||
func BenchmarkLuaRouteHookCall(b *testing.B) {
|
||||
r := new(Router)
|
||||
if err := r.Init(context.Background(), &Config{Rule: []*RoutingRule{{
|
||||
TargetTag: &RoutingRule_Tag{Tag: "out"},
|
||||
RuleTag: "rule",
|
||||
InboundTag: []string{"in"},
|
||||
Networks: []net.Network{net.Network_TCP},
|
||||
Ip: []*geodata.IPRule{{
|
||||
Value: &geodata.IPRule_Custom{Custom: &geodata.CIDRRule{
|
||||
Cidr: &geodata.CIDR{Ip: []byte{127, 0, 0, 0}, Prefix: 8},
|
||||
}},
|
||||
}},
|
||||
}}}, nil, nil, nil); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
L := lua.NewState()
|
||||
defer L.Close()
|
||||
r.RegisterLua(L)
|
||||
geodata.RegisterLua(L)
|
||||
if err := L.DoString(`
|
||||
local router = require("xray.router")
|
||||
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
|
||||
function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort,
|
||||
targetDomain, network, protocol, user, vlessRoute, skipDNSResolve)
|
||||
if inboundTag == "in" and network == router.NetworkTCP and matcher:AnyMatch(ctx:GetTargetIPs()) then
|
||||
return "out", "rule"
|
||||
end
|
||||
end
|
||||
`); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
routeCtx := newLuaRouteTestContext()
|
||||
for _, benchmark := range []struct {
|
||||
name string
|
||||
route func() (string, string, error)
|
||||
}{
|
||||
{"direct", func() (string, string, error) {
|
||||
route, err := r.PickRoute(routeCtx)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
return route.GetOutboundTag(), route.GetRuleTag(), nil
|
||||
}},
|
||||
{"lua_hook", func() (string, string, error) {
|
||||
return r.CallLuaHook(L, ctx, routeCtx)
|
||||
}},
|
||||
} {
|
||||
b.Run(benchmark.name, func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
var tag, rule string
|
||||
var err error
|
||||
for i := 0; i < b.N; i++ {
|
||||
tag, rule, err = benchmark.route()
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
b.StopTimer()
|
||||
if tag != "out" || rule != "rule" {
|
||||
b.Fatalf("route() = %q, %q; want out, rule", tag, rule)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
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,79 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"context"
|
||||
"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"
|
||||
luamgr "github.com/xtls/xray-core/common/lua"
|
||||
"github.com/xtls/xray-core/features/routing"
|
||||
lua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
const scriptExecutionTimeout = 10 * time.Second
|
||||
|
||||
type scriptEngine struct {
|
||||
router *Router
|
||||
pool *luamgr.Pool
|
||||
}
|
||||
|
||||
func newScriptEngine(path string, router *Router) (*scriptEngine, error) {
|
||||
program, err := luamgr.CompileFile(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
e := &scriptEngine{router: router}
|
||||
e.pool, err = luamgr.NewPool(router.ctx, func(poolCtx context.Context) (*lua.LState, error) {
|
||||
initCtx, cancel := context.WithTimeout(poolCtx, scriptExecutionTimeout)
|
||||
defer cancel()
|
||||
L, err := program.NewState(initCtx, func(L *lua.LState) {
|
||||
geodata.RegisterLua(L)
|
||||
log.RegisterLua(L)
|
||||
router.RegisterLua(L)
|
||||
dns.RegisterLua(L, router.dns)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if L.GetGlobal("HandleRoute").Type() != lua.LTFunction {
|
||||
L.Close()
|
||||
return nil, errors.New("routing script must define HandleRoute(...)")
|
||||
}
|
||||
return L, nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
errors.LogInfo(router.ctx, "routing script initialized from ", path)
|
||||
return e, nil
|
||||
}
|
||||
|
||||
func (e *scriptEngine) close() {
|
||||
e.pool.Close()
|
||||
}
|
||||
|
||||
func (e *scriptEngine) pickRoute(ctx routing.Context) (routing.Route, error) {
|
||||
L, err := e.pool.Acquire()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
reusable := false
|
||||
defer func() {
|
||||
e.pool.Release(L, reusable)
|
||||
}()
|
||||
callCtx, cancel := context.WithTimeout(e.pool.Context(), scriptExecutionTimeout)
|
||||
defer cancel()
|
||||
tag, ruleTag, err := e.router.CallLuaHook(L, callCtx, ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
reusable = true
|
||||
if tag == "" {
|
||||
return nil, common.ErrNoClue
|
||||
}
|
||||
return &Route{Context: ctx, outboundTag: tag, ruleTag: ruleTag}, nil
|
||||
}
|
||||
@@ -0,0 +1,372 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"context"
|
||||
stdnet "net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"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) {
|
||||
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
|
||||
}}
|
||||
r := startLuaRouter(t, `
|
||||
function HandleRoute(ctx, inbound)
|
||||
if inbound == "miss" then return nil end
|
||||
return "lua-out", "lua-rule"
|
||||
end`, 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)
|
||||
if err != nil || route.GetOutboundTag() != "lua-out" || route.GetRuleTag() != "lua-rule" || route.(*Route).Context != ctx {
|
||||
t.Fatalf("route = %v, %v", route, err)
|
||||
}
|
||||
ctx.Inbound.Tag = "miss"
|
||||
if route, err := r.PickRoute(ctx); route != nil || err != common.ErrNoClue {
|
||||
t.Fatalf("miss = %v, %v", route, err)
|
||||
}
|
||||
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 TestRouterScriptStateReuse(t *testing.T) {
|
||||
r := startLuaRouter(t, `
|
||||
local calls = 0
|
||||
function HandleRoute(ctx, inbound)
|
||||
calls = calls + 1
|
||||
if inbound == "miss" then return nil end
|
||||
if inbound == "fail" then error("failed") end
|
||||
return tostring(calls)
|
||||
end`, nil, nil)
|
||||
ctx := newLuaRouteTestContext()
|
||||
pick := func(want string) {
|
||||
t.Helper()
|
||||
route, err := r.PickRoute(ctx)
|
||||
if err != nil || route.GetOutboundTag() != want {
|
||||
t.Fatalf("route = %v, %v, want %q", route, err, want)
|
||||
}
|
||||
}
|
||||
|
||||
pick("1")
|
||||
ctx.Inbound.Tag = "miss"
|
||||
if _, err := r.PickRoute(ctx); err != common.ErrNoClue {
|
||||
t.Fatalf("miss = %v", err)
|
||||
}
|
||||
ctx.Inbound.Tag = "in"
|
||||
pick("3")
|
||||
ctx.Inbound.Tag = "fail"
|
||||
if _, err := r.PickRoute(ctx); err == nil {
|
||||
t.Fatal("script error was ignored")
|
||||
}
|
||||
ctx.Inbound.Tag = "in"
|
||||
pick("1")
|
||||
}
|
||||
|
||||
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")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
package geodata
|
||||
|
||||
import (
|
||||
lua "github.com/yuin/gopher-lua"
|
||||
luar "layeh.com/gopher-luar"
|
||||
)
|
||||
|
||||
// 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.NewTable()
|
||||
|
||||
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
|
||||
}
|
||||
L.Push(luar.New(L, matcher))
|
||||
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
|
||||
}
|
||||
L.Push(luar.New(L, matcher))
|
||||
return 1
|
||||
}))
|
||||
L.Push(module)
|
||||
return 1
|
||||
})
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
package geodata
|
||||
|
||||
import (
|
||||
"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")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
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.NewTable()
|
||||
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 {
|
||||
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 != "" {
|
||||
content.WriteString(filepath.Base(strings.TrimPrefix(caller.Source, "@")))
|
||||
content.WriteString(": ")
|
||||
}
|
||||
}
|
||||
for i := 1; i <= L.GetTop(); i++ {
|
||||
value := L.Get(i)
|
||||
// Use Error() for Go errors in userdata.
|
||||
if ud, ok := value.(*lua.LUserData); ok {
|
||||
if err, ok := ud.Value.(error); ok {
|
||||
content.WriteString(err.Error())
|
||||
continue
|
||||
}
|
||||
}
|
||||
content.WriteString(L.ToStringMeta(value).String())
|
||||
}
|
||||
Record(&GeneralMessage{
|
||||
Severity: severity,
|
||||
Content: content.String(),
|
||||
})
|
||||
return 0
|
||||
}))
|
||||
}
|
||||
L.Push(module)
|
||||
return 1
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,92 @@
|
||||
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) {
|
||||
logHandler.RLock()
|
||||
previous := logHandler.Handler
|
||||
logHandler.RUnlock()
|
||||
t.Cleanup(func() { RegisterHandler(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)
|
||||
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)
|
||||
}
|
||||
|
||||
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: hook"},
|
||||
{Severity_Info, "[Info] <string>: anonymous"},
|
||||
}
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,138 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
const maxIdleStates = 16
|
||||
|
||||
// LStateFactory must initialize a state fully and observe ctx while doing so.
|
||||
// The pool owns any non-nil state it returns, even when it also returns an error.
|
||||
type LStateFactory func(ctx context.Context) (*glua.LState, error)
|
||||
|
||||
// Pool lends each state to one caller at a time. It grows on contention and
|
||||
// keeps up to maxIdleStates idle states until Close. Callers decide whether a
|
||||
// state is reusable.
|
||||
type Pool struct {
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
|
||||
factory LStateFactory
|
||||
idle []*glua.LState
|
||||
|
||||
mu sync.Mutex
|
||||
active sync.WaitGroup
|
||||
closed bool
|
||||
}
|
||||
|
||||
// NewPool initializes one state before returning, so top-level errors surface at startup.
|
||||
func NewPool(ctx context.Context, factory LStateFactory) (*Pool, error) {
|
||||
poolCtx, cancel := context.WithCancel(ctx)
|
||||
|
||||
// Create one state now to catch factory errors at startup.
|
||||
state, err := factory(poolCtx)
|
||||
if err != nil {
|
||||
cancel()
|
||||
if state != nil {
|
||||
state.Close()
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
if state == nil {
|
||||
cancel()
|
||||
return nil, errors.New("Lua state factory returned nil")
|
||||
}
|
||||
if err := poolCtx.Err(); err != nil {
|
||||
state.Close()
|
||||
cancel()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &Pool{ctx: poolCtx, cancel: cancel, factory: factory, idle: []*glua.LState{state}}, nil
|
||||
}
|
||||
|
||||
// Context is cancelled by Close. Query contexts should derive from it.
|
||||
func (p *Pool) Context() context.Context {
|
||||
return p.ctx
|
||||
}
|
||||
|
||||
// Acquire returns an initialized exclusive state, growing the pool if necessary.
|
||||
func (p *Pool) Acquire() (*glua.LState, error) {
|
||||
p.mu.Lock()
|
||||
if p.closed || p.ctx.Err() != nil {
|
||||
p.mu.Unlock()
|
||||
return nil, p.ctx.Err()
|
||||
}
|
||||
|
||||
p.active.Add(1)
|
||||
|
||||
n := len(p.idle)
|
||||
if n != 0 {
|
||||
state := p.idle[n-1]
|
||||
p.idle = p.idle[:n-1]
|
||||
p.mu.Unlock()
|
||||
return state, nil
|
||||
}
|
||||
p.mu.Unlock()
|
||||
|
||||
// TODO: Limit the total number of states. When the limit is reached, wait
|
||||
// for a Release instead of creating another state; allow the wait to be
|
||||
// cancelled by the caller or by Close.
|
||||
state, err := p.factory(p.ctx)
|
||||
if err == nil && state == nil {
|
||||
err = errors.New("Lua state factory returned nil")
|
||||
}
|
||||
if err != nil {
|
||||
if state != nil {
|
||||
state.Close()
|
||||
}
|
||||
p.active.Done()
|
||||
return nil, err
|
||||
}
|
||||
if err := p.ctx.Err(); err != nil {
|
||||
state.Close()
|
||||
p.active.Done()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return state, nil
|
||||
}
|
||||
|
||||
// Release returns a healthy state to the pool and closes a failed or cancelled one.
|
||||
func (p *Pool) Release(state *glua.LState, reusable bool) {
|
||||
if reusable {
|
||||
p.mu.Lock()
|
||||
if !p.closed && p.ctx.Err() == nil && len(p.idle) < maxIdleStates {
|
||||
p.idle = append(p.idle, state)
|
||||
} else {
|
||||
reusable = false
|
||||
}
|
||||
p.mu.Unlock()
|
||||
}
|
||||
|
||||
if !reusable {
|
||||
state.Close()
|
||||
}
|
||||
|
||||
p.active.Done()
|
||||
}
|
||||
|
||||
// Close cancels active work, closes idle states, and waits for borrowed states.
|
||||
func (p *Pool) Close() {
|
||||
p.mu.Lock()
|
||||
if !p.closed {
|
||||
p.closed = true
|
||||
p.cancel()
|
||||
for _, state := range p.idle {
|
||||
state.Close()
|
||||
}
|
||||
p.idle = nil
|
||||
}
|
||||
p.mu.Unlock()
|
||||
|
||||
p.active.Wait()
|
||||
}
|
||||
@@ -0,0 +1,197 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
func TestPoolFactoryFailureClosesReturnedState(t *testing.T) {
|
||||
failure := errors.New("factory failed")
|
||||
state := glua.NewState()
|
||||
_, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) {
|
||||
return state, failure
|
||||
})
|
||||
if !errors.Is(err, failure) || !state.IsClosed() {
|
||||
t.Fatalf("NewPool error = %v, state closed = %t", err, state.IsClosed())
|
||||
}
|
||||
|
||||
var failedState *glua.LState
|
||||
calls := 0
|
||||
pool, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) {
|
||||
calls++
|
||||
if calls == 1 {
|
||||
return glua.NewState(), nil
|
||||
}
|
||||
failedState = glua.NewState()
|
||||
return failedState, failure
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer pool.Close()
|
||||
borrowed, err := pool.Acquire()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer pool.Release(borrowed, true)
|
||||
_, err = pool.Acquire()
|
||||
if !errors.Is(err, failure) || !failedState.IsClosed() {
|
||||
t.Fatalf("Acquire error = %v, state closed = %t", err, failedState.IsClosed())
|
||||
}
|
||||
}
|
||||
|
||||
func TestPoolReusesStatesAndLimitsIdle(t *testing.T) {
|
||||
created := 0
|
||||
pool, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) {
|
||||
created++
|
||||
return glua.NewState(), nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer pool.Close()
|
||||
|
||||
states := make([]*glua.LState, maxIdleStates+3)
|
||||
for i := range states {
|
||||
states[i], err = pool.Acquire()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
for _, state := range states {
|
||||
pool.Release(state, true)
|
||||
}
|
||||
for i, state := range states {
|
||||
if got, want := state.IsClosed(), i >= maxIdleStates; got != want {
|
||||
t.Fatalf("state %d closed = %t, want %t", i, got, want)
|
||||
}
|
||||
}
|
||||
borrowed, err := pool.Acquire()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if created != len(states) {
|
||||
t.Fatalf("Acquire created %d states, want %d", created, len(states))
|
||||
}
|
||||
pool.Release(borrowed, true)
|
||||
}
|
||||
|
||||
func TestPoolCloseCancelsAndWaitsForBorrowedState(t *testing.T) {
|
||||
pool, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) {
|
||||
return glua.NewState(), nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
state, err := pool.Acquire()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
pool.Close()
|
||||
close(done)
|
||||
}()
|
||||
select {
|
||||
case <-pool.Context().Done():
|
||||
case <-time.After(time.Second):
|
||||
pool.Release(state, false)
|
||||
t.Fatal("Close did not cancel the pool context")
|
||||
}
|
||||
select {
|
||||
case <-done:
|
||||
pool.Release(state, false)
|
||||
t.Fatal("Close returned while a state was borrowed")
|
||||
default:
|
||||
}
|
||||
pool.Release(state, true)
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("Close did not finish after Release")
|
||||
}
|
||||
if !state.IsClosed() {
|
||||
t.Fatal("borrowed state was not closed")
|
||||
}
|
||||
if _, err := pool.Acquire(); !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("Acquire after Close = %v, want context.Canceled", err)
|
||||
}
|
||||
pool.Close()
|
||||
}
|
||||
|
||||
func TestPoolCloseCancelsStateCreation(t *testing.T) {
|
||||
started := make(chan struct{})
|
||||
first := true
|
||||
pool, err := NewPool(context.Background(), func(ctx context.Context) (*glua.LState, error) {
|
||||
if first {
|
||||
first = false
|
||||
return glua.NewState(), nil
|
||||
}
|
||||
close(started)
|
||||
<-ctx.Done()
|
||||
return nil, ctx.Err()
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
borrowed, err := pool.Acquire()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
acquireDone := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := pool.Acquire()
|
||||
acquireDone <- err
|
||||
}()
|
||||
<-started
|
||||
closeDone := make(chan struct{})
|
||||
go func() {
|
||||
pool.Close()
|
||||
close(closeDone)
|
||||
}()
|
||||
select {
|
||||
case err := <-acquireDone:
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("Acquire during Close = %v, want context.Canceled", err)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("state creation did not stop after Close")
|
||||
}
|
||||
select {
|
||||
case <-closeDone:
|
||||
pool.Release(borrowed, false)
|
||||
t.Fatal("Close returned while the initial state was borrowed")
|
||||
default:
|
||||
}
|
||||
pool.Release(borrowed, true)
|
||||
select {
|
||||
case <-closeDone:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("Close did not finish after Release")
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkPoolAcquireRelease(b *testing.B) {
|
||||
pool, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) {
|
||||
return glua.NewState(), nil
|
||||
})
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
defer pool.Close()
|
||||
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
state, err := pool.Acquire()
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
pool.Release(state, true)
|
||||
}
|
||||
b.StopTimer()
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
// Package lua provides shared GopherLua programs and state management for Xray scripts.
|
||||
package lua
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"os"
|
||||
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
"github.com/yuin/gopher-lua/parse"
|
||||
)
|
||||
|
||||
// Program holds immutable bytecode that can be run by independent LStates.
|
||||
type Program struct {
|
||||
proto *glua.FunctionProto
|
||||
}
|
||||
|
||||
// CompileFile reads and compiles a Lua file once.
|
||||
func CompileFile(path string) (*Program, error) {
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer f.Close()
|
||||
chunk, err := parse.Parse(bufio.NewReader(f), path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
proto, err := glua.Compile(chunk, path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Program{proto: proto}, nil
|
||||
}
|
||||
|
||||
// NewState creates a VM, makes modules available, and executes the file top level.
|
||||
// Module loaders run only when Lua calls require. Each state gets its own globals.
|
||||
// The caller owns the returned state.
|
||||
func (p *Program) NewState(ctx context.Context, register func(*glua.LState)) (*glua.LState, error) {
|
||||
L := glua.NewState()
|
||||
if register != nil {
|
||||
register(L)
|
||||
}
|
||||
L.SetContext(ctx)
|
||||
L.Push(L.NewFunctionFromProto(p.proto))
|
||||
err := L.PCall(0, 0, nil)
|
||||
L.RemoveContext()
|
||||
if err == nil {
|
||||
err = ctx.Err()
|
||||
}
|
||||
if err != nil {
|
||||
L.Close()
|
||||
return nil, err
|
||||
}
|
||||
return L, nil
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"context"
|
||||
"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)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer first.Close()
|
||||
first.SetGlobal("value", glua.LNumber(42))
|
||||
second, err := program.NewState(context.Background(), 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)
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -23,6 +23,7 @@ require (
|
||||
github.com/stretchr/testify v1.12.1
|
||||
github.com/vishvananda/netlink v1.3.1
|
||||
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0
|
||||
github.com/yuin/gopher-lua v1.1.2
|
||||
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba
|
||||
golang.org/x/crypto v0.57.0
|
||||
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842
|
||||
@@ -36,6 +37,7 @@ require (
|
||||
google.golang.org/protobuf v1.36.12
|
||||
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0
|
||||
h12.io/socks v1.0.3
|
||||
layeh.com/gopher-luar v1.0.11
|
||||
lukechampine.com/blake3 v1.4.1
|
||||
mvdan.cc/gofumpt v0.12.0
|
||||
)
|
||||
|
||||
@@ -2,6 +2,9 @@ github.com/andybalholm/brotli v1.0.6 h1:Yf9fFpf49Zrxb9NlQaluyE92/+X7UVHlhMNJN2sx
|
||||
github.com/andybalholm/brotli v1.0.6/go.mod h1:fO7iG3H7G2nSZ7m0zPUDn85XEX2GTukHGRSepvi9Eig=
|
||||
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e h1:5mgtR5gwIgBKMiGI1QdXldZZ+SNor06Nbu1wCBulQBg=
|
||||
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e/go.mod h1:x7qxEvX6MCVtDuBKHj3E+88+BtrbEMuAL5qGUKItjW8=
|
||||
github.com/chzyer/logex v1.1.10/go.mod h1:+Ywpsq7O8HXn0nuIou7OrIPyXbp3wmkHB+jjWRnGsAI=
|
||||
github.com/chzyer/readline v0.0.0-20180603132655-2972be24d48e/go.mod h1:nSuG5e5PlCu98SY8svDHJxuZscDgtXS6KTTbou5AhLI=
|
||||
github.com/chzyer/test v0.0.0-20180213035817-a1ea475d72b1/go.mod h1:Q3SI9o4m/ZMnBNeIyt5eFwwo7qiLfzFZmjNmxjkiQlU=
|
||||
github.com/cloudflare/circl v1.6.5 h1:O64F26HEqNhznd/hrC5KZXVKYuKM2rx4deZDTc4ihQA=
|
||||
github.com/cloudflare/circl v1.6.5/go.mod h1:h5LNyxAc5nTue9DS5jT+48en2PSDYt3zdGnz5OstK6c=
|
||||
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344 h1:Arcl6UOIS/kgO2nW3A65HN+7CMjSDP/gofXL4CZt1V4=
|
||||
@@ -85,6 +88,9 @@ github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguH
|
||||
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0 h1:rb+fKQFhz+5I2PPuQsNYxI5mUU840XWYtRF0ZBjvkws=
|
||||
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0/go.mod h1:DsJblcWDGt76+FVqBVwbwRhxyyNJsGV48gJLch0OOWI=
|
||||
github.com/yuin/goldmark v1.4.1/go.mod h1:mwnBkeHKe2W/ZEtQ+71ViKU8L12m81fl3OWwC1Zlc8k=
|
||||
github.com/yuin/gopher-lua v0.0.0-20190206043414-8bfc7677f583/go.mod h1:gqRgreBUhTSL0GeU64rtZ3Uq3wtjOa/TB2YfrtkCbVQ=
|
||||
github.com/yuin/gopher-lua v1.1.2 h1:yF/FjE3hD65tBbt0VXLE13HWS9h34fdzJmrWRXwobGA=
|
||||
github.com/yuin/gopher-lua v1.1.2/go.mod h1:7aRmXIWl37SqRf0koeyylBEzJ+aPt8A+mmkQ4f1ntR8=
|
||||
go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
|
||||
go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
|
||||
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
|
||||
@@ -109,6 +115,7 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ
|
||||
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk=
|
||||
golang.org/x/sync v0.23.0/go.mod h1:sUUOizhqBxiL6pEWpqNLUiaJn1ShEbZ6BBqskPbjZm0=
|
||||
golang.org/x/sys v0.0.0-20190204203706-41f3e6584952/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
@@ -159,6 +166,8 @@ gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0 h1:Lk6hARj5UPY47dBep70OD/TI
|
||||
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0/go.mod h1:QkHjoMIBaYtpVufgwv3keYAbln78mBoCuShZrPrer1Q=
|
||||
h12.io/socks v1.0.3 h1:Ka3qaQewws4j4/eDQnOdpr4wXsC//dXtWvftlIcCQUo=
|
||||
h12.io/socks v1.0.3/go.mod h1:AIhxy1jOId/XCz9BO+EIgNL2rQiPTBNnOfnVnQ+3Eck=
|
||||
layeh.com/gopher-luar v1.0.11 h1:8zJudpKI6HWkoh9eyyNFaTM79PY6CAPcIr6X/KTiliw=
|
||||
layeh.com/gopher-luar v1.0.11/go.mod h1:TPnIVCZ2RJBndm7ohXyaqfhzjlZ+OA2SZR/YwL8tECk=
|
||||
lukechampine.com/blake3 v1.4.1 h1:I3Smz7gso8w4/TunLKec6K2fn+kyKtDxr/xcQEN84Wg=
|
||||
lukechampine.com/blake3 v1.4.1/go.mod h1:QFosUxmjB8mnrWFSNwKmvxHpfY72bmD2tQ0kBMM3kwo=
|
||||
mvdan.cc/gofumpt v0.12.0 h1:1Lbudkz2kpM9Cjz2pL4M19u7q+GaEhCTNf7N9mfpcho=
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user