mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-01 13:35:53 +00:00
feat(router): add Lua scripting for routing script
This commit is contained in:
@@ -470,6 +470,9 @@ func (d *DefaultDispatcher) routedDispatch(ctx context.Context, link *transport.
|
||||
return // DO NOT CHANGE: the traffic shouldn't be processed by default outbound if the specified outbound tag doesn't exist (yet), e.g., VLESS Reverse Proxy
|
||||
}
|
||||
} else {
|
||||
if err != common.ErrNoClue {
|
||||
errors.LogErrorInner(ctx, err, "failed to pick route for ", destination)
|
||||
}
|
||||
errors.LogInfo(ctx, "default route for ", destination)
|
||||
}
|
||||
}
|
||||
|
||||
+59
-9
@@ -11,8 +11,27 @@ import (
|
||||
lua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
// RegisterLua makes xray.dns available to require in an LState.
|
||||
// RegisterLua makes xray.dns available with Query and optional Servers.
|
||||
func RegisterLua(L *lua.LState, client featureDNS.Client) {
|
||||
// A configured DNS app passes its *DNS instance here.
|
||||
if s, ok := client.(*DNS); ok {
|
||||
registerLua(L, s, true)
|
||||
return
|
||||
}
|
||||
L.PreloadModule("xray.dns", func(L *lua.LState) int {
|
||||
module := L.NewTable()
|
||||
module.RawSetString("Query", newLuaClientQuery(L, client))
|
||||
L.Push(module)
|
||||
return 1
|
||||
})
|
||||
}
|
||||
|
||||
// RegisterLua makes xray.dns available to DNS scripts with Servers so no Query.
|
||||
func (s *DNS) RegisterLua(L *lua.LState) {
|
||||
registerLua(L, s, false)
|
||||
}
|
||||
|
||||
func registerLua(L *lua.LState, s *DNS, exposeQuery bool) {
|
||||
L.PreloadModule("xray.dns", func(L *lua.LState) int {
|
||||
servers := L.NewTable()
|
||||
for i, client := range s.clients {
|
||||
@@ -59,19 +78,56 @@ func (s *DNS) RegisterLua(L *lua.LState) {
|
||||
}))
|
||||
servers.RawSetInt(i+1, server)
|
||||
}
|
||||
|
||||
module := L.NewTable()
|
||||
module.RawSetString("Servers", servers)
|
||||
if exposeQuery {
|
||||
module.RawSetString("Query", newLuaClientQuery(L, s))
|
||||
}
|
||||
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 := L.Context()
|
||||
previous, top := L.Context(), L.GetTop()
|
||||
L.SetContext(ctx)
|
||||
defer func() {
|
||||
L.SetTop(top)
|
||||
if previous == nil {
|
||||
L.RemoveContext()
|
||||
} else {
|
||||
@@ -87,13 +143,7 @@ func (s *DNS) CallLuaHook(L *lua.LState, ctx context.Context, domain string, opt
|
||||
lua.LBool(option.IPv6Enable), lua.LBool(option.FakeEnable)); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
addresses, ttlValue, errorValue := L.Get(-3), L.Get(-2), L.Get(-1)
|
||||
L.Pop(3)
|
||||
ips, ttl, err := readLuaDNSResult(addresses, ttlValue, errorValue)
|
||||
if ctx.Err() != nil {
|
||||
return nil, 0, ctx.Err()
|
||||
}
|
||||
return ips, ttl, err
|
||||
return readLuaDNSResult(L.Get(-3), L.Get(-2), L.Get(-1))
|
||||
}
|
||||
|
||||
func readLuaDNSResult(addresses, ttlValue, errorValue lua.LValue) ([]net.IP, uint32, error) {
|
||||
|
||||
+14
-4
@@ -587,8 +587,10 @@ type Config struct {
|
||||
DomainStrategy Config_DomainStrategy `protobuf:"varint,1,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.app.router.Config_DomainStrategy" json:"domain_strategy,omitempty"`
|
||||
Rule []*RoutingRule `protobuf:"bytes,2,rep,name=rule,proto3" json:"rule,omitempty"`
|
||||
BalancingRule []*BalancingRule `protobuf:"bytes,3,rep,name=balancing_rule,json=balancingRule,proto3" json:"balancing_rule,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
// Absolute path to the Lua routing script.
|
||||
Script string `protobuf:"bytes,4,opt,name=script,proto3" json:"script,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *Config) Reset() {
|
||||
@@ -642,6 +644,13 @@ func (x *Config) GetBalancingRule() []*BalancingRule {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *Config) GetScript() string {
|
||||
if x != nil {
|
||||
return x.Script
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
var File_app_router_config_proto protoreflect.FileDescriptor
|
||||
|
||||
const file_app_router_config_proto_rawDesc = "" +
|
||||
@@ -699,11 +708,12 @@ const file_app_router_config_proto_rawDesc = "" +
|
||||
"\tbaselines\x18\x03 \x03(\x03R\tbaselines\x12\x1a\n" +
|
||||
"\bexpected\x18\x04 \x01(\x05R\bexpected\x12\x16\n" +
|
||||
"\x06maxRTT\x18\x05 \x01(\x03R\x06maxRTT\x12\x1c\n" +
|
||||
"\ttolerance\x18\x06 \x01(\x02R\ttolerance\"\x96\x02\n" +
|
||||
"\ttolerance\x18\x06 \x01(\x02R\ttolerance\"\xae\x02\n" +
|
||||
"\x06Config\x12O\n" +
|
||||
"\x0fdomain_strategy\x18\x01 \x01(\x0e2&.xray.app.router.Config.DomainStrategyR\x0edomainStrategy\x120\n" +
|
||||
"\x04rule\x18\x02 \x03(\v2\x1c.xray.app.router.RoutingRuleR\x04rule\x12E\n" +
|
||||
"\x0ebalancing_rule\x18\x03 \x03(\v2\x1e.xray.app.router.BalancingRuleR\rbalancingRule\"B\n" +
|
||||
"\x0ebalancing_rule\x18\x03 \x03(\v2\x1e.xray.app.router.BalancingRuleR\rbalancingRule\x12\x16\n" +
|
||||
"\x06script\x18\x04 \x01(\tR\x06script\"B\n" +
|
||||
"\x0eDomainStrategy\x12\b\n" +
|
||||
"\x04AsIs\x10\x00\x12\x10\n" +
|
||||
"\fIpIfNonMatch\x10\x02\x12\x0e\n" +
|
||||
|
||||
@@ -110,4 +110,6 @@ message Config {
|
||||
DomainStrategy domain_strategy = 1;
|
||||
repeated RoutingRule rule = 2;
|
||||
repeated BalancingRule balancing_rule = 3;
|
||||
// Absolute path to the Lua routing script.
|
||||
string script = 4;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,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")
|
||||
}
|
||||
}
|
||||
@@ -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