mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-03 20:38:03 +03:00
lua: refactor to simplify state creation and pooled script execution
This commit is contained in:
+1
-1
@@ -157,7 +157,7 @@ func (s *DNS) CallLuaHook(L *lua.LState, ctx context.Context, domain string, opt
|
||||
}()
|
||||
fn := L.GetGlobal("HandleDNSQuery")
|
||||
if fn.Type() != lua.LTFunction {
|
||||
return nil, 0, errors.New("DNS script must define HandleDNSQuery(domain, ipv4, ipv6, fake)")
|
||||
return nil, 0, errors.New("DNS script must define HandleDNSQuery(...)")
|
||||
}
|
||||
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
|
||||
lua.LString(strings.ToLower(domain)), lua.LBool(option.IPv4Enable),
|
||||
|
||||
+19
-30
@@ -26,23 +26,19 @@ func newScriptEngine(path string, server *DNS) (*scriptEngine, error) {
|
||||
return nil, err
|
||||
}
|
||||
e := &scriptEngine{dns: server}
|
||||
e.pool, err = luamgr.NewPool(server.ctx, func(poolCtx context.Context) (*lua.LState, error) {
|
||||
initCtx, cancel := context.WithTimeout(poolCtx, scriptExecutionTimeout)
|
||||
defer cancel()
|
||||
L, err := program.NewState(initCtx, func(L *lua.LState) {
|
||||
e.pool, err = luamgr.NewPool(server.ctx, program.NewStateFactory(
|
||||
scriptExecutionTimeout,
|
||||
func(L *lua.LState) {
|
||||
geodata.RegisterLua(L)
|
||||
log.RegisterLua(L)
|
||||
server.RegisterLua(L)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if L.GetGlobal("HandleDNSQuery").Type() != lua.LTFunction {
|
||||
L.Close()
|
||||
return nil, errors.New("DNS script must define HandleDNSQuery(domain, ipv4, ipv6, fake)")
|
||||
}
|
||||
return L, nil
|
||||
})
|
||||
},
|
||||
func(L *lua.LState) error {
|
||||
if L.GetGlobal("HandleDNSQuery").Type() != lua.LTFunction {
|
||||
return errors.New("DNS script must define HandleDNSQuery(...)")
|
||||
}
|
||||
return nil
|
||||
}))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -54,20 +50,13 @@ func (e *scriptEngine) close() {
|
||||
e.pool.Close()
|
||||
}
|
||||
|
||||
func (e *scriptEngine) query(domain string, option dns.IPOption) ([]net.IP, uint32, error) {
|
||||
L, err := e.pool.Acquire()
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
reusable := false
|
||||
defer func() {
|
||||
e.pool.Release(L, reusable)
|
||||
}()
|
||||
queryCtx, cancel := context.WithTimeout(e.pool.Context(), scriptExecutionTimeout)
|
||||
defer cancel()
|
||||
ips, ttl, err := e.dns.CallLuaHook(L, queryCtx, domain, option)
|
||||
if err == nil {
|
||||
reusable = true
|
||||
}
|
||||
return ips, ttl, err
|
||||
func (e *scriptEngine) query(domain string, option dns.IPOption) (ips []net.IP, ttl uint32, err error) {
|
||||
err = e.pool.WithState(func(L *lua.LState) error {
|
||||
luaCtx, cancel := context.WithTimeout(e.pool.Context(), scriptExecutionTimeout)
|
||||
defer cancel()
|
||||
var luaErr error
|
||||
ips, ttl, luaErr = e.dns.CallLuaHook(L, luaCtx, domain, option)
|
||||
return luaErr
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
+2
-2
@@ -121,9 +121,9 @@ func pushLuaError(L *lua.LState, err error) {
|
||||
}
|
||||
|
||||
// CallLuaHook invokes HandleRoute in the supplied state.
|
||||
func (r *Router) CallLuaHook(L *lua.LState, ctx context.Context, routeCtx routing.Context) (string, string, error) {
|
||||
func (r *Router) CallLuaHook(L *lua.LState, luaCtx context.Context, routeCtx routing.Context) (string, string, error) {
|
||||
previous, top := L.Context(), L.GetTop()
|
||||
L.SetContext(ctx)
|
||||
L.SetContext(luaCtx)
|
||||
defer func() {
|
||||
L.SetTop(top)
|
||||
if previous == nil {
|
||||
|
||||
+18
-26
@@ -27,24 +27,20 @@ func newScriptEngine(path string, router *Router) (*scriptEngine, error) {
|
||||
return nil, err
|
||||
}
|
||||
e := &scriptEngine{router: router}
|
||||
e.pool, err = luamgr.NewPool(router.ctx, func(poolCtx context.Context) (*lua.LState, error) {
|
||||
initCtx, cancel := context.WithTimeout(poolCtx, scriptExecutionTimeout)
|
||||
defer cancel()
|
||||
L, err := program.NewState(initCtx, func(L *lua.LState) {
|
||||
e.pool, err = luamgr.NewPool(router.ctx, program.NewStateFactory(
|
||||
scriptExecutionTimeout,
|
||||
func(L *lua.LState) {
|
||||
geodata.RegisterLua(L)
|
||||
log.RegisterLua(L)
|
||||
router.RegisterLua(L)
|
||||
dns.RegisterLua(L, router.dns)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if L.GetGlobal("HandleRoute").Type() != lua.LTFunction {
|
||||
L.Close()
|
||||
return nil, errors.New("routing script must define HandleRoute(...)")
|
||||
}
|
||||
return L, nil
|
||||
})
|
||||
},
|
||||
func(L *lua.LState) error {
|
||||
if L.GetGlobal("HandleRoute").Type() != lua.LTFunction {
|
||||
return errors.New("routing script must define HandleRoute(...)")
|
||||
}
|
||||
return nil
|
||||
}))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -57,21 +53,17 @@ func (e *scriptEngine) close() {
|
||||
}
|
||||
|
||||
func (e *scriptEngine) pickRoute(ctx routing.Context) (routing.Route, error) {
|
||||
L, err := e.pool.Acquire()
|
||||
var tag, ruleTag string
|
||||
err := e.pool.WithState(func(L *lua.LState) error {
|
||||
luaCtx, cancel := context.WithTimeout(e.pool.Context(), scriptExecutionTimeout)
|
||||
defer cancel()
|
||||
var luaErr error
|
||||
tag, ruleTag, luaErr = e.router.CallLuaHook(L, luaCtx, ctx)
|
||||
return luaErr
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
reusable := false
|
||||
defer func() {
|
||||
e.pool.Release(L, reusable)
|
||||
}()
|
||||
callCtx, cancel := context.WithTimeout(e.pool.Context(), scriptExecutionTimeout)
|
||||
defer cancel()
|
||||
tag, ruleTag, err := e.router.CallLuaHook(L, callCtx, ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
reusable = true
|
||||
if tag == "" {
|
||||
return nil, common.ErrNoClue
|
||||
}
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
// Package lua provides shared GopherLua programs and state management for Xray scripts.
|
||||
package lua
|
||||
+19
-32
@@ -2,7 +2,6 @@ package lua
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
@@ -10,13 +9,9 @@ import (
|
||||
|
||||
const maxIdleStates = 16
|
||||
|
||||
// LStateFactory must initialize a state fully and observe ctx while doing so.
|
||||
// The pool owns any non-nil state it returns, even when it also returns an error.
|
||||
type LStateFactory func(ctx context.Context) (*glua.LState, error)
|
||||
|
||||
// Pool lends each state to one caller at a time. It grows on contention and
|
||||
// keeps up to maxIdleStates idle states until Close. Callers decide whether a
|
||||
// state is reusable.
|
||||
// keeps up to maxIdleStates idle states until Close. Acquire/Release callers
|
||||
// decide reusability; WithState uses its callback's error.
|
||||
type Pool struct {
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
@@ -29,25 +24,12 @@ type Pool struct {
|
||||
closed bool
|
||||
}
|
||||
|
||||
// NewPool initializes one state before returning, so top-level errors surface at startup.
|
||||
// NewPool tests the factory by creating one state during initialization.
|
||||
func NewPool(ctx context.Context, factory LStateFactory) (*Pool, error) {
|
||||
poolCtx, cancel := context.WithCancel(ctx)
|
||||
|
||||
// Create one state now to catch factory errors at startup.
|
||||
state, err := factory(poolCtx)
|
||||
if err != nil {
|
||||
cancel()
|
||||
if state != nil {
|
||||
state.Close()
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
if state == nil {
|
||||
cancel()
|
||||
return nil, errors.New("Lua state factory returned nil")
|
||||
}
|
||||
if err := poolCtx.Err(); err != nil {
|
||||
state.Close()
|
||||
cancel()
|
||||
return nil, err
|
||||
}
|
||||
@@ -83,18 +65,7 @@ func (p *Pool) Acquire() (*glua.LState, error) {
|
||||
// for a Release instead of creating another state; allow the wait to be
|
||||
// cancelled by the caller or by Close.
|
||||
state, err := p.factory(p.ctx)
|
||||
if err == nil && state == nil {
|
||||
err = errors.New("Lua state factory returned nil")
|
||||
}
|
||||
if err != nil {
|
||||
if state != nil {
|
||||
state.Close()
|
||||
}
|
||||
p.active.Done()
|
||||
return nil, err
|
||||
}
|
||||
if err := p.ctx.Err(); err != nil {
|
||||
state.Close()
|
||||
p.active.Done()
|
||||
return nil, err
|
||||
}
|
||||
@@ -102,6 +73,22 @@ func (p *Pool) Acquire() (*glua.LState, error) {
|
||||
return state, nil
|
||||
}
|
||||
|
||||
// WithState runs work on an exclusive state and releases it afterward. A state
|
||||
// is reusable only when work succeeds; a panic closes it before propagating.
|
||||
func (p *Pool) WithState(work func(*glua.LState) error) error {
|
||||
state, err := p.Acquire()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
reusable := false
|
||||
defer func() {
|
||||
p.Release(state, reusable)
|
||||
}()
|
||||
err = work(state)
|
||||
reusable = err == nil
|
||||
return err
|
||||
}
|
||||
|
||||
// Release returns a healthy state to the pool and closes a failed or cancelled one.
|
||||
func (p *Pool) Release(state *glua.LState, reusable bool) {
|
||||
if reusable {
|
||||
|
||||
+7
-10
@@ -9,25 +9,22 @@ import (
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
func TestPoolFactoryFailureClosesReturnedState(t *testing.T) {
|
||||
func TestPoolFactoryFailure(t *testing.T) {
|
||||
failure := errors.New("factory failed")
|
||||
state := glua.NewState()
|
||||
_, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) {
|
||||
return state, failure
|
||||
return nil, failure
|
||||
})
|
||||
if !errors.Is(err, failure) || !state.IsClosed() {
|
||||
t.Fatalf("NewPool error = %v, state closed = %t", err, state.IsClosed())
|
||||
if !errors.Is(err, failure) {
|
||||
t.Fatalf("NewPool error = %v, want %v", err, failure)
|
||||
}
|
||||
|
||||
var failedState *glua.LState
|
||||
calls := 0
|
||||
pool, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) {
|
||||
calls++
|
||||
if calls == 1 {
|
||||
return glua.NewState(), nil
|
||||
}
|
||||
failedState = glua.NewState()
|
||||
return failedState, failure
|
||||
return nil, failure
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -39,8 +36,8 @@ func TestPoolFactoryFailureClosesReturnedState(t *testing.T) {
|
||||
}
|
||||
defer pool.Release(borrowed, true)
|
||||
_, err = pool.Acquire()
|
||||
if !errors.Is(err, failure) || !failedState.IsClosed() {
|
||||
t.Fatalf("Acquire error = %v, state closed = %t", err, failedState.IsClosed())
|
||||
if !errors.Is(err, failure) {
|
||||
t.Fatalf("Acquire error = %v, want %v", err, failure)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+34
-13
@@ -1,10 +1,10 @@
|
||||
// Package lua provides shared GopherLua programs and state management for Xray scripts.
|
||||
package lua
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"os"
|
||||
"time"
|
||||
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
"github.com/yuin/gopher-lua/parse"
|
||||
@@ -15,6 +15,10 @@ type Program struct {
|
||||
proto *glua.FunctionProto
|
||||
}
|
||||
|
||||
// LStateFactory returns a fully initialized state or nil and an error.
|
||||
// Implementations must close partial states on failure; callers own successful states.
|
||||
type LStateFactory func(context.Context) (*glua.LState, error)
|
||||
|
||||
// CompileFile reads and compiles a Lua file once.
|
||||
func CompileFile(path string) (*Program, error) {
|
||||
f, err := os.Open(path)
|
||||
@@ -33,24 +37,41 @@ func CompileFile(path string) (*Program, error) {
|
||||
return &Program{proto: proto}, nil
|
||||
}
|
||||
|
||||
// NewState creates a VM, makes modules available, and executes the file top level.
|
||||
// Module loaders run only when Lua calls require. Each state gets its own globals.
|
||||
// The caller owns the returned state.
|
||||
func (p *Program) NewState(ctx context.Context, register func(*glua.LState)) (*glua.LState, error) {
|
||||
// NewState creates a state, runs register, executes the program under ctx, and
|
||||
// runs validate. It removes the initialization context before returning a state
|
||||
// owned by the caller.
|
||||
func (p *Program) NewState(ctx context.Context, register func(*glua.LState), validate func(*glua.LState) error) (*glua.LState, error) {
|
||||
L := glua.NewState()
|
||||
valid := false
|
||||
defer func() {
|
||||
if !valid {
|
||||
L.Close()
|
||||
}
|
||||
}()
|
||||
L.SetContext(ctx)
|
||||
defer L.RemoveContext()
|
||||
if register != nil {
|
||||
register(L)
|
||||
}
|
||||
L.SetContext(ctx)
|
||||
L.Push(L.NewFunctionFromProto(p.proto))
|
||||
err := L.PCall(0, 0, nil)
|
||||
L.RemoveContext()
|
||||
if err == nil {
|
||||
err = ctx.Err()
|
||||
}
|
||||
if err != nil {
|
||||
L.Close()
|
||||
// Execute the Lua script's top level.
|
||||
if err := L.PCall(0, 0, nil); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if validate != nil {
|
||||
if err := validate(L); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
valid = true
|
||||
return L, nil
|
||||
}
|
||||
|
||||
// NewStateFactory returns a factory that gives each state an initialization timeout.
|
||||
func (p *Program) NewStateFactory(initTimeout time.Duration, register func(*glua.LState), validate func(*glua.LState) error) LStateFactory {
|
||||
return func(ctx context.Context) (*glua.LState, error) {
|
||||
initCtx, cancel := context.WithTimeout(ctx, initTimeout)
|
||||
defer cancel()
|
||||
return p.NewState(initCtx, register, validate)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,6 +2,7 @@ package lua
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
@@ -18,13 +19,13 @@ func TestProgramStatesAreIndependent(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
first, err := program.NewState(context.Background(), nil)
|
||||
first, err := program.NewState(context.Background(), nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer first.Close()
|
||||
first.SetGlobal("value", glua.LNumber(42))
|
||||
second, err := program.NewState(context.Background(), nil)
|
||||
second, err := program.NewState(context.Background(), nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -45,7 +46,7 @@ func TestProgramInitializationObservesCancellation(t *testing.T) {
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
state, err := program.NewState(ctx, nil)
|
||||
state, err := program.NewState(ctx, nil, nil)
|
||||
if err == nil || state != nil {
|
||||
if state != nil {
|
||||
state.Close()
|
||||
@@ -53,3 +54,23 @@ func TestProgramInitializationObservesCancellation(t *testing.T) {
|
||||
t.Fatalf("NewState with canceled context = %v, %v; want nil state and error", state, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewStateClosesFailedValidation(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "state.lua")
|
||||
if err := os.WriteFile(path, []byte("value = 1"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
program, err := CompileFile(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
wantErr := errors.New("invalid script")
|
||||
var checked *glua.LState
|
||||
L, err := program.NewState(context.Background(), nil, func(L *glua.LState) error {
|
||||
checked = L
|
||||
return wantErr
|
||||
})
|
||||
if L != nil || !errors.Is(err, wantErr) || checked == nil || !checked.IsClosed() {
|
||||
t.Fatalf("state = %v, error = %v, checked state closed = %t", L, err, checked != nil && checked.IsClosed())
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user