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
+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())
}
}