refactor: extract shared result validation and userdata helpers

This commit is contained in:
Meo597
2026-10-04 06:29:00 +08:00
parent 745526f14c
commit 73fb3e8f4a
5 changed files with 237 additions and 92 deletions
+14 -42
View File
@@ -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
}
+19 -50
View File
@@ -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)
+9
View File
@@ -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{})