mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-05 07:13:36 +00:00
router: preserve Lua states after recoverable route errors & refactor
This commit is contained in:
+28
-21
@@ -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)
|
||||
|
||||
+69
-59
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
+22
-14
@@ -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
|
||||
}
|
||||
|
||||
+70
-60
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user