From ad4dfff19bd3163a9dc8058cfd5038bfd89639e6 Mon Sep 17 00:00:00 2001 From: yiming <1965768941@qq.com> Date: Sun, 6 Sep 2026 00:31:09 +0800 Subject: [PATCH] fix: guard wg balance and worker goroutine against hook panics --- hook_test.go | 163 +++++++++++++++++++++++++++++++++++++++++++++++++++ pool.go | 47 +++++++++------ worker.go | 31 +++++----- 3 files changed, 207 insertions(+), 34 deletions(-) diff --git a/hook_test.go b/hook_test.go index 8b10274..5bc0b3f 100644 --- a/hook_test.go +++ b/hook_test.go @@ -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 +} diff --git a/pool.go b/pool.go index 230f4cc..172a601 100644 --- a/pool.go +++ b/pool.go @@ -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: @@ -238,9 +238,7 @@ 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: } @@ -248,9 +246,7 @@ func (p *Pool) submit(ctx context.Context, task Task) bool { 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 @@ -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() @@ -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 { @@ -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() { @@ -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) +} diff --git a/worker.go b/worker.go index 46a5937..8fb48cd 100644 --- a/worker.go +++ b/worker.go @@ -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() }