mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-02 13:56:39 +00:00
293 lines
9.8 KiB
Go
293 lines
9.8 KiB
Go
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)
|