diff --git a/internal/limits/limiters/concurrency.go b/internal/limits/limiters/concurrency.go
index 18923e89..2955b128 100644
--- a/internal/limits/limiters/concurrency.go
+++ b/internal/limits/limiters/concurrency.go
@@ -18,7 +18,11 @@ along with this program. If not, see .
package limiters
-import "context"
+import (
+ "context"
+
+ "golang.org/x/sync/semaphore"
+)
// Semaphore is a convenience wrapper for a channel that implements
// semaphore-kind synchronization.
@@ -26,43 +30,63 @@ import "context"
// If the argument given to the NewSemaphore is negative or zero,
// all methods are no-op.
type Semaphore struct {
- c chan struct{}
+ weighted *semaphore.Weighted
+ ctx context.Context
+ cancel context.CancelFunc
}
func NewSemaphore(max int) Semaphore {
- return Semaphore{c: make(chan struct{}, max)}
+ ctx, cancel := context.WithCancel(context.TODO())
+ s := Semaphore{weighted: nil, ctx: ctx, cancel: cancel}
+ if max > 0 {
+ s.weighted = semaphore.NewWeighted(int64(max))
+ }
+ return s
}
func (s Semaphore) Take() bool {
- if cap(s.c) <= 0 {
+ if s.weighted == nil {
return true
}
- s.c <- struct{}{}
+
+ if err := s.weighted.Acquire(s.ctx, 1); err != nil {
+ return false
+ }
return true
}
func (s Semaphore) TakeContext(ctx context.Context) error {
- if cap(s.c) <= 0 {
+ if s.weighted == nil {
return nil
}
select {
- case s.c <- struct{}{}:
- return nil
- case <-ctx.Done():
- return ctx.Err()
+ case <-s.ctx.Done():
+ return ErrClosed
+ default:
}
+ reqCtx, reqCancel := context.WithCancel(ctx)
+ defer reqCancel()
+
+ stop := context.AfterFunc(s.ctx, func() {
+ reqCancel()
+ })
+ defer stop()
+
+ return s.weighted.Acquire(reqCtx, 1)
}
func (s Semaphore) Release() {
- if cap(s.c) <= 0 {
+ if s.weighted == nil {
return
}
select {
- case <-s.c:
+ case <-s.ctx.Done():
+ return
default:
- panic("limiters: mismatched Release call")
+ s.weighted.Release(1)
}
}
func (s Semaphore) Close() {
+ s.cancel()
}
diff --git a/internal/limits/limiters/limiters.go b/internal/limits/limiters/limiters.go
index 9b11761d..691580b7 100644
--- a/internal/limits/limiters/limiters.go
+++ b/internal/limits/limiters/limiters.go
@@ -20,7 +20,10 @@ along with this program. If not, see .
// of resources consumed by the server.
package limiters
-import "context"
+import (
+ "context"
+ "errors"
+)
// The L interface represents a blocking limiter that has some upper bound of
// resource use and blocks when it is exceeded until enough resources are
@@ -33,3 +36,5 @@ type L interface {
// Close frees any resources used internally by Limiter for book-keeping.
Close()
}
+
+var ErrClosed = errors.New("limiters: Bucket is closed")
diff --git a/internal/limits/limiters/rate.go b/internal/limits/limiters/rate.go
index a774187f..ee0c44ac 100644
--- a/internal/limits/limiters/rate.go
+++ b/internal/limits/limiters/rate.go
@@ -20,11 +20,10 @@ package limiters
import (
"context"
- "errors"
"time"
-)
-var ErrClosed = errors.New("limiters: Rate bucket is closed")
+ "golang.org/x/time/rate"
+)
// Rate structure implements a basic rate-limiter for requests using the token
// bucket approach.
@@ -37,81 +36,54 @@ var ErrClosed = errors.New("limiters: Rate bucket is closed")
//
// If burstSize = 0, all methods are no-op and always succeed.
type Rate struct {
- bucket chan struct{}
- stop chan struct{}
+ limiter *rate.Limiter
+ ctx context.Context
+ cancel context.CancelFunc
}
func NewRate(burstSize int, interval time.Duration) Rate {
- r := Rate{
- bucket: make(chan struct{}, burstSize),
- stop: make(chan struct{}),
- }
-
- if burstSize == 0 {
- return r
+ ctx, cancel := context.WithCancel(context.TODO())
+ r := Rate{limiter: nil, ctx: ctx, cancel: cancel}
+ if burstSize > 0 {
+ r.limiter = rate.NewLimiter(rate.Every(interval), burstSize)
}
-
- for i := 0; i < burstSize; i++ {
- r.bucket <- struct{}{}
- }
-
- go r.fill(burstSize, interval)
return r
}
-func (r Rate) fill(burstSize int, interval time.Duration) {
- t := time.NewTimer(interval)
- defer t.Stop()
- for {
- t.Reset(interval)
- select {
- case <-t.C:
- case <-r.stop:
- close(r.bucket)
- return
- }
-
- fill:
- for i := 0; i < burstSize; i++ {
- select {
- case r.bucket <- struct{}{}:
- default:
- // If there are no Take pending and the bucket is already
- // full - don't block.
- break fill
- }
- }
- }
-}
-
func (r Rate) Take() bool {
- if cap(r.bucket) == 0 {
+ if r.limiter == nil {
return true
}
- _, ok := <-r.bucket
- return ok
+ if err := r.limiter.Wait(r.ctx); err != nil {
+ return false
+ }
+ return true
}
func (r Rate) TakeContext(ctx context.Context) error {
- if cap(r.bucket) == 0 {
+ if r.limiter == nil {
return nil
}
-
select {
- case _, ok := <-r.bucket:
- if !ok {
- return ErrClosed
- }
- return nil
- case <-ctx.Done():
- return ctx.Err()
+ case <-r.ctx.Done():
+ return ErrClosed
+ default:
}
+ reqCtx, reqCancel := context.WithCancel(ctx)
+ defer reqCancel()
+
+ stop := context.AfterFunc(r.ctx, func() {
+ reqCancel()
+ })
+ defer stop()
+
+ return r.limiter.Wait(reqCtx)
}
func (r Rate) Release() {
}
func (r Rate) Close() {
- close(r.stop)
+ r.cancel()
}