diff --git a/app/dns/lua.go b/app/dns/lua.go index 9e4a58ca2..4a459c74f 100644 --- a/app/dns/lua.go +++ b/app/dns/lua.go @@ -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), diff --git a/app/dns/script.go b/app/dns/script.go index e40d044bc..e46c669a8 100644 --- a/app/dns/script.go +++ b/app/dns/script.go @@ -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 } diff --git a/app/router/lua.go b/app/router/lua.go index 8e723b954..ac8f13ef5 100644 --- a/app/router/lua.go +++ b/app/router/lua.go @@ -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 { diff --git a/app/router/script.go b/app/router/script.go index db17e9d4f..6829d3308 100644 --- a/app/router/script.go +++ b/app/router/script.go @@ -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 } diff --git a/common/lua/lua.go b/common/lua/lua.go new file mode 100644 index 000000000..57e562f01 --- /dev/null +++ b/common/lua/lua.go @@ -0,0 +1,2 @@ +// Package lua provides shared GopherLua programs and state management for Xray scripts. +package lua diff --git a/common/lua/pool.go b/common/lua/pool.go index dff82f733..41ab213ab 100644 --- a/common/lua/pool.go +++ b/common/lua/pool.go @@ -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 { diff --git a/common/lua/pool_test.go b/common/lua/pool_test.go index 7b4d45bbe..1f1a80b77 100644 --- a/common/lua/pool_test.go +++ b/common/lua/pool_test.go @@ -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) } } diff --git a/common/lua/program.go b/common/lua/program.go index 95bf10cc1..38eb225fc 100644 --- a/common/lua/program.go +++ b/common/lua/program.go @@ -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) + } +} diff --git a/common/lua/program_test.go b/common/lua/program_test.go index d60b892fc..f4593f7e6 100644 --- a/common/lua/program_test.go +++ b/common/lua/program_test.go @@ -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()) + } +}