Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
163 changes: 163 additions & 0 deletions hook_test.go
Original file line number Diff line number Diff line change
@@ -1 +1,164 @@
package agilepool

import (
"context"
"sync/atomic"
"testing"
"time"
)

// panicHooks is a hooks implementation that deliberately panics at selected
// lifecycle points. It mimics a custom hooks implementation that does not
// recover its own callbacks (the bundled internal/hook.Hooks does).
type panicHooks struct {
panicSubmitted bool
panicEnqueued bool
panicStarted bool
panicCompleted bool
panicPoolClosed bool
}

func (h *panicHooks) DispatchTaskSubmitted(context.Context, Task) {
if h.panicSubmitted {
panic("submitted hook panic")
}
}

func (h *panicHooks) DispatchTaskEnqueued(context.Context, Task) {
if h.panicEnqueued {
panic("enqueued hook panic")
}
}

func (h *panicHooks) DispatchTaskStarted(context.Context, Task) {
if h.panicStarted {
panic("started hook panic")
}
}

func (h *panicHooks) DispatchTaskCompleted(context.Context, Task, any) {
if h.panicCompleted {
panic("completed hook panic")
}
}

func (h *panicHooks) DispatchPoolClosed(*Pool) {
if h.panicPoolClosed {
panic("pool closed hook panic")
}
}

// waitPoolDone reports whether p.Wait() returns within timeout. A pool whose
// wg was leaked by a panicking hook (or a dead worker goroutine) blocks here.
func waitPoolDone(p *Pool, timeout time.Duration) bool {
done := make(chan struct{})
go func() {
p.Wait()
close(done)
}()
select {
case <-done:
return true
case <-time.After(timeout):
return false
}
}

// TestHookCompletedPanicDoesNotSkipDone is the regression test for the
// runTask defer ordering bug: p.done() used to live in the same deferred
// closure after DispatchTaskCompleted, so a panicking Completed hook skipped
// wg.Done and left Wait() blocked forever.
func TestHookCompletedPanicDoesNotSkipDone(t *testing.T) {
p := NewPool(NewConfig())
defer p.Close()
if err := p.SetHook(&panicHooks{panicCompleted: true}); err != nil {
t.Fatal(err)
}

const n = 10
var executed atomic.Int32
for i := 0; i < n; i++ {
p.Submit(TaskFunc(func() error {
executed.Add(1)
return nil
}))
}
if !waitPoolDone(p, 5*time.Second) {
t.Fatal("Wait() blocked: a panicking Completed hook skipped wg.Done")
}
if got := executed.Load(); got != n {
t.Fatalf("executed = %d, want %d", got, n)
}

// The worker goroutine must survive the hook panic and keep serving.
p.Submit(TaskFunc(func() error {
executed.Add(1)
return nil
}))
if !waitPoolDone(p, 5*time.Second) {
t.Fatal("Wait() blocked after a Completed hook panic")
}
if got := executed.Load(); got != n+1 {
t.Fatalf("executed = %d, want %d", got, n+1)
}
}

// TestHookStartedPanicDoesNotPreventExecution guards the dispatch order in
// runTask: the Started hook used to be dispatched before any defer was
// registered, so its panic crashed the worker goroutine before the task ran.
func TestHookStartedPanicDoesNotPreventExecution(t *testing.T) {
p := NewPool(NewConfig())
defer p.Close()
if err := p.SetHook(&panicHooks{panicStarted: true}); err != nil {
t.Fatal(err)
}

var executed atomic.Int32
p.Submit(TaskFunc(func() error {
executed.Add(1)
return nil
}))
if !waitPoolDone(p, 5*time.Second) {
t.Fatal("Wait() blocked")
}
if got := executed.Load(); got != 1 {
t.Fatalf("executed = %d, want 1", got)
}
}

// TestHookSubmittedPanicDoesNotAbortSubmit guards the submission path: the
// Submitted hook fires after wg.Add(1), so a panic there used to leak the
// WaitGroup before the task was enqueued.
func TestHookSubmittedPanicDoesNotAbortSubmit(t *testing.T) {
p := NewPool(NewConfig())
defer p.Close()
if err := p.SetHook(&panicHooks{panicSubmitted: true}); err != nil {
t.Fatal(err)
}

var executed atomic.Int32
p.Submit(TaskFunc(func() error {
executed.Add(1)
return nil
}))
if !waitPoolDone(p, 5*time.Second) {
t.Fatal("Wait() blocked: a panicking Submitted hook leaked the wg")
}
if got := executed.Load(); got != 1 {
t.Fatalf("executed = %d, want 1", got)
}
}

// TestHookPoolClosedPanicDoesNotAbortClose guards Pool.Close: a panicking
// PoolClosed hook must not propagate to the Close caller.
func TestHookPoolClosedPanicDoesNotAbortClose(t *testing.T) {
p := NewPool(NewConfig())
if err := p.SetHook(&panicHooks{panicPoolClosed: true}); err != nil {
t.Fatal(err)
}
p.Submit(TaskFunc(func() error { return nil }))
if !waitPoolDone(p, 5*time.Second) {
t.Fatal("Wait() blocked")
}
p.Close() // must return normally despite the panicking hook
}
47 changes: 29 additions & 18 deletions pool.go
Original file line number Diff line number Diff line change
Expand Up @@ -222,9 +222,9 @@ func (p *Pool) submit(ctx context.Context, task Task) bool {
hookCtx = wrapped.ctx
hookTask = wrapped.task
}
if p.hooks != nil {
p.hooks.DispatchTaskSubmitted(hookCtx, hookTask)
}
p.dispatchHook(func(h hooks) {
h.DispatchTaskSubmitted(hookCtx, hookTask)
})
if p.config.workMode == NONBLOCK {
select {
case p.taskQueue <- task:
Expand All @@ -238,19 +238,15 @@ func (p *Pool) submit(ctx context.Context, task Task) bool {

select {
case p.taskQueue <- task:
if p.hooks != nil {
p.dispatchTaskEnqueuedFor(task)
}
p.dispatchTaskEnqueuedFor(task)
return true
default:
}

result := p.taskBuf.PushAndForward(task, func(t Task) bool {
select {
case p.taskQueue <- t:
if p.hooks != nil {
p.dispatchTaskEnqueuedFor(t)
}
p.dispatchTaskEnqueuedFor(t)
return true
default:
return false
Expand All @@ -268,9 +264,7 @@ func (p *Pool) submit(ctx context.Context, task Task) bool {
// behind after Close.
select {
case p.taskQueue <- task:
if p.hooks != nil {
p.dispatchTaskEnqueuedFor(task)
}
p.dispatchTaskEnqueuedFor(task)
return true
case <-ctx.Done():
p.done()
Expand All @@ -290,9 +284,9 @@ func (p *Pool) dispatchTaskEnqueuedFor(task Task) {
ctx = wrapped.ctx
task = wrapped.task
}
if p.hooks != nil {
p.hooks.DispatchTaskEnqueued(ctx, task)
}
p.dispatchHook(func(h hooks) {
h.DispatchTaskEnqueued(ctx, task)
})
}

type contextTask struct {
Expand Down Expand Up @@ -542,9 +536,9 @@ func (p *Pool) Close() {
p.taskBuf.Close()

close(p.closePoolCn)
if p.hooks != nil {
p.hooks.DispatchPoolClosed(p) // fire OnPoolClosed hooks
}
p.dispatchHook(func(h hooks) {
h.DispatchPoolClosed(p) // fire OnPoolClosed hooks
})
}

func (p *Pool) Wait() {
Expand Down Expand Up @@ -601,3 +595,20 @@ func (p *Pool) SetHook(hooks hooks) error {
p.hooks = hooks
return nil
}

// dispatchHook invokes fn with the pool's hook set and absorbs any panic a
// hooks implementation raises. internal/hook.Hooks already recovers each
// registered callback; this guard additionally keeps a custom implementation
// that panics mid-dispatch from crashing the submitting goroutine, a worker,
// or Close, and from skipping pool bookkeeping such as wg.Done.
func (p *Pool) dispatchHook(fn func(h hooks)) {
if p.hooks == nil {
return
}
defer func() {
if r := recover(); r != nil {
p.logger.Printf("hook dispatch panicked: %v\n%s\n", r, Stack(1))
}
}()
fn(p.hooks)
}
31 changes: 15 additions & 16 deletions worker.go
Original file line number Diff line number Diff line change
Expand Up @@ -113,35 +113,34 @@ loop:
func (w *worker) runTask(task Task) {
atomic.AddInt64(&w.pool.consumeCount, 1)

// Balance the submit-side wg.Add exactly once per task. Registered before
// any hook dispatch and kept on its own defer, so a hook implementation
// that panics cannot skip done() and deadlock Wait()/Close().
defer w.pool.done()

hookCtx := context.Background()
hookTask := task
if wrapped, ok := task.(*contextTask); ok {
hookCtx = wrapped.ctx
hookTask = wrapped.task
}
if w.pool.hooks != nil {
w.pool.hooks.DispatchTaskStarted(hookCtx, hookTask)
}

// Capture task panics so the Completed hook can observe the recovered
// value. dispatchHook guards the hook call itself; w.pool.done() above
// runs afterwards regardless of what the hooks do.
var recovered any
defer func() {
hookCtx := context.Background()
hookTask := task
if wrapped, ok := task.(*contextTask); ok {
hookCtx = wrapped.ctx
hookTask = wrapped.task
}
if w.pool.hooks != nil {
w.pool.hooks.DispatchTaskCompleted(hookCtx, hookTask, recovered)
}
w.pool.done()
}()

defer func() {
if p := recover(); p != nil {
recovered = p
w.pool.logger.Printf("worker exits from panic: %v\n%s\n", p, Stack(1))
}
w.pool.dispatchHook(func(h hooks) {
h.DispatchTaskCompleted(hookCtx, hookTask, recovered)
})
}()

w.pool.dispatchHook(func(h hooks) {
h.DispatchTaskStarted(hookCtx, hookTask)
})
task.process()
}
Loading