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
50 changes: 37 additions & 13 deletions internal/limits/limiters/concurrency.go
Original file line number Diff line number Diff line change
Expand Up @@ -18,51 +18,75 @@ along with this program. If not, see <https://www.gnu.org/licenses/>.

package limiters

import "context"
import (
"context"

"golang.org/x/sync/semaphore"
)

// Semaphore is a convenience wrapper for a channel that implements
// semaphore-kind synchronization.
//
// 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()
}
7 changes: 6 additions & 1 deletion internal/limits/limiters/limiters.go
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,10 @@ along with this program. If not, see <https://www.gnu.org/licenses/>.
// 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
Expand All @@ -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")
84 changes: 28 additions & 56 deletions internal/limits/limiters/rate.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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()
}
Loading