diff --git a/packages/shared/pkg/utils/resizable_semaphore.go b/packages/shared/pkg/utils/resizable_semaphore.go index db1542335d..38626b839a 100644 --- a/packages/shared/pkg/utils/resizable_semaphore.go +++ b/packages/shared/pkg/utils/resizable_semaphore.go @@ -40,7 +40,11 @@ func (s *AdjustableSemaphore) Acquire(ctx context.Context, n int64) error { } // Wake ->cond.Wait when ctx is canceled. - stop := context.AfterFunc(ctx, s.cond.Broadcast) + stop := context.AfterFunc(ctx, func() { + s.mu.Lock() + defer s.mu.Unlock() + s.cond.Broadcast() + }) defer stop() // ensure we don’t leak the callback for s.used+n > s.limit { diff --git a/packages/shared/pkg/utils/resizable_semaphore_test.go b/packages/shared/pkg/utils/resizable_semaphore_test.go index 22230be37e..9ddc9837f5 100644 --- a/packages/shared/pkg/utils/resizable_semaphore_test.go +++ b/packages/shared/pkg/utils/resizable_semaphore_test.go @@ -6,6 +6,7 @@ import ( "runtime" "sync" "testing" + "testing/synctest" "time" "github.com/stretchr/testify/assert" @@ -274,6 +275,52 @@ func TestAcquireRespectsContextCancel(t *testing.T) { } } +type semaphoreCancelContext struct { + context.Context + cancel context.CancelFunc +} + +func (c *semaphoreCancelContext) Err() error { + err := c.Context.Err() + if err == nil { + c.cancel() + runtime.Gosched() + } + + return err +} + +func TestAcquireCancellationBeforeWait(t *testing.T) { + t.Parallel() + + synctest.Test(t, func(t *testing.T) { + for range 100 { + s, err := NewAdjustableSemaphore(1) + require.NoError(t, err) + require.True(t, s.TryAcquire(1)) + + ctx, cancel := context.WithCancel(t.Context()) + cancelOnCheck := &semaphoreCancelContext{Context: ctx, cancel: cancel} + result := make(chan error, 1) + go func() { result <- s.Acquire(cancelOnCheck, 1) }() + synctest.Wait() + cancel() + + select { + case err := <-result: + require.ErrorIs(t, err, context.Canceled) + default: + s.mu.Lock() + s.cond.Broadcast() + s.mu.Unlock() + synctest.Wait() + <-result + t.Fatal("Acquire missed cancellation before entering Wait") + } + } + }) +} + // ----------------------------------------------------------------------------- // race-detector stress // -----------------------------------------------------------------------------