feat(router): add Lua scripting for routing script

This commit is contained in:
Meo597
2026-10-01 18:25:53 +08:00
parent 1c52c65872
commit 2440f53cdd
11 changed files with 1093 additions and 13 deletions
+3
View File
@@ -470,6 +470,9 @@ func (d *DefaultDispatcher) routedDispatch(ctx context.Context, link *transport.
return // DO NOT CHANGE: the traffic shouldn't be processed by default outbound if the specified outbound tag doesn't exist (yet), e.g., VLESS Reverse Proxy
}
} else {
if err != common.ErrNoClue {
errors.LogErrorInner(ctx, err, "failed to pick route for ", destination)
}
errors.LogInfo(ctx, "default route for ", destination)
}
}
+59 -9
View File
@@ -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
View File
@@ -587,8 +587,10 @@ type Config struct {
DomainStrategy Config_DomainStrategy `protobuf:"varint,1,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.app.router.Config_DomainStrategy" json:"domain_strategy,omitempty"`
Rule []*RoutingRule `protobuf:"bytes,2,rep,name=rule,proto3" json:"rule,omitempty"`
BalancingRule []*BalancingRule `protobuf:"bytes,3,rep,name=balancing_rule,json=balancingRule,proto3" json:"balancing_rule,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
// Absolute path to the Lua routing script.
Script string `protobuf:"bytes,4,opt,name=script,proto3" json:"script,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *Config) Reset() {
@@ -642,6 +644,13 @@ func (x *Config) GetBalancingRule() []*BalancingRule {
return nil
}
func (x *Config) GetScript() string {
if x != nil {
return x.Script
}
return ""
}
var File_app_router_config_proto protoreflect.FileDescriptor
const file_app_router_config_proto_rawDesc = "" +
@@ -699,11 +708,12 @@ const file_app_router_config_proto_rawDesc = "" +
"\tbaselines\x18\x03 \x03(\x03R\tbaselines\x12\x1a\n" +
"\bexpected\x18\x04 \x01(\x05R\bexpected\x12\x16\n" +
"\x06maxRTT\x18\x05 \x01(\x03R\x06maxRTT\x12\x1c\n" +
"\ttolerance\x18\x06 \x01(\x02R\ttolerance\"\x96\x02\n" +
"\ttolerance\x18\x06 \x01(\x02R\ttolerance\"\xae\x02\n" +
"\x06Config\x12O\n" +
"\x0fdomain_strategy\x18\x01 \x01(\x0e2&.xray.app.router.Config.DomainStrategyR\x0edomainStrategy\x120\n" +
"\x04rule\x18\x02 \x03(\v2\x1c.xray.app.router.RoutingRuleR\x04rule\x12E\n" +
"\x0ebalancing_rule\x18\x03 \x03(\v2\x1e.xray.app.router.BalancingRuleR\rbalancingRule\"B\n" +
"\x0ebalancing_rule\x18\x03 \x03(\v2\x1e.xray.app.router.BalancingRuleR\rbalancingRule\x12\x16\n" +
"\x06script\x18\x04 \x01(\tR\x06script\"B\n" +
"\x0eDomainStrategy\x12\b\n" +
"\x04AsIs\x10\x00\x12\x10\n" +
"\fIpIfNonMatch\x10\x02\x12\x0e\n" +
+2
View File
@@ -110,4 +110,6 @@ message Config {
DomainStrategy domain_strategy = 1;
repeated RoutingRule rule = 2;
repeated BalancingRule balancing_rule = 3;
// Absolute path to the Lua routing script.
string script = 4;
}
+206
View File
@@ -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)
}
+292
View File
@@ -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)
+17
View File
@@ -20,6 +20,8 @@ import (
type Router struct {
domainStrategy Config_DomainStrategy
rules atomic.Pointer[[]*Rule]
scriptPath string
script *scriptEngine
balancers atomic.Pointer[map[string]*Balancer]
dns dns.Client
@@ -40,6 +42,7 @@ type Route struct {
// Init initializes the Router.
func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm outbound.Manager, dispatcher routing.Dispatcher) error {
r.domainStrategy = config.DomainStrategy
r.scriptPath = config.Script
r.dns = d
r.ctx = ctx
r.ohm = ohm
@@ -52,6 +55,10 @@ func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm out
// PickRoute implements routing.Router.
func (r *Router) PickRoute(ctx routing.Context) (routing.Route, error) {
if r.script != nil {
return r.script.pickRoute(ctx)
}
originalCtx := ctx
rule, ctx, err := r.pickRouteInternal(ctx)
if err != nil {
@@ -221,6 +228,13 @@ func (r *Router) pickRouteInternal(ctx routing.Context) (*Rule, routing.Context,
// Start implements common.Runnable.
func (r *Router) Start() error {
if r.scriptPath != "" {
engine, err := newScriptEngine(r.scriptPath, r)
if err != nil {
return errors.New("failed to initialize routing script").Base(err)
}
r.script = engine
}
return nil
}
@@ -235,6 +249,9 @@ func closeWebhooks(rules []*Rule) {
// Close implements common.Closable.
func (r *Router) Close() error {
if r.script != nil {
r.script.close()
}
r.mu.Lock()
defer r.mu.Unlock()
closeWebhooks(*r.rules.Load())
+79
View File
@@ -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
}
+372
View File
@@ -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")
}
}
+11
View File
@@ -7,6 +7,7 @@ import (
"github.com/xtls/xray-core/app/router"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/platform"
"github.com/xtls/xray-core/common/serial"
"google.golang.org/protobuf/proto"
@@ -72,6 +73,7 @@ type RouterConfig struct {
RuleList []json.RawMessage `json:"rules"`
DomainStrategy *string `json:"domainStrategy"`
Balancers []*BalancingRule `json:"balancers"`
Script string `json:"script"`
}
func (c *RouterConfig) getDomainStrategy() router.Config_DomainStrategy {
@@ -92,6 +94,15 @@ func (c *RouterConfig) getDomainStrategy() router.Config_DomainStrategy {
func (c *RouterConfig) Build() (*router.Config, error) {
config := new(router.Config)
if c.Script != "" {
path, err := platform.ResolveLuaFile(c.Script)
if err != nil {
return nil, errors.New("failed to resolve routing script").Base(err)
}
config.Script = path
}
config.DomainStrategy = c.getDomainStrategy()
var rawRuleList []json.RawMessage
+38
View File
@@ -2,6 +2,8 @@ package conf_test
import (
"encoding/json"
"os"
"path/filepath"
"testing"
"time"
_ "unsafe"
@@ -236,3 +238,39 @@ func TestRouterConfig(t *testing.T) {
},
})
}
func TestRouterScriptConfig(t *testing.T) {
dir := t.TempDir()
t.Setenv("xray.location.confdir", dir)
path := filepath.Join(dir, "route.lua")
if err := os.WriteFile(path, []byte("function HandleRoute() end"), 0o600); err != nil {
t.Fatal(err)
}
for _, tc := range []struct {
name string
script string
wantError bool
}{
{"relative", "route.lua", false},
{"absolute", path, false},
{"missing", "missing.lua", true},
{"directory", dir, true},
} {
t.Run(tc.name, func(t *testing.T) {
built, err := (&RouterConfig{Script: tc.script}).Build()
if tc.wantError {
if err == nil {
t.Fatal("Build accepted invalid script path")
}
return
}
if err != nil {
t.Fatal(err)
}
if built.Script != path {
t.Fatalf("script path = %q, want %q", built.Script, path)
}
})
}
}