Skip to content
Merged
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
5 changes: 4 additions & 1 deletion auto_ratelimit_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -187,7 +187,10 @@ func TestAutoLimiterCreateOrDefault(t *testing.T) {
require.True(t, recreated.CanTake())
recreated.Take()
}
require.False(t, recreated.CanTake())
// Take returns at the unbuffered handoff; the producer decrements its
// accounting immediately afterwards. Wait for that scheduler-sized window
// instead of making the test depend on goroutine ordering.
require.Eventually(t, func() bool { return !recreated.CanTake() }, time.Second, time.Millisecond)
}

func TestAutoLimiterAddAndTake(t *testing.T) {
Expand Down
22 changes: 18 additions & 4 deletions ratelimit.go
Original file line number Diff line number Diff line change
Expand Up @@ -92,7 +92,7 @@ func (limiter *Limiter) Take() {

switch limiter.strategy {
case LeakyBucket:
_ = limiter.leakyBucketLimiter.Wait(context.TODO())
_ = limiter.leakyBucketLimiter.Wait(limiter.ctx)
default:
<-limiter.tokens
}
Expand Down Expand Up @@ -142,6 +142,7 @@ func (limiter *Limiter) SetLimit(max uint) {

switch limiter.strategy {
case LeakyBucket:
limiter.leakyBucketLimiter.SetLimit(leakyBucketRate(max, limiter.interval))
limiter.leakyBucketLimiter.SetBurst(int(max))
default:
}
Expand Down Expand Up @@ -169,7 +170,7 @@ func (limiter *Limiter) SetDuration(d time.Duration) {
limiter.interval = d
switch limiter.strategy {
case LeakyBucket:
limiter.leakyBucketLimiter.SetLimit(rate.Every(d))
limiter.leakyBucketLimiter.SetLimit(leakyBucketRate(limiter.GetLimit(), d))
default:
limiter.ticker.Reset(d)
}
Expand All @@ -190,7 +191,10 @@ func (limiter *Limiter) Stop() {
}

switch limiter.strategy {
case LeakyBucket: // NOP
case LeakyBucket:
if limiter.cancelFunc != nil {
limiter.cancelFunc()
}
default:
if limiter.cancelFunc != nil {
limiter.cancelFunc()
Expand Down Expand Up @@ -239,13 +243,23 @@ func NewUnlimited(ctx context.Context) *Limiter {

// NewLeakyBucket creates a limiter that uses golang.org/x/time/rate.
func NewLeakyBucket(ctx context.Context, max uint, duration time.Duration) *Limiter {
internalctx, cancel := context.WithCancel(ctx)
limiter := &Limiter{
strategy: LeakyBucket,
leakyBucketLimiter: rate.NewLimiter(rate.Every(duration), int(max)),
leakyBucketLimiter: rate.NewLimiter(leakyBucketRate(max, duration), int(max)),
ctx: internalctx,
cancelFunc: cancel,
}

limiter.maxCount.Store(uint32(max))
limiter.interval = duration

return limiter
}

func leakyBucketRate(max uint, duration time.Duration) rate.Limit {
if duration <= 0 {
return rate.Inf
}
return rate.Limit(float64(max) / duration.Seconds())
}
57 changes: 55 additions & 2 deletions ratelimit_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,10 @@ func TestRateLimit(t *testing.T) {
// take another one above max
limiter.Take()
took = time.Since(start).Nanoseconds()
require.GreaterOrEqual(t, took, expected.Nanoseconds())
// Runtime timers may wake fractionally before their nominal deadline on
// some platforms. A small tolerance still proves that the full refill
// window was enforced without making CI depend on timer granularity.
require.GreaterOrEqual(t, took, (expected - 10*time.Millisecond).Nanoseconds())
})

t.Run("Unlimited Rate Limit", func(t *testing.T) {
Expand Down Expand Up @@ -93,7 +96,10 @@ func TestRateLimit(t *testing.T) {
limiter.Take()
limiter.Take()
limiter.Take()
require.False(t, limiter.CanTake())
// The token producer decrements its atomic count immediately after the
// unbuffered handoff. Under the race detector the receiver can resume in
// that tiny window, so wait for the producer-side accounting to settle.
require.Eventually(t, func() bool { return !limiter.CanTake() }, time.Second, time.Millisecond)
})

t.Run("LeakyBucket", func(t *testing.T) {
Expand All @@ -108,4 +114,51 @@ func TestRateLimit(t *testing.T) {
expected := 3 * time.Second
require.True(t, took >= expected)
})

t.Run("LeakyBucket applies max per duration", func(t *testing.T) {
limiter := NewLeakyBucket(context.Background(), 50, time.Second)
require.InDelta(t, 50, float64(limiter.leakyBucketLimiter.Limit()), 0.001)

limiter.SetLimit(25)
require.InDelta(t, 25, float64(limiter.leakyBucketLimiter.Limit()), 0.001)
limiter.SetDuration(500 * time.Millisecond)
require.InDelta(t, 50, float64(limiter.leakyBucketLimiter.Limit()), 0.001)
})

t.Run("LeakyBucket observes cancellation", func(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
limiter := NewLeakyBucket(ctx, 1, time.Hour)
limiter.Take()

done := make(chan struct{})
go func() {
limiter.Take()
close(done)
}()
cancel()

select {
case <-done:
case <-time.After(time.Second):
t.Fatal("Take remained blocked after limiter context cancellation")
}
})

t.Run("LeakyBucket stop unblocks waiters", func(t *testing.T) {
limiter := NewLeakyBucket(context.Background(), 1, time.Hour)
limiter.Take()

done := make(chan struct{})
go func() {
limiter.Take()
close(done)
}()
limiter.Stop()

select {
case <-done:
case <-time.After(time.Second):
t.Fatal("Take remained blocked after limiter stop")
}
})
}
Loading