Skip to content
Open
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
26 changes: 22 additions & 4 deletions internal/pool/pool.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,12 +16,26 @@
// holding it, the rest keep draining the queue instead of idling behind it.
package pool

import "sync"
import (
"context"
"sync"
)

// Each runs fn for every item in items, concurrently, across at most workers
// goroutines. order isn't preserved - fn must be safe to call from multiple
// goroutines and guard any shared state itself. blocks until every item is done.
func Each[T any](items []T, workers int, fn func(T)) {
EachCtx(context.Background(), items, workers, func(_ context.Context, item T) {
fn(item)
})
}

// EachCtx is Each with a context: once ctx is cancelled, workers stop pulling
// new items off the queue (fn still gets called with ctx for the item it's
// already holding, so it can bail out mid-item too). this is what lets a
// fan-out scanner - a full port sweep, a directory brute-force - stop
// promptly on ctrl-c / -max-time instead of draining the whole queue first.
func EachCtx[T any](ctx context.Context, items []T, workers int, fn func(context.Context, T)) {
if len(items) == 0 {
return
}
Expand All @@ -46,10 +60,14 @@ func Each[T any](items []T, workers int, fn func(T)) {
for i := 0; i < workers; i++ {
go func() {
defer wg.Done()
// pull until the queue is drained; a worker that finishes its
// current item just grabs the next, which is the work-stealing.
// pull until the queue is drained or ctx is cancelled; a worker
// that finishes its current item just grabs the next, which is
// the work-stealing.
for item := range queue {
fn(item)
if ctx.Err() != nil {
return
}
fn(ctx, item)
}
}()
}
Expand Down
49 changes: 49 additions & 0 deletions internal/pool/pool_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -13,9 +13,11 @@
package pool

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

// every item runs exactly once across a spread of sizes and worker counts,
Expand Down Expand Up @@ -143,3 +145,50 @@ func TestEachCapsAtItemCount(t *testing.T) {
t.Fatalf("peak concurrency %d exceeded item count %d", got, items)
}
}

// TestEachCtxStopsEarlyOnCancel proves the cancellation contract documented on
// EachCtx (pool.go), for the -max-time / ctrl-c fan-out case.
func TestEachCtxStopsEarlyOnCancel(t *testing.T) {
const (
items = 1000
workers = 4
)
work := make([]int, items)

ctx, cancel := context.WithCancel(context.Background())

var processed int64
var wg sync.WaitGroup
wg.Add(1)
go func() {
defer wg.Done()
EachCtx(ctx, work, workers, func(_ context.Context, _ int) {
// cancel partway through so most items are still queued when it fires.
if atomic.AddInt64(&processed, 1) == 20 {
cancel()
}
// give the other workers a chance to observe the cancellation
// before grabbing their next item, instead of racing straight
// through the buffered channel.
time.Sleep(time.Millisecond)
})
}()

done := make(chan struct{})
go func() {
wg.Wait()
close(done)
}()

select {
case <-done:
case <-time.After(5 * time.Second):
t.Fatal("EachCtx did not return after cancellation; workers kept draining the queue")
}

got := atomic.LoadInt64(&processed)
if got >= items {
t.Fatalf("EachCtx processed all %d items despite cancellation (processed=%d): cancel had no effect", items, got)
}
t.Logf("EachCtx stopped after %d/%d items once ctx was cancelled", got, items)
}
Loading