From 73fb3e8f4aff49dc6da4fd110a7f6491e5d58f57 Mon Sep 17 00:00:00 2001 From: Meo597 <197331664+Meo597@users.noreply.github.com> Date: Sun, 4 Oct 2026 06:29:00 +0800 Subject: [PATCH] refactor: extract shared result validation and userdata helpers --- app/dns/lua.go | 56 +++++------------- app/router/lua.go | 69 ++++++---------------- app/router/lua_test.go | 9 +++ common/lua/utils.go | 74 ++++++++++++++++++++++++ common/lua/utils_test.go | 121 +++++++++++++++++++++++++++++++++++++++ 5 files changed, 237 insertions(+), 92 deletions(-) create mode 100644 common/lua/utils.go create mode 100644 common/lua/utils_test.go diff --git a/app/dns/lua.go b/app/dns/lua.go index f74f197d1..6eb7c1c3a 100644 --- a/app/dns/lua.go +++ b/app/dns/lua.go @@ -2,10 +2,10 @@ package dns import ( "context" - "math" "strings" "github.com/xtls/xray-core/common/errors" + luamgr "github.com/xtls/xray-core/common/lua" "github.com/xtls/xray-core/common/net" featureDNS "github.com/xtls/xray-core/features/dns" "github.com/xtls/xray-core/features/dns/localdns" @@ -82,17 +82,9 @@ func registerLua(L *lua.LState, servers []luaDNSServer, client featureDNS.Client } else { ips, ttl, err = client.query(ctx, string(domain), option) } - addresses := L.NewUserData() - addresses.Value = ips - L.Push(addresses) + luamgr.PushUserData(L, ips) L.Push(lua.LNumber(ttl)) - if err != nil { - ud := L.NewUserData() - ud.Value = err - L.Push(ud) - } else { - L.Push(lua.LNil) - } + luamgr.PushError(L, err) return 3 })) serverList.RawSetInt(i+1, server) @@ -127,17 +119,9 @@ func newLuaClientQuery(L *lua.LState, client featureDNS.Client) *lua.LFunction { return 0 } ips, ttl, err := client.LookupIP(string(domain), option) - addresses := L.NewUserData() - addresses.Value = ips - L.Push(addresses) + luamgr.PushUserData(L, ips) L.Push(lua.LNumber(ttl)) - if err != nil { - ud := L.NewUserData() - ud.Value = err - L.Push(ud) - } else { - L.Push(lua.LNil) - } + luamgr.PushError(L, err) return 3 }) } @@ -160,34 +144,22 @@ func (s *DNS) callLuaHook(L *lua.LState, domain string, option featureDNS.IPOpti } func readLuaDNSResult(addresses, ttlValue, errorValue lua.LValue) ([]net.IP, uint32, error) { - if errorValue != lua.LNil { - if ud, ok := errorValue.(*lua.LUserData); ok { - if err, ok := ud.Value.(error); ok { - return nil, 0, err - } - } - if s, ok := errorValue.(lua.LString); ok { - return nil, 0, errors.New(string(s)) - } - return nil, 0, errors.New("DNS script error must be an error or string") + if err := luamgr.ReadError(errorValue, "DNS script error must be an error or string"); err != nil { + return nil, 0, err } - ttl, ok := ttlValue.(lua.LNumber) - if !ok || ttl < 0 || ttl > math.MaxUint32 || math.Trunc(float64(ttl)) != float64(ttl) { - return nil, 0, errors.New("DNS script returned invalid TTL") + ttl, err := luamgr.ReadUint32(ttlValue, "DNS script returned invalid TTL") + if err != nil { + return nil, 0, err } if addresses == lua.LNil { return nil, 0, featureDNS.ErrEmptyResponse } - ud, ok := addresses.(*lua.LUserData) - if !ok { - return nil, 0, errors.New("DNS script IPs must be native IP slice userdata") - } - ips, ok := ud.Value.([]net.IP) - if !ok { - return nil, 0, errors.New("DNS script IPs must be native IP slice userdata") + ips, err := luamgr.ReadUserData[[]net.IP](addresses, "DNS script IPs must be native IP slice userdata") + if err != nil { + return nil, 0, err } if len(ips) == 0 { return nil, 0, featureDNS.ErrEmptyResponse } - return ips, uint32(ttl), nil + return ips, ttl, nil } diff --git a/app/router/lua.go b/app/router/lua.go index a424b869a..118b9a55e 100644 --- a/app/router/lua.go +++ b/app/router/lua.go @@ -5,6 +5,7 @@ import ( "strings" "github.com/xtls/xray-core/common/errors" + luamgr "github.com/xtls/xray-core/common/lua" "github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/features/routing" lua "github.com/yuin/gopher-lua" @@ -37,12 +38,12 @@ func (r *Router) RegisterLua(L *lua.LState) { balancer, found := (*r.balancers.Load())[string(tag)] if !found { L.Push(lua.LNil) - pushLuaError(L, errors.New("balancer ", tag, " not found")) + luamgr.PushError(L, errors.New("balancer ", tag, " not found")) return 2 } outboundTag, err := balancer.PickOutbound() L.Push(lua.LString(outboundTag)) - pushLuaError(L, err) + luamgr.PushError(L, err) return 2 })) @@ -51,7 +52,7 @@ func (r *Router) RegisterLua(L *lua.LState) { L.Push(lua.LNumber(pid)) L.Push(lua.LString(name)) L.Push(lua.LString(path)) - pushLuaError(L, err) + luamgr.PushError(L, err) return 4 })) @@ -75,13 +76,16 @@ func registerLuaContext(L *lua.LState) { methods := L.NewTable() L.SetFuncs(methods, map[string]lua.LGFunction{ "GetSourceIPs": func(L *lua.LState) int { - return pushLuaIPs(L, checkLuaContext(L).GetSourceIPs()) + luamgr.PushUserData(L, checkLuaContext(L).GetSourceIPs()) + return 1 }, "GetTargetIPs": func(L *lua.LState) int { - return pushLuaIPs(L, checkLuaContext(L).GetTargetIPs()) + luamgr.PushUserData(L, checkLuaContext(L).GetTargetIPs()) + return 1 }, "GetLocalIPs": func(L *lua.LState) int { - return pushLuaIPs(L, checkLuaContext(L).GetLocalIPs()) + luamgr.PushUserData(L, checkLuaContext(L).GetLocalIPs()) + return 1 }, "GetAttributes": func(L *lua.LState) int { values := L.NewUserData() @@ -102,23 +106,6 @@ func checkLuaContext(L *lua.LState) routing.Context { 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, routeCtx routing.Context) (string, string, error) { top := L.GetTop() @@ -142,36 +129,18 @@ func (r *Router) callLuaHook(L *lua.LState, routeCtx routing.Context) (string, s } 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 err := luamgr.ReadError(errorValue, "routing script error must be an error or string"); err != nil { + return "", "", err } - if tagValue == lua.LNil { - return "", "", nil + tag, err := luamgr.ReadOptionalString(tagValue, "routing script outboundTag must be a string or nil") + if err != nil || tag == "" { + return "", "", err } - tag, ok := tagValue.(lua.LString) - if !ok { - return "", "", errors.New("routing script outboundTag must be a string or nil") + ruleTag, err := luamgr.ReadOptionalString(ruleValue, "routing script ruleTag must be a string") + if err != nil { + return "", "", err } - 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 + return tag, 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 58dbe359b..1229dadfd 100644 --- a/app/router/lua_test.go +++ b/app/router/lua_test.go @@ -126,10 +126,16 @@ func TestLuaRouteResult(t *testing.T) { {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"}, } { t.Run(tc.name, func(t *testing.T) { @@ -137,6 +143,9 @@ func TestLuaRouteResult(t *testing.T) { 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{}) diff --git a/common/lua/utils.go b/common/lua/utils.go new file mode 100644 index 000000000..89ef7553e --- /dev/null +++ b/common/lua/utils.go @@ -0,0 +1,74 @@ +package lua + +import ( + "math" + + "github.com/xtls/xray-core/common/errors" + glua "github.com/yuin/gopher-lua" +) + +// PushUserData pushes a native Go value without copying it. +func PushUserData(L *glua.LState, value any) { + ud := L.NewUserData() + ud.Value = value + L.Push(ud) +} + +// PushError pushes nil or the original Go error as userdata. +func PushError(L *glua.LState, err error) { + if err == nil { + L.Push(glua.LNil) + return + } + PushUserData(L, err) +} + +// ReadUserData reads a native Go value of type T without copying it. +// Other Lua values or userdata containing a different type return invalidMessage. +func ReadUserData[T any](value glua.LValue, invalidMessage string) (T, error) { + if ud, ok := value.(*glua.LUserData); ok { + if result, ok := ud.Value.(T); ok { + return result, nil + } + } + var zero T + return zero, errors.New(invalidMessage) +} + +// ReadError accepts nil, a native Go error, or a Lua string. +// Native errors retain their identity; other values return invalidMessage. +func ReadError(value glua.LValue, invalidMessage string) error { + if value == glua.LNil { + return nil + } + if ud, ok := value.(*glua.LUserData); ok { + if err, ok := ud.Value.(error); ok { + return err + } + } + if message, ok := value.(glua.LString); ok { + return errors.New(string(message)) + } + return errors.New(invalidMessage) +} + +// ReadUint32 accepts only integral Lua numbers in the uint32 range. +func ReadUint32(value glua.LValue, invalidMessage string) (uint32, error) { + number, ok := value.(glua.LNumber) + if !ok || number < 0 || number > math.MaxUint32 || math.Trunc(float64(number)) != float64(number) { + return 0, errors.New(invalidMessage) + } + return uint32(number), nil +} + +// ReadOptionalString accepts a Lua string or nil, which becomes an empty string. +// It does not coerce other values to strings. +func ReadOptionalString(value glua.LValue, invalidMessage string) (string, error) { + if value == glua.LNil { + return "", nil + } + if result, ok := value.(glua.LString); ok { + return string(result), nil + } + return "", errors.New(invalidMessage) +} diff --git a/common/lua/utils_test.go b/common/lua/utils_test.go new file mode 100644 index 000000000..3a6f6a97a --- /dev/null +++ b/common/lua/utils_test.go @@ -0,0 +1,121 @@ +package lua + +import ( + "errors" + "math" + "strings" + "testing" + + glua "github.com/yuin/gopher-lua" +) + +func TestReadUint32(t *testing.T) { + for _, tc := range []struct { + name string + value glua.LValue + want uint32 + wantErr bool + }{ + {name: "zero", value: glua.LNumber(0)}, + {name: "integer", value: glua.LNumber(45), want: 45}, + {name: "maximum", value: glua.LNumber(math.MaxUint32), want: math.MaxUint32}, + {name: "fraction", value: glua.LNumber(1.5), wantErr: true}, + {name: "negative", value: glua.LNumber(-1), wantErr: true}, + {name: "overflow", value: glua.LNumber(math.MaxUint32 + 1), wantErr: true}, + {name: "NaN", value: glua.LNumber(math.NaN()), wantErr: true}, + {name: "positive infinity", value: glua.LNumber(math.Inf(1)), wantErr: true}, + {name: "negative infinity", value: glua.LNumber(math.Inf(-1)), wantErr: true}, + {name: "nil", value: glua.LNil, wantErr: true}, + {name: "numeric string", value: glua.LString("45"), wantErr: true}, + {name: "boolean", value: glua.LTrue, wantErr: true}, + } { + t.Run(tc.name, func(t *testing.T) { + got, err := ReadUint32(tc.value, "invalid number") + if got != tc.want || (err != nil) != tc.wantErr { + t.Fatalf("ReadUint32() = %d, %v; want %d, error %t", got, err, tc.want, tc.wantErr) + } + if err != nil && !strings.Contains(err.Error(), "invalid number") { + t.Fatalf("error = %v, want invalid number", err) + } + }) + } +} + +func TestReadOptionalString(t *testing.T) { + for _, tc := range []struct { + name string + value glua.LValue + want string + wantErr bool + }{ + {name: "nil", value: glua.LNil}, + {name: "empty", value: glua.LString("")}, + {name: "string", value: glua.LString("out"), want: "out"}, + {name: "number", value: glua.LNumber(1), wantErr: true}, + {name: "boolean", value: glua.LFalse, wantErr: true}, + } { + t.Run(tc.name, func(t *testing.T) { + got, err := ReadOptionalString(tc.value, "invalid string") + if got != tc.want || (err != nil) != tc.wantErr { + t.Fatalf("ReadOptionalString() = %q, %v; want %q, error %t", got, err, tc.want, tc.wantErr) + } + if err != nil && !strings.Contains(err.Error(), "invalid string") { + t.Fatalf("error = %v, want invalid string", err) + } + }) + } +} + +func TestUserDataRoundTrip(t *testing.T) { + L := glua.NewState() + defer L.Close() + want := []int{1, 2} + PushUserData(L, want) + if L.GetTop() != 1 { + t.Fatalf("stack top = %d, want 1", L.GetTop()) + } + got, err := ReadUserData[[]int](L.Get(-1), "invalid userdata") + if err != nil || len(got) != len(want) || &got[0] != &want[0] { + t.Fatalf("userdata = %v, %v; want original slice", got, err) + } + PushUserData(L, []int(nil)) + if got, err := ReadUserData[[]int](L.Get(-1), "invalid userdata"); err != nil || got != nil { + t.Fatalf("nil slice userdata = %v, %v", got, err) + } + for _, value := range []glua.LValue{glua.LNil, glua.LString("1"), L.NewTable(), L.Get(1)} { + if got, err := ReadUserData[int](value, "invalid userdata"); got != 0 || err == nil || !strings.Contains(err.Error(), "invalid userdata") { + t.Fatalf("ReadUserData(%v) = %d, %v; want invalid userdata", value, got, err) + } + } +} + +func TestErrorRoundTrip(t *testing.T) { + L := glua.NewState() + defer L.Close() + want := errors.New("upstream failed") + for _, err := range []error{nil, want} { + PushError(L, err) + if L.GetTop() != 1 { + t.Fatalf("stack top = %d, want 1", L.GetTop()) + } + if err == nil && L.Get(-1) != glua.LNil { + t.Fatalf("nil error pushed as %v", L.Get(-1)) + } + if got := ReadError(L.Get(-1), "invalid error"); got != err { + t.Fatalf("ReadError() = %v, want original error %v", got, err) + } + L.Pop(1) + } + for _, message := range []string{"script failed", ""} { + if err := ReadError(glua.LString(message), "invalid error"); err == nil || !strings.Contains(err.Error(), message) { + t.Fatalf("string error = %v, want %q", err, message) + } + } + wrong := L.NewUserData() + wrong.Value = "not a native error" + for _, value := range []glua.LValue{glua.LTrue, glua.LNumber(1), L.NewTable(), wrong, L.NewUserData()} { + if err := ReadError(value, "invalid error"); err == nil || !strings.Contains(err.Error(), "invalid error") { + t.Fatalf("ReadError(%v) = %v, want invalid error", value, err) + } + } +}