diff --git a/go.mod b/go.mod index 7de6d34..dedc13e 100644 --- a/go.mod +++ b/go.mod
@@ -1,3 +1,3 @@ module golang.org/x/sync -go 1.25.0 +go 1.26.0
diff --git a/semaphore/semaphore.go b/semaphore/semaphore.go index 040c5bc..f2a7d3f 100644 --- a/semaphore/semaphore.go +++ b/semaphore/semaphore.go
@@ -17,14 +17,17 @@ } // NewWeighted creates a new weighted semaphore with the given -// maximum combined weight for concurrent access. +// maximum combined weight for concurrent access. NewWeighted panics if n is +// negative. func NewWeighted(n int64) *Weighted { - w := &Weighted{size: n} - return w + if n < 0 { + panic("semaphore: size < 0") + } + return &Weighted{size: n} } // Weighted provides a way to bound concurrent access to a resource. -// The callers can request access with a given weight. +// The callers can request access with a given non-negative weight. type Weighted struct { size int64 cur int64 @@ -32,10 +35,13 @@ waiters list.List } -// Acquire acquires the semaphore with a weight of n, blocking until resources +// Acquire acquires the semaphore with a non-negative weight of n, blocking until resources // are available or ctx is done. On success, returns nil. On failure, returns // ctx.Err() and leaves the semaphore unchanged. func (s *Weighted) Acquire(ctx context.Context, n int64) error { + if n < 0 { + panic("semaphore: n < 0") + } done := ctx.Done() s.mu.Lock() @@ -106,9 +112,12 @@ } } -// TryAcquire acquires the semaphore with a weight of n without blocking. +// TryAcquire acquires the semaphore with a non-negative weight of n without blocking. // On success, returns true. On failure, returns false and leaves the semaphore unchanged. func (s *Weighted) TryAcquire(n int64) bool { + if n < 0 { + panic("semaphore: n < 0") + } s.mu.Lock() success := s.size-s.cur >= n && s.waiters.Len() == 0 if success { @@ -118,8 +127,11 @@ return success } -// Release releases the semaphore with a weight of n. +// Release releases the semaphore with a non-negative weight of n. func (s *Weighted) Release(n int64) { + if n < 0 { + panic("semaphore: n < 0") + } s.mu.Lock() s.cur -= n if s.cur < 0 {
diff --git a/semaphore/semaphore_test.go b/semaphore/semaphore_test.go index 0fa08a9..058a382 100644 --- a/semaphore/semaphore_test.go +++ b/semaphore/semaphore_test.go
@@ -56,6 +56,60 @@ w.Release(1) } +func TestWeightedNegativeSizePanic(t *testing.T) { + t.Parallel() + + defer func() { + if recover() == nil { + t.Fatal("NewWeighted with negative size did not panic") + } + }() + semaphore.NewWeighted(-1) +} + +func TestWeightedZeroSize(t *testing.T) { + t.Parallel() + + w := semaphore.NewWeighted(0) + if !w.TryAcquire(0) { + t.Fatal("TryAcquire on a zero-sized semaphore with zero weight failed") + } +} + +func TestWeightedNegativeWeightPanic(t *testing.T) { + t.Parallel() + + ctx := context.Background() + w := semaphore.NewWeighted(1) + + func() { + defer func() { + if recover() == nil { + t.Fatal("Acquire with negative weight did not panic") + } + }() + w.Acquire(ctx, -1) + }() + + func() { + defer func() { + if recover() == nil { + t.Fatal("TryAcquire with negative weight did not panic") + } + }() + w.TryAcquire(-1) + }() + + func() { + defer func() { + if recover() == nil { + t.Fatal("Release with negative weight did not panic") + } + }() + w.Release(-1) + }() +} + func TestWeightedTryAcquire(t *testing.T) { t.Parallel()