From 1d56b3f15c761ed60cc7455800f73403a556160b Mon Sep 17 00:00:00 2001 From: Ray Liu Date: Sun, 20 Sep 2026 19:59:49 -0400 Subject: [PATCH] [4/9][worker] execute worker generations concurrently --- worker.go | 44 +++++++++++++++++++++++++++++---- worker_test.go | 67 +++++++++++++++++++++++++++++++++++++++++++++++--- 2 files changed, 102 insertions(+), 9 deletions(-) diff --git a/worker.go b/worker.go index f6cc11e..4b49479 100644 --- a/worker.go +++ b/worker.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "sync" "time" kqpb "github.com/raydatray/kq/internal/proto/kq" @@ -21,6 +22,15 @@ type Worker struct { handler Handler } +type workerJob struct { + record *kgo.Record +} + +type workerResult struct { + record *kgo.Record + err error +} + func NewWorker(config WorkerConfig, handler Handler) (*Worker, error) { if handler == nil { return nil, errors.New("kq: handler cannot be nil") @@ -52,6 +62,26 @@ func NewWorker(config WorkerConfig, handler Handler) (*Worker, error) { } func (w *Worker) Run(ctx context.Context) error { + jobs := make(chan workerJob, w.config.Concurrency) + results := make(chan workerResult, w.config.Concurrency) + var workers sync.WaitGroup + for range w.config.Concurrency { + workers.Add(1) + go func() { + defer workers.Done() + for job := range jobs { + results <- workerResult{ + record: job.record, + err: w.handle(ctx, job.record.Value), + } + } + }() + } + defer func() { + close(jobs) + workers.Wait() + }() + for { result := w.consumer.Poll(ctx, w.config.Concurrency) @@ -68,15 +98,19 @@ func (w *Worker) Run(ctx context.Context) error { } } - var taskErr error for _, record := range result.records { - err := w.handle(ctx, record.Value) + jobs <- workerJob{record: record} + } + + var taskErr error + for range result.records { + result := <-results status := kgo.AckAccept - if err != nil { + if result.err != nil { status = kgo.AckRelease - taskErr = errors.Join(taskErr, err) + taskErr = errors.Join(taskErr, result.err) } - w.consumer.Ack(record, status) + w.consumer.Ack(result.record, status) } ackErr := w.flushAcks() diff --git a/worker_test.go b/worker_test.go index e9df0a5..5357583 100644 --- a/worker_test.go +++ b/worker_test.go @@ -3,8 +3,10 @@ package kq import ( "context" "errors" + "fmt" "strings" "sync" + "sync/atomic" "testing" "time" @@ -62,13 +64,13 @@ func TestWorkerProcessesBoundedPollGeneration(t *testing.T) { {records: records}, {closed: true}, }} - var handled []string + handled := make(chan string, len(records)) worker := &Worker{ consumer: consumer, producer: new(fakeProducer), config: WorkerConfig{Concurrency: 3}, handler: func(_ context.Context, task Task) error { - handled = append(handled, task.Type) + handled <- task.Type return nil }, } @@ -76,8 +78,12 @@ func TestWorkerProcessesBoundedPollGeneration(t *testing.T) { if err := worker.Run(context.Background()); err != nil { t.Fatal(err) } - if got := strings.Join(handled, ","); got != "first,second" { - t.Fatalf("handled = %q, want first,second", got) + seen := make(map[string]bool) + for range records { + seen[<-handled] = true + } + if !seen["first"] || !seen["second"] { + t.Fatalf("handled = %v, want first and second", seen) } if len(consumer.pollLimits) != 2 || consumer.pollLimits[0] != 3 || consumer.pollLimits[1] != 3 { t.Fatalf("poll limits = %v, want [3 3]", consumer.pollLimits) @@ -93,6 +99,59 @@ func TestWorkerProcessesBoundedPollGeneration(t *testing.T) { } } +func TestWorkerExecutesGenerationConcurrently(t *testing.T) { + const concurrency = 4 + records := make([]*kgo.Record, 0, concurrency) + for i := range concurrency { + records = append(records, &kgo.Record{ + Value: encodeWorkerTestTask(t, Task{Type: fmt.Sprintf("task-%d", i)}, 0), + }) + } + consumer := &fakeShareGroupConsumer{results: []sharePollResult{ + {records: records}, + {closed: true}, + }} + started := make(chan struct{}, concurrency) + release := make(chan struct{}) + var active atomic.Int32 + var maximum atomic.Int32 + worker := &Worker{ + consumer: consumer, + producer: new(fakeProducer), + config: WorkerConfig{Concurrency: concurrency}, + handler: func(context.Context, Task) error { + current := active.Add(1) + for { + prior := maximum.Load() + if current <= prior || maximum.CompareAndSwap(prior, current) { + break + } + } + started <- struct{}{} + <-release + active.Add(-1) + return nil + }, + } + done := make(chan error, 1) + go func() { done <- worker.Run(context.Background()) }() + + for range concurrency { + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("generation did not execute concurrently") + } + } + if got := maximum.Load(); got != concurrency { + t.Fatalf("maximum concurrency = %d, want %d", got, concurrency) + } + close(release) + if err := <-done; err != nil { + t.Fatal(err) + } +} + func TestWorkerHandlesTask(t *testing.T) { producer := new(fakeProducer) value := encodeWorkerTestTask(t, Task{