From d4c965a7c8cb27bb89f4e6e06f2f270a79613a94 Mon Sep 17 00:00:00 2001 From: bensynapse <118375461+bensynapse@users.noreply.github.com> Date: Sat, 3 Oct 2026 10:57:50 +0300 Subject: [PATCH] fix(ring): serialize capacity and shutdown checks --- ring.go | 11 ++++--- ring_test.go | 88 ++++++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 94 insertions(+), 5 deletions(-) diff --git a/ring.go b/ring.go index 233293c..3c51a8e 100644 --- a/ring.go +++ b/ring.go @@ -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 @@ -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) @@ -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 } @@ -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 { @@ -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 diff --git a/ring_test.go b/ring_test.go index d2024eb..79ddebb 100644 --- a/ring_test.go +++ b/ring_test.go @@ -6,6 +6,7 @@ import ( "fmt" "log" "runtime" + "sync" "testing" "time" @@ -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,