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
44 changes: 39 additions & 5 deletions worker.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import (
"context"
"errors"
"fmt"
"sync"
"time"

kqpb "github.com/raydatray/kq/internal/proto/kq"
Expand All @@ -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")
Expand Down Expand Up @@ -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)

Expand All @@ -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()
Expand Down
67 changes: 63 additions & 4 deletions worker_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,8 +3,10 @@ package kq
import (
"context"
"errors"
"fmt"
"strings"
"sync"
"sync/atomic"
"testing"
"time"

Expand Down Expand Up @@ -62,22 +64,26 @@ 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
},
}

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)
Expand All @@ -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{
Expand Down
Loading