lua: refactor to simplify state creation and pooled script execution

This commit is contained in:
Meo597
2026-10-03 21:57:54 +08:00
parent 5afe260f10
commit 2610e57ecf
9 changed files with 126 additions and 117 deletions
+1 -1
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
}
+2
View File
@@ -0,0 +1,2 @@
// Package lua provides shared GopherLua programs and state management for Xray scripts.
package lua
+19 -32
View File
@@ -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
View File
@@ -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
View File
@@ -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)
}
}
+24 -3
View File
@@ -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())
}
}