diff --git a/app/dns/lua.go b/app/dns/lua.go index 4a459c74f..deb96b394 100644 --- a/app/dns/lua.go +++ b/app/dns/lua.go @@ -142,19 +142,11 @@ func newLuaClientQuery(L *lua.LState, client featureDNS.Client) *lua.LFunction { }) } -// CallLuaHook invokes HandleDNSQuery in the supplied state. +// callLuaHook invokes HandleDNSQuery in the supplied state. // Returned slices and IP bytes may share storage with DNS caches or matcher inputs. -func (s *DNS) CallLuaHook(L *lua.LState, ctx context.Context, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) { - previous, top := L.Context(), L.GetTop() - L.SetContext(ctx) - defer func() { - L.SetTop(top) - if previous == nil { - L.RemoveContext() - } else { - L.SetContext(previous) - } - }() +func (s *DNS) callLuaHook(L *lua.LState, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) { + top := L.GetTop() + defer L.SetTop(top) fn := L.GetGlobal("HandleDNSQuery") if fn.Type() != lua.LTFunction { return nil, 0, errors.New("DNS script must define HandleDNSQuery(...)") diff --git a/app/dns/lua_test.go b/app/dns/lua_test.go index bf6c2feea..0b59dd4e4 100644 --- a/app/dns/lua_test.go +++ b/app/dns/lua_test.go @@ -84,14 +84,15 @@ func TestCallLuaHookCancellation(t *testing.T) { if err := L.DoString(`function HandleDNSQuery(domain, ipv4, ipv6, fake) while true do end end`); err != nil { t.Fatal(err) } - ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) - defer cancel() - _, _, err := (&DNS{}).CallLuaHook(L, ctx, "example.com", featureDNS.IPOption{IPv4Enable: true}) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + L.SetContext(ctx) + _, _, err := (&DNS{}).callLuaHook(L, "example.com", featureDNS.IPOption{IPv4Enable: true}) if err == nil { t.Fatal("CallLuaHook did not stop after context cancellation") } - if L.Context() != nil { - t.Fatal("CallLuaHook left the canceled context on the Lua state") + if L.Context() != ctx { + t.Fatal("CallLuaHook changed the Lua state's context") } } @@ -111,12 +112,12 @@ func TestCallLuaHookNormalizesDomain(t *testing.T) { t.Fatal(err) } s := &DNS{} - if _, _, err := s.CallLuaHook(L, context.Background(), "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil { + if _, _, err := s.callLuaHook(L, "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil { t.Fatal(err) } } -func TestCallLuaHookRestoresState(t *testing.T) { +func TestCallLuaHookRestoresStack(t *testing.T) { for _, tc := range []struct { name string body string @@ -134,16 +135,13 @@ func TestCallLuaHookRestoresState(t *testing.T) { if err := L.DoString("function HandleDNSQuery() " + tc.body + " end"); err != nil { t.Fatal(err) } - previous, cancel := context.WithCancel(context.Background()) - defer cancel() - L.SetContext(previous) L.Push(lua.LTrue) - _, _, err := (&DNS{}).CallLuaHook(L, context.Background(), "example.com", featureDNS.IPOption{IPv4Enable: true}) + _, _, err := (&DNS{}).callLuaHook(L, "example.com", featureDNS.IPOption{IPv4Enable: true}) if (err != nil) != tc.wantErr { t.Fatalf("hook error = %v, want error %t", err, tc.wantErr) } - if L.Context() != previous || L.GetTop() != 1 || L.Get(1) != lua.LTrue { - t.Fatal("hook did not restore the previous context and stack") + if L.GetTop() != 1 || L.Get(1) != lua.LTrue { + t.Fatal("hook did not restore the stack") } }) } @@ -170,7 +168,8 @@ end `); err != nil { t.Fatal(err) } - got, ttl, err := server.CallLuaHook(L, context.Background(), "example.com", option) + L.SetContext(context.Background()) + got, ttl, err := server.callLuaHook(L, "example.com", option) if err != nil || ttl != 60 || len(got) != 1 || !got[0].Equal(ips[0]) { t.Fatalf("server query = %v, TTL %d, %v", got, ttl, err) } @@ -189,7 +188,6 @@ func TestLuaDNSClientQuery(t *testing.T) { L := lua.NewState() defer L.Close() L.SetContext(context.Background()) - defer L.RemoveContext() geodata.RegisterLua(L) want := []net.IP{{127, 0, 0, 1}} client := &luaDNSClient{lookup: func(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) { @@ -218,7 +216,6 @@ func TestLuaDNSLocalClient(t *testing.T) { L := lua.NewState() defer L.Close() L.SetContext(context.Background()) - defer L.RemoveContext() RegisterLua(L, localdns.New()) if err := L.DoString(` local dns = require("xray.dns") @@ -268,12 +265,13 @@ end } ctx := context.Background() + L.SetContext(ctx) for _, bench := range []struct { name string query func() ([]net.IP, uint32, error) }{ {"direct", func() ([]net.IP, uint32, error) { return client.QueryIP(ctx, "example.com", option) }}, - {"lua_hook", func() ([]net.IP, uint32, error) { return server.CallLuaHook(L, ctx, "example.com", option) }}, + {"lua_hook", func() ([]net.IP, uint32, error) { return server.callLuaHook(L, "example.com", option) }}, } { b.Run(bench.name, func(b *testing.B) { b.ReportAllocs() diff --git a/app/dns/script.go b/app/dns/script.go index e46c669a8..6391a1fd9 100644 --- a/app/dns/script.go +++ b/app/dns/script.go @@ -1,7 +1,6 @@ package dns import ( - "context" "time" "github.com/xtls/xray-core/common/errors" @@ -13,7 +12,7 @@ import ( lua "github.com/yuin/gopher-lua" ) -const scriptExecutionTimeout = 10 * time.Second +const scriptExecutionTimeout = 6 * time.Second type scriptEngine struct { dns *DNS @@ -26,8 +25,8 @@ func newScriptEngine(path string, server *DNS) (*scriptEngine, error) { return nil, err } e := &scriptEngine{dns: server} - e.pool, err = luamgr.NewPool(server.ctx, program.NewStateFactory( - scriptExecutionTimeout, + e.pool, err = luamgr.NewPool(server.ctx, scriptExecutionTimeout, program.NewStateFactory( + scriptExecutionTimeout*20, func(L *lua.LState) { geodata.RegisterLua(L) log.RegisterLua(L) @@ -51,12 +50,10 @@ func (e *scriptEngine) close() { } 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 + err = e.pool.WithState(nil, 0, func(L *lua.LState) error { + var hookErr error + ips, ttl, hookErr = e.dns.callLuaHook(L, domain, option) + return hookErr }) return } diff --git a/app/router/lua.go b/app/router/lua.go index ac8f13ef5..a424b869a 100644 --- a/app/router/lua.go +++ b/app/router/lua.go @@ -1,7 +1,6 @@ package router import ( - "context" "runtime" "strings" @@ -120,18 +119,10 @@ func pushLuaError(L *lua.LState, err error) { L.Push(value) } -// CallLuaHook invokes HandleRoute in the supplied state. -func (r *Router) CallLuaHook(L *lua.LState, luaCtx context.Context, routeCtx routing.Context) (string, string, error) { - previous, top := L.Context(), L.GetTop() - L.SetContext(luaCtx) - defer func() { - L.SetTop(top) - if previous == nil { - L.RemoveContext() - } else { - L.SetContext(previous) - } - }() +// callLuaHook invokes HandleRoute in the supplied state. +func (r *Router) callLuaHook(L *lua.LState, routeCtx routing.Context) (string, string, error) { + top := L.GetTop() + defer L.SetTop(top) fn := L.GetGlobal("HandleRoute") if fn.Type() != lua.LTFunction { return "", "", errors.New("routing script must define HandleRoute(...)") diff --git a/app/router/lua_test.go b/app/router/lua_test.go index 518acdb83..58dbe359b 100644 --- a/app/router/lua_test.go +++ b/app/router/lua_test.go @@ -88,7 +88,7 @@ function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort, end`) ctx := newLuaRouteTestContext() - tag, rule, err := r.CallLuaHook(L, context.Background(), ctx) + tag, rule, err := r.callLuaHook(L, ctx) if err != nil || tag != "out" || rule != "rule" { t.Fatalf("hook = %q, %q, %v", tag, rule, err) } @@ -137,11 +137,9 @@ func TestLuaRouteResult(t *testing.T) { value := L.NewUserData() value.Value = nativeErr L.SetGlobal("nativeError", value) - previous := context.WithValue(context.Background(), struct{}{}, true) - L.SetContext(previous) L.Push(lua.LTrue) - tag, rule, err := r.CallLuaHook(L, context.Background(), &routing_session.Context{}) + tag, rule, err := r.callLuaHook(L, &routing_session.Context{}) if tag != tc.tag || rule != tc.rule { t.Fatalf("result = %q, %q, %v", tag, rule, err) } @@ -157,8 +155,8 @@ func TestLuaRouteResult(t *testing.T) { case err != nil: t.Fatal(err) } - if L.Context() != previous || L.GetTop() != 1 || L.Get(1) != lua.LTrue { - t.Fatal("hook did not restore the previous context and stack") + if L.GetTop() != 1 || L.Get(1) != lua.LTrue { + t.Fatal("hook did not restore the stack") } }) } @@ -168,10 +166,11 @@ func TestLuaRouteCancellation(t *testing.T) { r, L := newLuaRouterState(t, `function HandleRoute() while true do end end`) ctx, cancel := context.WithCancel(context.Background()) cancel() - if _, _, err := r.CallLuaHook(L, ctx, &routing_session.Context{}); err == nil { + L.SetContext(ctx) + if _, _, err := r.callLuaHook(L, &routing_session.Context{}); err == nil { t.Fatal("CallLuaHook did not stop after context cancellation") } - if L.Context() != nil || L.GetTop() != 0 { + if L.Context() != ctx || L.GetTop() != 0 { t.Fatal("CallLuaHook did not restore the Lua state") } } @@ -253,7 +252,7 @@ end b.Fatal(err) } - ctx := context.Background() + L.SetContext(context.Background()) routeCtx := newLuaRouteTestContext() for _, benchmark := range []struct { name string @@ -267,7 +266,7 @@ end return route.GetOutboundTag(), route.GetRuleTag(), nil }}, {"lua_hook", func() (string, string, error) { - return r.CallLuaHook(L, ctx, routeCtx) + return r.callLuaHook(L, routeCtx) }}, } { b.Run(benchmark.name, func(b *testing.B) { diff --git a/app/router/script.go b/app/router/script.go index 6829d3308..da23d3bd7 100644 --- a/app/router/script.go +++ b/app/router/script.go @@ -1,7 +1,6 @@ package router import ( - "context" "time" "github.com/xtls/xray-core/app/dns" @@ -14,7 +13,7 @@ import ( lua "github.com/yuin/gopher-lua" ) -const scriptExecutionTimeout = 10 * time.Second +const scriptExecutionTimeout = 6 * time.Second type scriptEngine struct { router *Router @@ -27,8 +26,8 @@ func newScriptEngine(path string, router *Router) (*scriptEngine, error) { return nil, err } e := &scriptEngine{router: router} - e.pool, err = luamgr.NewPool(router.ctx, program.NewStateFactory( - scriptExecutionTimeout, + e.pool, err = luamgr.NewPool(router.ctx, scriptExecutionTimeout, program.NewStateFactory( + scriptExecutionTimeout*20, func(L *lua.LState) { geodata.RegisterLua(L) log.RegisterLua(L) @@ -54,12 +53,10 @@ func (e *scriptEngine) close() { func (e *scriptEngine) pickRoute(ctx routing.Context) (routing.Route, error) { 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 + err := e.pool.WithState(nil, 0, func(L *lua.LState) error { + var hookErr error + tag, ruleTag, hookErr = e.router.callLuaHook(L, ctx) + return hookErr }) if err != nil { return nil, err diff --git a/common/lua/pool.go b/common/lua/pool.go index 974850c7c..5457ebcf6 100644 --- a/common/lua/pool.go +++ b/common/lua/pool.go @@ -2,7 +2,9 @@ package lua import ( "context" + "errors" "sync" + "time" glua "github.com/yuin/gopher-lua" ) @@ -13,8 +15,9 @@ const maxIdleStates = 16 // 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 + ctx context.Context + cancel context.CancelFunc + timeout time.Duration factory LStateFactory idle []*glua.LState @@ -26,7 +29,11 @@ type Pool struct { } // NewPool tests the factory by creating one state during initialization. -func NewPool(ctx context.Context, factory LStateFactory) (*Pool, error) { +func NewPool(ctx context.Context, timeout time.Duration, factory LStateFactory) (*Pool, error) { + if timeout <= 0 { + return nil, errors.New("Lua pool timeout must be positive") + } + poolCtx, cancel := context.WithCancel(ctx) state, err := factory(poolCtx) @@ -35,20 +42,26 @@ func NewPool(ctx context.Context, factory LStateFactory) (*Pool, error) { return nil, err } - return &Pool{ctx: poolCtx, cancel: cancel, factory: factory, idle: []*glua.LState{state}, top: state.GetTop()}, nil -} - -// Context is cancelled by Close. Query contexts should derive from it. -func (p *Pool) Context() context.Context { - return p.ctx + return &Pool{ctx: poolCtx, cancel: cancel, timeout: timeout, factory: factory, idle: []*glua.LState{state}, top: state.GetTop()}, nil } // Acquire returns an initialized exclusive state, growing the pool if necessary. -func (p *Pool) Acquire() (*glua.LState, error) { +// ctx is passed to the factory for state creation; nil uses the pool context. +func (p *Pool) Acquire(ctx context.Context) (*glua.LState, error) { p.mu.Lock() - if p.closed || p.ctx.Err() != nil { + if p.closed { p.mu.Unlock() - return nil, p.ctx.Err() + return nil, errors.New("Lua pool is closed") + } + if err := p.ctx.Err(); err != nil { + p.mu.Unlock() + return nil, err + } + if ctx == nil { + ctx = p.ctx + } else if err := ctx.Err(); err != nil { + p.mu.Unlock() + return nil, err } p.active.Add(1) @@ -65,7 +78,7 @@ func (p *Pool) Acquire() (*glua.LState, error) { // TODO: Limit the total number of states. When the limit is reached, wait // 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) + state, err := p.factory(ctx) if err != nil { p.active.Done() return nil, err @@ -74,15 +87,24 @@ 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() +// WithState runs work on an exclusive state and releases it afterward. +// Nil ctx and zero timeout use pool defaults. The timeout starts after acquisition. +func (p *Pool) WithState(ctx context.Context, timeout time.Duration, work func(*glua.LState) error) error { + state, err := p.Acquire(ctx) if err != nil { return err } + if ctx == nil { + ctx = p.ctx + } + if timeout == 0 { + timeout = p.timeout + } + ctx, cancel := context.WithTimeout(ctx, timeout) + state.SetContext(ctx) reusable := false defer func() { + cancel() p.Release(state, reusable) }() err = work(state) @@ -90,9 +112,10 @@ func (p *Pool) WithState(work func(*glua.LState) error) error { return err } -// Release returns a healthy state to the pool and closes a failed or cancelled one. +// Release resets a state for reuse or closes it. func (p *Pool) Release(state *glua.LState, reusable bool) { if reusable { + state.RemoveContext() state.SetTop(p.top) p.mu.Lock() if !p.closed && p.ctx.Err() == nil && len(p.idle) < maxIdleStates { @@ -110,7 +133,7 @@ func (p *Pool) Release(state *glua.LState, reusable bool) { p.active.Done() } -// Close cancels active work, closes idle states, and waits for borrowed states. +// Close cancels the pool context, closes idle states, and waits for borrowed states. func (p *Pool) Close() { p.mu.Lock() if !p.closed { diff --git a/common/lua/pool_test.go b/common/lua/pool_test.go index 1f1a80b77..49bb25bbd 100644 --- a/common/lua/pool_test.go +++ b/common/lua/pool_test.go @@ -9,186 +9,458 @@ import ( glua "github.com/yuin/gopher-lua" ) +func newTestPool(t testing.TB, ctx context.Context, timeout time.Duration, factory LStateFactory) *Pool { + t.Helper() + pool, err := NewPool(ctx, timeout, factory) + if err != nil { + t.Fatal(err) + } + t.Cleanup(pool.Close) + return pool +} + +func assertPoolCloseBlocked(t *testing.T, done <-chan struct{}) { + t.Helper() + select { + case <-done: + t.Fatal("Close returned while work was still active") + case <-time.After(20 * time.Millisecond): + } +} + +func TestPoolTimeoutValidation(t *testing.T) { + for _, tc := range []struct { + name string + timeout time.Duration + wantErr bool + }{ + {"zero", 0, true}, + {"negative", -time.Nanosecond, true}, + {"positive", time.Nanosecond, false}, + } { + t.Run(tc.name, func(t *testing.T) { + called := false + pool, err := NewPool(context.Background(), tc.timeout, func(context.Context) (*glua.LState, error) { + called = true + return glua.NewState(), nil + }) + if pool != nil { + t.Cleanup(pool.Close) + } + if (err != nil) != tc.wantErr { + t.Fatalf("NewPool error = %v, want error %t", err, tc.wantErr) + } + if tc.wantErr && (pool != nil || called) { + t.Fatal("invalid timeout created a pool or called the factory") + } + }) + } +} + func TestPoolFactoryFailure(t *testing.T) { failure := errors.New("factory failed") - _, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) { + _, err := NewPool(context.Background(), time.Second, func(context.Context) (*glua.LState, error) { return nil, failure }) if !errors.Is(err, failure) { - t.Fatalf("NewPool error = %v, want %v", err, failure) + t.Fatalf("NewPool error = %v, want original factory error", err) } calls := 0 - pool, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) { + pool := newTestPool(t, context.Background(), time.Second, func(context.Context) (*glua.LState, error) { calls++ if calls == 1 { return glua.NewState(), nil } return nil, failure }) + state, err := pool.Acquire(nil) if err != nil { t.Fatal(err) } - defer pool.Close() - borrowed, err := pool.Acquire() - if err != nil { - t.Fatal(err) - } - defer pool.Release(borrowed, true) - _, err = pool.Acquire() + defer pool.Release(state, true) + err = pool.WithState(nil, 0, func(*glua.LState) error { + t.Error("work ran after factory failure") + return nil + }) if !errors.Is(err, failure) { - t.Fatalf("Acquire error = %v, want %v", err, failure) + t.Fatalf("WithState error = %v, want original factory error", err) } } func TestPoolReusesStatesAndLimitsIdle(t *testing.T) { created := 0 - pool, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) { + pool := newTestPool(t, context.Background(), time.Second, func(context.Context) (*glua.LState, error) { created++ return glua.NewState(), nil }) - if err != nil { - t.Fatal(err) - } - defer pool.Close() - - states := make([]*glua.LState, maxIdleStates+3) - for i := range states { - states[i], err = pool.Acquire() + var borrowed []*glua.LState + defer func() { + for _, state := range borrowed { + pool.Release(state, false) + } + }() + for range maxIdleStates + 3 { + state, err := pool.Acquire(nil) if err != nil { t.Fatal(err) } + borrowed = append(borrowed, state) + state.SetContext(context.Background()) } + states := borrowed for _, state := range states { pool.Release(state, true) } - for i, state := range states { - if got, want := state.IsClosed(), i >= maxIdleStates; got != want { - t.Fatalf("state %d closed = %t, want %t", i, got, want) + borrowed = nil + open := 0 + for _, state := range states { + if !state.IsClosed() { + if state.Context() != nil { + t.Fatal("Release left a context on a reusable state") + } + open++ } } - borrowed, err := pool.Acquire() - if err != nil { + if open != maxIdleStates { + t.Fatalf("retained %d states, want %d", open, maxIdleStates) + } + if err := pool.WithState(nil, 0, func(*glua.LState) error { return nil }); err != nil { t.Fatal(err) } if created != len(states) { - t.Fatalf("Acquire created %d states, want %d", created, len(states)) + t.Fatalf("created %d states, want %d", created, len(states)) + } + pool.Close() + for _, state := range states { + if !state.IsClosed() { + t.Fatal("Close left an idle state open") + } } - pool.Release(borrowed, true) } -func TestPoolCloseCancelsAndWaitsForBorrowedState(t *testing.T) { - pool, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) { +func TestPoolWithStateOptions(t *testing.T) { + key := struct{}{} + parent := context.WithValue(context.Background(), key, "pool") + caller := context.WithValue(context.Background(), key, "caller") + pool := newTestPool(t, parent, time.Second, func(context.Context) (*glua.LState, error) { return glua.NewState(), nil }) - if err != nil { - t.Fatal(err) + for _, tc := range []struct { + name string + ctx context.Context + timeout time.Duration + wantValue string + wantTimeout time.Duration + }{ + {"defaults", nil, 0, "pool", time.Second}, + {"context", caller, 0, "caller", time.Second}, + {"timeout", nil, 2 * time.Second, "pool", 2 * time.Second}, + {"both", caller, 2 * time.Second, "caller", 2 * time.Second}, + } { + t.Run(tc.name, func(t *testing.T) { + started := time.Now() + err := pool.WithState(tc.ctx, tc.timeout, func(L *glua.LState) error { + ctx := L.Context() + if ctx.Value(key) != tc.wantValue { + t.Errorf("context value = %v, want %q", ctx.Value(key), tc.wantValue) + } + deadline, ok := ctx.Deadline() + if !ok || deadline.Before(started.Add(tc.wantTimeout)) || deadline.After(time.Now().Add(tc.wantTimeout)) { + t.Errorf("deadline = %v, want timeout %v", deadline, tc.wantTimeout) + } + return nil + }) + if err != nil { + t.Fatal(err) + } + }) } - state, err := pool.Acquire() - if err != nil { - t.Fatal(err) +} + +func TestPoolFactoryContext(t *testing.T) { + caller, cancel := context.WithTimeout(context.Background(), time.Minute) + defer cancel() + for _, tc := range []struct { + name string + ctx context.Context + }{ + {"default", nil}, + {"caller", caller}, + } { + t.Run(tc.name, func(t *testing.T) { + var contexts []context.Context + pool := newTestPool(t, context.Background(), time.Second, func(ctx context.Context) (*glua.LState, error) { + contexts = append(contexts, ctx) + return glua.NewState(), nil + }) + state, err := pool.Acquire(nil) + if err != nil { + t.Fatal(err) + } + defer pool.Release(state, true) + if err := pool.WithState(tc.ctx, 2*time.Second, func(*glua.LState) error { return nil }); err != nil { + t.Fatal(err) + } + want := tc.ctx + if want == nil { + want = pool.ctx + } + if len(contexts) != 2 || contexts[0] != pool.ctx || contexts[1] != want { + t.Fatal("factory did not receive the initialization and acquisition contexts unchanged") + } + }) } - done := make(chan struct{}) +} + +func TestPoolWithStateLifecycle(t *testing.T) { + failure := errors.New("work failed") + for _, tc := range []struct { + name string + work func(*glua.LState, context.CancelFunc) error + reusable bool + wantPanic bool + wantErr error + }{ + {"success", func(*glua.LState, context.CancelFunc) error { return nil }, true, false, nil}, + {"canceled success", func(_ *glua.LState, cancel context.CancelFunc) error { + cancel() + return nil + }, true, false, nil}, + {"error", func(*glua.LState, context.CancelFunc) error { return failure }, false, false, failure}, + {"timeout", func(L *glua.LState, _ context.CancelFunc) error { return L.DoString("while true do end") }, false, false, nil}, + {"panic", func(*glua.LState, context.CancelFunc) error { panic(failure) }, false, true, nil}, + } { + t.Run(tc.name, func(t *testing.T) { + pool := newTestPool(t, context.Background(), 10*time.Millisecond, func(context.Context) (*glua.LState, error) { + state := glua.NewState() + state.Push(glua.LTrue) + return state, nil + }) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + var state *glua.LState + var workCtx context.Context + var recovered any + err := func() (err error) { + defer func() { recovered = recover() }() + return pool.WithState(ctx, 0, func(L *glua.LState) error { + state, workCtx = L, L.Context() + L.Push(glua.LFalse) + return tc.work(L, cancel) + }) + }() + if tc.wantPanic { + if recovered != failure { + t.Fatalf("panic = %v, want original panic", recovered) + } + } else { + if recovered != nil || (err == nil) != tc.reusable { + t.Fatalf("WithState error = %v, panic = %v", err, recovered) + } + if tc.wantErr != nil && !errors.Is(err, tc.wantErr) { + t.Fatalf("WithState error = %v, want %v", err, tc.wantErr) + } + } + if workCtx.Err() == nil { + t.Fatal("WithState did not cancel the execution context") + } + if closed := state.IsClosed(); closed == tc.reusable { + t.Fatalf("state closed = %t, want %t", closed, !tc.reusable) + } + if tc.reusable && (state.Context() != nil || state.GetTop() != 1 || state.Get(1) != glua.LTrue) { + t.Fatal("WithState did not reset the state for reuse") + } + if err := pool.WithState(nil, 0, func(L *glua.LState) error { + if (L == state) != tc.reusable { + t.Error("unexpected state reuse") + } + return nil + }); err != nil { + t.Fatal(err) + } + }) + } +} + +func TestPoolClose(t *testing.T) { + pool := newTestPool(t, context.Background(), time.Second, func(context.Context) (*glua.LState, error) { + return glua.NewState(), nil + }) + finishCtx, finish := context.WithCancel(context.Background()) + t.Cleanup(finish) + started, done := make(chan *glua.LState, 1), make(chan error, 1) + var workCtx context.Context + go func() { + done <- pool.WithState(nil, 0, func(L *glua.LState) error { + workCtx = L.Context() + started <- L + <-finishCtx.Done() + return nil + }) + }() + var state *glua.LState + select { + case state = <-started: + case <-time.After(time.Second): + t.Fatal("WithState did not start") + } + closed := make(chan struct{}) go func() { pool.Close() - close(done) + close(closed) }() select { - case <-pool.Context().Done(): + case <-workCtx.Done(): case <-time.After(time.Second): - pool.Release(state, false) - t.Fatal("Close did not cancel the pool context") + t.Fatal("Close did not cancel work using the pool context") + } + if !errors.Is(workCtx.Err(), context.Canceled) { + t.Fatalf("work context error = %v, want context.Canceled", workCtx.Err()) + } + assertPoolCloseBlocked(t, closed) + finish() + select { + case err := <-done: + if err != nil { + t.Fatalf("successful work returned an error: %v", err) + } + case <-time.After(time.Second): + t.Fatal("WithState did not finish") } select { - case <-done: - pool.Release(state, false) - t.Fatal("Close returned while a state was borrowed") - default: - } - pool.Release(state, true) - select { - case <-done: + case <-closed: case <-time.After(time.Second): - t.Fatal("Close did not finish after Release") + t.Fatal("Close did not finish after WithState") } if !state.IsClosed() { - t.Fatal("borrowed state was not closed") + t.Fatal("Release returned a state to a closed pool") } - if _, err := pool.Acquire(); !errors.Is(err, context.Canceled) { - t.Fatalf("Acquire after Close = %v, want context.Canceled", err) + if state, err := pool.Acquire(nil); state != nil || err == nil || errors.Is(err, context.Canceled) { + t.Fatalf("Acquire after Close = %v, %v; want closed pool error", state, err) } pool.Close() } -func TestPoolCloseCancelsStateCreation(t *testing.T) { - started := make(chan struct{}) +func TestPoolCloseWaitsForFactory(t *testing.T) { + finishCtx, finish := context.WithCancel(context.Background()) + started, canceled := make(chan struct{}), make(chan struct{}) first := true - pool, err := NewPool(context.Background(), func(ctx context.Context) (*glua.LState, error) { + pool := newTestPool(t, context.Background(), time.Second, func(ctx context.Context) (*glua.LState, error) { if first { first = false return glua.NewState(), nil } close(started) <-ctx.Done() + close(canceled) + <-finishCtx.Done() return nil, ctx.Err() }) + t.Cleanup(finish) + state, err := pool.Acquire(nil) if err != nil { t.Fatal(err) } - borrowed, err := pool.Acquire() - if err != nil { - t.Fatal(err) - } + pool.Release(state, false) acquireDone := make(chan error, 1) go func() { - _, err := pool.Acquire() + _, err := pool.Acquire(nil) acquireDone <- err }() - <-started - closeDone := make(chan struct{}) + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("state creation did not start") + } + closed := make(chan struct{}) go func() { pool.Close() - close(closeDone) + close(closed) }() select { + case <-canceled: + case <-time.After(time.Second): + t.Fatal("Close did not cancel state creation") + } + assertPoolCloseBlocked(t, closed) + finish() + select { case err := <-acquireDone: if !errors.Is(err, context.Canceled) { - t.Fatalf("Acquire during Close = %v, want context.Canceled", err) + t.Fatalf("Acquire error = %v, want context.Canceled", err) } case <-time.After(time.Second): - t.Fatal("state creation did not stop after Close") + t.Fatal("state creation did not finish") } select { - case <-closeDone: - pool.Release(borrowed, false) - t.Fatal("Close returned while the initial state was borrowed") - default: - } - pool.Release(borrowed, true) - select { - case <-closeDone: + case <-closed: case <-time.After(time.Second): - t.Fatal("Close did not finish after Release") + t.Fatal("Close did not finish after state creation") + } +} + +func TestPoolCloseWaitsForCallerContext(t *testing.T) { + pool := newTestPool(t, context.Background(), time.Minute, func(context.Context) (*glua.LState, error) { + return glua.NewState(), nil + }) + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + started, done := make(chan context.Context, 1), make(chan error, 1) + go func() { + done <- pool.WithState(ctx, 0, func(L *glua.LState) error { + started <- L.Context() + <-L.Context().Done() + return L.Context().Err() + }) + }() + var workCtx context.Context + select { + case workCtx = <-started: + case <-time.After(time.Second): + t.Fatal("WithState did not start") + } + closed := make(chan struct{}) + go func() { + pool.Close() + close(closed) + }() + select { + case <-pool.ctx.Done(): + case <-time.After(time.Second): + t.Fatal("Close did not cancel the pool context") + } + assertPoolCloseBlocked(t, closed) + if workCtx.Err() != nil || ctx.Err() != nil { + t.Fatal("Close canceled the caller's execution context") + } + cancel() + select { + case err := <-done: + if !errors.Is(err, context.Canceled) { + t.Fatalf("WithState error = %v, want context.Canceled", err) + } + case <-time.After(time.Second): + t.Fatal("WithState did not stop after caller cancellation") + } + select { + case <-closed: + case <-time.After(time.Second): + t.Fatal("Close did not finish after WithState") } } func BenchmarkPoolAcquireRelease(b *testing.B) { - pool, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) { + pool := newTestPool(b, context.Background(), time.Second, func(context.Context) (*glua.LState, error) { return glua.NewState(), nil }) - if err != nil { - b.Fatal(err) - } - defer pool.Close() - b.ReportAllocs() b.ResetTimer() for i := 0; i < b.N; i++ { - state, err := pool.Acquire() + state, err := pool.Acquire(nil) if err != nil { b.Fatal(err) } pool.Release(state, true) } - b.StopTimer() }