diff --git a/app/dispatcher/default.go b/app/dispatcher/default.go index 2717df5fe..cae973517 100644 --- a/app/dispatcher/default.go +++ b/app/dispatcher/default.go @@ -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) } } diff --git a/app/dns/lua.go b/app/dns/lua.go index 6f5dda79c..ed0295942 100644 --- a/app/dns/lua.go +++ b/app/dns/lua.go @@ -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) { diff --git a/app/router/config.pb.go b/app/router/config.pb.go index 636096a7f..2eea46d5f 100644 --- a/app/router/config.pb.go +++ b/app/router/config.pb.go @@ -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" + diff --git a/app/router/config.proto b/app/router/config.proto index 67e4e47f1..09b64b0a3 100644 --- a/app/router/config.proto +++ b/app/router/config.proto @@ -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; } diff --git a/app/router/lua.go b/app/router/lua.go new file mode 100644 index 000000000..259f3a96b --- /dev/null +++ b/app/router/lua.go @@ -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) +} diff --git a/app/router/lua_test.go b/app/router/lua_test.go new file mode 100644 index 000000000..52ccb0e04 --- /dev/null +++ b/app/router/lua_test.go @@ -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) diff --git a/app/router/router.go b/app/router/router.go index 406f932fb..97ffba13c 100644 --- a/app/router/router.go +++ b/app/router/router.go @@ -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()) diff --git a/app/router/script.go b/app/router/script.go new file mode 100644 index 000000000..db17e9d4f --- /dev/null +++ b/app/router/script.go @@ -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 +} diff --git a/app/router/script_test.go b/app/router/script_test.go new file mode 100644 index 000000000..553812cd7 --- /dev/null +++ b/app/router/script_test.go @@ -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") + } +} diff --git a/infra/conf/router.go b/infra/conf/router.go index 3f85ce8fa..ba5cfeef2 100644 --- a/infra/conf/router.go +++ b/infra/conf/router.go @@ -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 diff --git a/infra/conf/router_test.go b/infra/conf/router_test.go index 130cf4f78..04e4bcd83 100644 --- a/infra/conf/router_test.go +++ b/infra/conf/router_test.go @@ -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) + } + }) + } +}