From 52be8f81703d7310c14b3e6f4bf2ee219b09737b Mon Sep 17 00:00:00 2001 From: Meo597 <197331664+Meo597@users.noreply.github.com> Date: Mon, 5 Oct 2026 06:38:24 +0800 Subject: [PATCH] router: preserve Lua states after recoverable route errors & refactor --- app/router/lua.go | 49 ++++++++------ app/router/lua_test.go | 128 ++++++++++++++++++++----------------- app/router/script.go | 36 +++++++---- app/router/script_test.go | 130 ++++++++++++++++++++------------------ 4 files changed, 189 insertions(+), 154 deletions(-) diff --git a/app/router/lua.go b/app/router/lua.go index e0ce0afb4..c09106228 100644 --- a/app/router/lua.go +++ b/app/router/lua.go @@ -106,41 +106,48 @@ func checkLuaContext(L *lua.LState) routing.Context { return ctx } -// callLuaHook invokes HandleRoute in the supplied state. -func (r *Router) callLuaHook(L *lua.LState, routeCtx routing.Context) (string, string, error) { - top := L.GetTop() - defer L.SetTop(top) +// callLuaRoute runs HandleRoute and leaves (outboundTag, ruleTag, err) on the stack. +func callLuaRoute(L *lua.LState, ctx routing.Context) error { fn := L.GetGlobal("HandleRoute") if fn.Type() != lua.LTFunction { - return "", "", errors.New("routing script must define HandleRoute(...)") + return errors.New("routing script must define HandleRoute(...)") } + value := L.NewUserData() - value.Value = routeCtx + value.Value = ctx 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(strings.ToLower(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)) + + return L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true}, + value, + lua.LString(ctx.GetInboundTag()), + lua.LNumber(ctx.GetSourcePort()), + lua.LNumber(ctx.GetTargetPort()), + lua.LNumber(ctx.GetLocalPort()), + lua.LString(strings.ToLower(ctx.GetTargetDomain())), + lua.LNumber(ctx.GetNetwork()), + lua.LString(ctx.GetProtocol()), + lua.LString(ctx.GetUser()), + lua.LNumber(ctx.GetVlessRoute()), + lua.LBool(ctx.GetSkipDNSResolve())) } -func readLuaRouteResult(tagValue, ruleValue, errorValue lua.LValue) (string, string, error) { - if err := xlua.ReadError(errorValue, "routing script error must be an error or string"); err != nil { +// readLuaRouteResult reads (outboundTag, ruleTag, err) from the stack. +func readLuaRouteResult(L *lua.LState) (string, string, error) { + if err := xlua.ReadError(L.Get(-1), "routing script error must be an error or string"); err != nil { return "", "", err } - tag, err := xlua.ReadOptionalString(tagValue, "routing script outboundTag must be a string or nil") - if err != nil || tag == "" { + + outboundTag, err := xlua.ReadOptionalString(L.Get(-3), "routing script outboundTag must be a string or nil") + if err != nil || outboundTag == "" { return "", "", err } - ruleTag, err := xlua.ReadOptionalString(ruleValue, "routing script ruleTag must be a string") + + ruleTag, err := xlua.ReadOptionalString(L.Get(-2), "routing script ruleTag must be a string") if err != nil { return "", "", err } - return tag, ruleTag, nil + + return outboundTag, ruleTag, nil } type processFinder func(string, string, uint16, string, uint16) (int, string, string, error) diff --git a/app/router/lua_test.go b/app/router/lua_test.go index 1229dadfd..f90b1b867 100644 --- a/app/router/lua_test.go +++ b/app/router/lua_test.go @@ -48,7 +48,7 @@ func newLuaRouteTestContext() *luaRouteTestContext { } } -func newLuaRouterState(t *testing.T, script string) (*Router, *lua.LState) { +func newLuaRouterState(t *testing.T, script string) *lua.LState { t.Helper() r := new(Router) if err := r.Init(context.Background(), &Config{}, nil, nil, nil); err != nil { @@ -61,11 +61,11 @@ func newLuaRouterState(t *testing.T, script string) (*Router, *lua.LState) { if err := L.DoString(script); err != nil { t.Fatal(err) } - return r, L + return L } func TestLuaRouteBinding(t *testing.T) { - r, L := newLuaRouterState(t, ` + 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) @@ -88,9 +88,11 @@ function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort, end`) ctx := newLuaRouteTestContext() - tag, rule, err := r.callLuaHook(L, ctx) - if err != nil || tag != "out" || rule != "rule" { - t.Fatalf("hook = %q, %q, %v", tag, rule, err) + if err := callLuaRoute(L, ctx); err != nil { + t.Fatal(err) + } + if L.GetTop() != 3 || L.Get(1) != lua.LString("out") || L.Get(2) != lua.LString("rule") || L.Get(3) != lua.LNil { + t.Fatal("callLuaRoute did not leave the three route results on the stack") } if L.GetGlobal("savedContext").(*lua.LUserData).Value != ctx { t.Fatal("routing context was copied") @@ -117,70 +119,73 @@ assert(require("xray.router").LocalOS == expectedOS)`); err != nil { } } -func TestLuaRouteResult(t *testing.T) { +func TestReadLuaRouteResult(t *testing.T) { nativeErr := go_errors.New("native failure") for _, tc := range []struct { - name, body, tag, rule, wantErr string - native bool + name, values string + wantTag, wantRule string + wantErr error + wantMessage string }{ - {name: "route", body: `return "out", "rule"`, tag: "out", rule: "rule"}, - {name: "no match", body: `return nil`}, - {name: "empty tag", body: `return ""`}, - {name: "no match ignores rule", body: `return nil, false`}, - {name: "empty tag ignores rule", body: `return "", false`}, - {name: "missing rule", body: `return "out"`, tag: "out"}, - {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: "error overrides invalid tags", body: `return false, false, nativeError`, native: true}, - {name: "invalid error", body: `return "out", "rule", false`, wantErr: "error or string"}, - {name: "wrong error userdata", body: `return "out", "rule", wrongError`, wantErr: "error or string"}, - {name: "runtime error", body: `error("runtime failure")`, wantErr: "runtime failure"}, + {name: "route", values: `"out", "rule"`, wantTag: "out", wantRule: "rule"}, + {name: "no match", values: `nil`}, + {name: "empty tag", values: `""`}, + {name: "no match ignores rule", values: `nil, false`}, + {name: "empty tag ignores rule", values: `"", false`}, + {name: "missing rule", values: `"out"`, wantTag: "out"}, + {name: "invalid tag", values: `1`, wantMessage: "outboundTag"}, + {name: "invalid rule", values: `"out", false`, wantMessage: "ruleTag"}, + {name: "string error", values: `nil, nil, "script failure"`, wantMessage: "script failure"}, + {name: "native error", values: `nil, nil, nativeError`, wantErr: nativeErr}, + {name: "error overrides invalid tags", values: `false, false, nativeError`, wantErr: nativeErr}, + {name: "invalid error", values: `"out", "rule", false`, wantMessage: "error or string"}, + {name: "wrong error userdata", values: `"out", "rule", wrongError`, wantMessage: "error or string"}, } { 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) - wrong := L.NewUserData() - wrong.Value = "not a native error" - L.SetGlobal("wrongError", wrong) - L.Push(lua.LTrue) - - tag, rule, err := r.callLuaHook(L, &routing_session.Context{}) - if tag != tc.tag || rule != tc.rule { - t.Fatalf("result = %q, %q, %v", tag, rule, err) + L := lua.NewState() + defer L.Close() + for name, value := range map[string]any{"nativeError": nativeErr, "wrongError": "not a native error"} { + ud := L.NewUserData() + ud.Value = value + L.SetGlobal(name, ud) + } + fn, err := L.LoadString("return " + tc.values) + if err != nil { + t.Fatal(err) + } + if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true}); err != nil { + t.Fatal(err) + } + outboundTag, ruleTag, err := readLuaRouteResult(L) + if outboundTag != tc.wantTag || ruleTag != tc.wantRule { + t.Fatalf("result = %q, %q, %v; want %q, %q", outboundTag, ruleTag, err, tc.wantTag, tc.wantRule) } switch { - case tc.native: - if err != nativeErr { + case tc.wantErr != nil: + if err != tc.wantErr { 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 tc.wantMessage != "": + if err == nil || !strings.Contains(err.Error(), tc.wantMessage) { + t.Fatalf("error = %v, want %q", err, tc.wantMessage) } case err != nil: t.Fatal(err) } - if L.GetTop() != 1 || L.Get(1) != lua.LTrue { - t.Fatal("hook did not restore the stack") - } }) } } -func TestLuaRouteCancellation(t *testing.T) { - r, L := newLuaRouterState(t, `function HandleRoute() while true do end end`) +func TestCallLuaRouteCancellation(t *testing.T) { + L := newLuaRouterState(t, `function HandleRoute() while true do end end`) ctx, cancel := context.WithCancel(context.Background()) cancel() L.SetContext(ctx) - if _, _, err := r.callLuaHook(L, &routing_session.Context{}); err == nil { - t.Fatal("CallLuaHook did not stop after context cancellation") + if err := callLuaRoute(L, &routing_session.Context{}); err == nil { + t.Fatal("callLuaRoute did not stop after context cancellation") } - if L.Context() != ctx || L.GetTop() != 0 { - t.Fatal("CallLuaHook did not restore the Lua state") + if L.Context() != ctx { + t.Fatal("callLuaRoute changed the Lua state's context") } } @@ -227,9 +232,9 @@ func TestFindProcess(t *testing.T) { } } -// BenchmarkLuaRouteHookCall isolates a preloaded Lua hook and its routing context bridge. +// BenchmarkLuaRoute measures a preloaded routing script using its context bridge. // The direct case runs an equivalent native routing rule. -func BenchmarkLuaRouteHookCall(b *testing.B) { +func BenchmarkLuaRoute(b *testing.B) { r := new(Router) if err := r.Init(context.Background(), &Config{Rule: []*RoutingRule{{ TargetTag: &RoutingRule_Tag{Tag: "out"}, @@ -262,36 +267,41 @@ end } L.SetContext(context.Background()) - routeCtx := newLuaRouteTestContext() + ctx := newLuaRouteTestContext() for _, benchmark := range []struct { name string route func() (string, string, error) }{ {"direct", func() (string, string, error) { - route, err := r.PickRoute(routeCtx) + route, err := r.PickRoute(ctx) if err != nil { return "", "", err } return route.GetOutboundTag(), route.GetRuleTag(), nil }}, - {"lua_hook", func() (string, string, error) { - return r.callLuaHook(L, routeCtx) + {"lua_script", func() (string, string, error) { + if err := callLuaRoute(L, ctx); err != nil { + return "", "", err + } + outboundTag, ruleTag, err := readLuaRouteResult(L) + L.Pop(3) + return outboundTag, ruleTag, err }}, } { b.Run(benchmark.name, func(b *testing.B) { b.ReportAllocs() b.ResetTimer() - var tag, rule string + var outboundTag, ruleTag string var err error for i := 0; i < b.N; i++ { - tag, rule, err = benchmark.route() + outboundTag, ruleTag, 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) + if outboundTag != "out" || ruleTag != "rule" { + b.Fatalf("route() = %q, %q; want out, rule", outboundTag, ruleTag) } }) } diff --git a/app/router/script.go b/app/router/script.go index 164438ca6..b4a720b56 100644 --- a/app/router/script.go +++ b/app/router/script.go @@ -16,8 +16,7 @@ import ( const scriptExecutionTimeout = 6 * time.Second type scriptEngine struct { - router *Router - pool *xlua.Pool + pool *xlua.Pool } func newScriptEngine(path string, router *Router) (*scriptEngine, error) { @@ -25,8 +24,8 @@ func newScriptEngine(path string, router *Router) (*scriptEngine, error) { if err != nil { return nil, err } - e := &scriptEngine{router: router} - e.pool, err = xlua.NewPool(router.ctx, scriptExecutionTimeout, program.NewStateFactory( + + pool, err := xlua.NewPool(router.ctx, scriptExecutionTimeout, program.NewStateFactory( scriptExecutionTimeout*20, func(L *lua.LState) { geodata.RegisterLua(L) @@ -43,8 +42,9 @@ func newScriptEngine(path string, router *Router) (*scriptEngine, error) { if err != nil { return nil, err } + errors.LogInfo(router.ctx, "routing script initialized from ", path) - return e, nil + return &scriptEngine{pool: pool}, nil } func (e *scriptEngine) close() { @@ -52,17 +52,25 @@ func (e *scriptEngine) close() { } func (e *scriptEngine) pickRoute(ctx routing.Context) (routing.Route, error) { - var tag, ruleTag string - err := e.pool.WithState(nil, 0, func(L *lua.LState) error { - var hookErr error - tag, ruleTag, hookErr = e.router.callLuaHook(L, ctx) - return hookErr - }) - if err != nil { + var outboundTag, ruleTag string + var routeErr error + + if err := e.pool.WithState(nil, 0, func(L *lua.LState) error { + if err := callLuaRoute(L, ctx); err != nil { + return err + } + outboundTag, ruleTag, routeErr = readLuaRouteResult(L) + return nil + }); err != nil { return nil, err } - if tag == "" { + + if routeErr != nil { + return nil, routeErr + } + if outboundTag == "" { return nil, common.ErrNoClue } - return &Route{Context: ctx, outboundTag: tag, ruleTag: ruleTag}, nil + + return &Route{Context: ctx, outboundTag: outboundTag, ruleTag: ruleTag}, nil } diff --git a/app/router/script_test.go b/app/router/script_test.go index ebc0f5ecc..7daeb868a 100644 --- a/app/router/script_test.go +++ b/app/router/script_test.go @@ -5,6 +5,7 @@ import ( stdnet "net" "os" "path/filepath" + "strings" "sync" "sync/atomic" "testing" @@ -90,34 +91,76 @@ func TestRouterScriptStartup(t *testing.T) { } 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, ` + for _, tc := range []struct { + name, body string + wantTag, wantRule string + wantErr error + wantMessage string + wantCalls string + }{ + {name: "route", body: `return "lua-out", "lua-rule"`, wantTag: "lua-out", wantRule: "lua-rule", wantCalls: "2"}, + {name: "no match", body: `return nil`, wantErr: common.ErrNoClue, wantCalls: "2"}, + {name: "empty tag", body: `return ""`, wantErr: common.ErrNoClue, wantCalls: "2"}, + {name: "balancer error", body: `local tag, err = router:PickOutbound("missing"); return tag, nil, err`, wantMessage: "not found", wantCalls: "2"}, + {name: "string error", body: `return nil, nil, "blocked"`, wantMessage: "blocked", wantCalls: "2"}, + {name: "invalid tag", body: `return false`, wantMessage: "outboundTag", wantCalls: "2"}, + {name: "invalid rule", body: `return "lua-out", false`, wantMessage: "ruleTag", wantCalls: "2"}, + {name: "execution error", body: `error("execution failed")`, wantMessage: "execution failed", wantCalls: "1"}, + } { + t.Run(tc.name, func(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 + }} + script := ` +local router = require("xray.router") +local calls = 0 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") + calls = calls + 1 + if inbound == "count" then return "lua-out", tostring(calls) end + ` + tc.body + ` +end +` + r := startLuaRouter(t, script, 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) + switch { + case tc.wantErr != nil: + if err != tc.wantErr { + t.Fatalf("route error = %v, want %v", err, tc.wantErr) + } + case tc.wantMessage != "": + if err == nil || !strings.Contains(err.Error(), tc.wantMessage) { + t.Fatalf("route error = %v, want %q", err, tc.wantMessage) + } + case err != nil: + t.Fatal(err) + } + if tc.wantTag == "" { + if route != nil { + t.Fatalf("route = %v, want nil", route) + } + } else if route == nil || route.GetOutboundTag() != tc.wantTag || route.GetRuleTag() != tc.wantRule || route.(*Route).Context != ctx { + t.Fatalf("route = %v; want %q, %q and original context", route, tc.wantTag, tc.wantRule) + } + + ctx.Inbound.Tag = "count" + route, err = r.PickRoute(ctx) + if err != nil || route == nil || route.GetOutboundTag() != "lua-out" || route.GetRuleTag() != tc.wantCalls { + t.Fatalf("next route = %v, %v; want lua-out, calls %s", route, err, tc.wantCalls) + } + if dnsCalls.Load() != 0 { + t.Fatal("script routing implicitly resolved DNS") + } + }) } } @@ -225,39 +268,6 @@ end`, nil, config("a")) 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 {