mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-03 22:36:55 +00:00
refactor: extract shared result validation and userdata helpers
This commit is contained in:
+14
-42
@@ -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
@@ -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)
|
||||
|
||||
@@ -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{})
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user