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
11 changes: 6 additions & 5 deletions ring.go
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,9 @@ func (s *Ring) Shutdown() error {
//
// Thread-safety: This method is safe for concurrent calls.
func (s *Ring) Queue(task core.TaskMessage) error {
s.Lock()
defer s.Unlock()

// Reject new tasks if shutdown has been initiated
if s.stopFlag.Load() == 1 {
return ErrQueueShutdown
Expand All @@ -79,7 +82,6 @@ func (s *Ring) Queue(task core.TaskMessage) error {
return ErrMaxCapacity
}

s.Lock()
// Grow the buffer if it's full (before adding the new task)
if s.count == len(s.taskQueue) {
s.resize(s.count * 2)
Expand All @@ -88,7 +90,6 @@ func (s *Ring) Queue(task core.TaskMessage) error {
s.taskQueue[s.tail] = task
s.tail = (s.tail + 1) % len(s.taskQueue)
s.count++
s.Unlock()

return nil
}
Expand All @@ -105,6 +106,9 @@ func (s *Ring) Queue(task core.TaskMessage) error {
//
// Thread-safety: This method is safe for concurrent calls.
func (s *Ring) Request() (core.TaskMessage, error) {
s.Lock()
defer s.Unlock()

// If shutting down and queue is empty, signal exit and return closed error
if s.stopFlag.Load() == 1 && s.count == 0 {
select {
Expand All @@ -114,9 +118,6 @@ func (s *Ring) Request() (core.TaskMessage, error) {
return nil, ErrQueueHasBeenClosed
}

s.Lock()
defer s.Unlock()

// Return early if queue is empty (but not shutting down yet)
if s.count == 0 {
return nil, ErrNoTaskInQueue
Expand Down
88 changes: 88 additions & 0 deletions ring_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import (
"fmt"
"log"
"runtime"
"sync"
"testing"
"time"

Expand Down Expand Up @@ -34,6 +35,93 @@ func TestMaxCapacity(t *testing.T) {
assert.Equal(t, ErrMaxCapacity, err)
}

func TestConcurrentMaxCapacity(t *testing.T) {
for _, capacity := range []int{1, 2, 8} {
t.Run(fmt.Sprintf("capacity_%d", capacity), func(t *testing.T) {
const producers = 64
w := NewRing(WithQueueSize(capacity))
start := make(chan struct{})
results := make(chan error, producers)
var workers sync.WaitGroup

for range producers {
workers.Go(func() {
<-start
results <- w.Queue(&mockMessage{})
})
}
close(start)
workers.Wait()
close(results)

accepted := 0
for err := range results {
if err == nil {
accepted++
} else {
require.ErrorIs(t, err, ErrMaxCapacity)
}
}
require.Equal(t, capacity, accepted)
require.Equal(t, capacity, w.count)

for range capacity {
_, err := w.Request()
require.NoError(t, err)
}
_, err := w.Request()
require.ErrorIs(t, err, ErrNoTaskInQueue)
require.NoError(t, w.Queue(&mockMessage{}))
})
}
}

func TestConcurrentRequestsDuringShutdown(t *testing.T) {
const (
tasks = 256
consumers = 8
)
w := NewRing()
for range tasks {
require.NoError(t, w.Queue(&mockMessage{}))
}
// Set the shutdown state directly to isolate concurrent draining.
w.stopFlag.Store(1)
start := make(chan struct{})
results := make(chan error, tasks+consumers)
var workers sync.WaitGroup

for range consumers {
workers.Go(func() {
<-start
for {
_, err := w.Request()
results <- err
if err != nil {
return
}
}
})
}
close(start)
workers.Wait()
close(results)

drained := 0
for err := range results {
if err == nil {
drained++
} else {
require.ErrorIs(t, err, ErrQueueHasBeenClosed)
}
}
require.Equal(t, tasks, drained)
task, err := w.Request()
require.Nil(t, task)
require.ErrorIs(t, err, ErrQueueHasBeenClosed)
require.ErrorIs(t, w.Queue(&mockMessage{}), ErrQueueShutdown)
}

func TestCustomFuncAndWait(t *testing.T) {
m := mockMessage{
message: testMessageFoo,
Expand Down