diff --git a/README.md b/README.md index 7336fe8..3a4ff29 100644 --- a/README.md +++ b/README.md @@ -15,6 +15,20 @@ The original library i.e `golang.org/x/time/rate` implements classic **token buc This allows scanners to respect maximum defined rate limits, pause until the allowed interval hits, and then process again at maximum speed. The original library slowed down requests according to the refill ratio. +## Unlimited mode + +`NewUnlimited(ctx)` permits requests immediately without a ticker, token channel, +or background goroutine. `CanTake` returns true and `GetLimit` initially returns +`math.MaxUint32`. `MultiLimiter` and `AutoLimiter` use this behavior for unlimited keys. + +For compatibility, calling `SetLimit` or `SetDuration` creates a finite burst bucket. +The initial burst is `math.MaxUint32`; a changed limit applies at the next refill. +The default refill interval is 1 ms. The refill schedule starts when the first +setter is called. `SetDuration` requires a positive duration. + +`Stop` prevents later setters from starting a bucket. Requests remain nonblocking +when an unconfigured unlimited limiter is stopped or its context is canceled. + ## Example An Example showing usage of ratelimit as a library is specified below: diff --git a/ratelimit.go b/ratelimit.go index 813f464..7c09cff 100644 --- a/ratelimit.go +++ b/ratelimit.go @@ -3,6 +3,7 @@ package ratelimit import ( "context" "math" + "sync" "sync/atomic" "time" @@ -14,13 +15,14 @@ var minusOne = ^uint32(0) // Limiter allows a burst of request during the defined duration type Limiter struct { - strategy Strategy - maxCount atomic.Uint32 - interval time.Duration - count atomic.Uint32 - ticker *time.Ticker - tokens chan struct{} - ctx context.Context + unlimited *unlimitedLimiter + strategy Strategy + maxCount atomic.Uint32 + interval time.Duration + count atomic.Uint32 + ticker *time.Ticker + tokens chan struct{} + ctx context.Context // internal cancelFunc context.CancelFunc @@ -28,6 +30,32 @@ type Limiter struct { leakyBucketLimiter *rate.Limiter } +// unlimitedLimiter creates a finite bucket only when a caller changes its settings. +// Its mutex protects initialization and Stop; Take never holds it while waiting. +type unlimitedLimiter struct { + mu sync.Mutex + finite *Limiter + stopped bool +} + +func (u *unlimitedLimiter) bucket() *Limiter { + u.mu.Lock() + defer u.mu.Unlock() + + return u.finite +} + +// Called with unlimited.mu held. Preserve the initial burst until the first refill. +func (limiter *Limiter) initUnlimitedBucket() *Limiter { + u := limiter.unlimited + if u.finite == nil && !u.stopped { + u.finite = New(limiter.ctx, math.MaxUint32, limiter.interval) + u.finite.SetLimit(limiter.GetLimit()) + } + + return u.finite +} + func (limiter *Limiter) run(ctx context.Context) { defer close(limiter.tokens) for { @@ -35,6 +63,7 @@ func (limiter *Limiter) run(ctx context.Context) { <-limiter.ticker.C limiter.count.Store(limiter.maxCount.Load()) } + select { case <-ctx.Done(): // Internal Context @@ -53,6 +82,14 @@ func (limiter *Limiter) run(ctx context.Context) { // Take one token from the bucket func (limiter *Limiter) Take() { + if limiter.unlimited != nil { + if finite := limiter.unlimited.bucket(); finite != nil { + finite.Take() + } + + return + } + switch limiter.strategy { case LeakyBucket: _ = limiter.leakyBucketLimiter.Wait(context.TODO()) @@ -63,6 +100,14 @@ func (limiter *Limiter) Take() { // CanTake checks if the rate limiter has any token func (limiter *Limiter) CanTake() bool { + if limiter.unlimited != nil { + if finite := limiter.unlimited.bucket(); finite != nil { + return finite.CanTake() + } + + return true + } + switch limiter.strategy { case LeakyBucket: return limiter.leakyBucketLimiter.Tokens() > 0 @@ -76,9 +121,25 @@ func (limiter *Limiter) GetLimit() uint { return uint(limiter.maxCount.Load()) } -// GetLimit returns current rate limit per given duration +// SetLimit changes the rate limit. Burst buckets apply it at the next refill. +// On an unlimited limiter, it starts a finite bucket with a 1 ms interval unless +// SetDuration has already changed the interval. func (limiter *Limiter) SetLimit(max uint) { + if limiter.unlimited != nil { + limiter.unlimited.mu.Lock() + defer limiter.unlimited.mu.Unlock() + + limiter.maxCount.Store(uint32(max)) + + if finite := limiter.initUnlimitedBucket(); finite != nil { + finite.SetLimit(max) + } + + return + } + limiter.maxCount.Store(uint32(max)) + switch limiter.strategy { case LeakyBucket: limiter.leakyBucketLimiter.SetBurst(int(max)) @@ -86,8 +147,25 @@ func (limiter *Limiter) SetLimit(max uint) { } } -// GetLimit returns current rate limit per given duration +// SetDuration changes the refill interval. For burst buckets, it panics if d is not positive. +// On an unlimited limiter, it starts a finite bucket using the current limit. func (limiter *Limiter) SetDuration(d time.Duration) { + if limiter.unlimited != nil { + limiter.unlimited.mu.Lock() + defer limiter.unlimited.mu.Unlock() + + if d <= 0 { + panic("non-positive interval for Ticker.Reset") + } + limiter.interval = d + + if finite := limiter.initUnlimitedBucket(); finite != nil { + finite.SetDuration(d) + } + + return + } + limiter.interval = d switch limiter.strategy { case LeakyBucket: @@ -99,6 +177,18 @@ func (limiter *Limiter) SetDuration(d time.Duration) { // Stop the rate limiter canceling the internal context func (limiter *Limiter) Stop() { + if limiter.unlimited != nil { + limiter.unlimited.mu.Lock() + defer limiter.unlimited.mu.Unlock() + + limiter.unlimited.stopped = true + if limiter.unlimited.finite != nil { + limiter.unlimited.finite.Stop() + } + + return + } + switch limiter.strategy { case LeakyBucket: // NOP default: @@ -122,36 +212,40 @@ func New(ctx context.Context, max uint, duration time.Duration) *Limiter { strategy: None, interval: duration, } + limiter.maxCount.Store(uint32(max)) limiter.count.Store(uint32(max)) + go limiter.run(internalctx) return limiter } -// NewUnlimited create a bucket with approximated unlimited tokens +// NewUnlimited creates a limiter whose Take never waits and CanTake returns true +// until SetLimit or SetDuration is called. +// It creates no timer, token channel, or background goroutine until SetLimit or +// SetDuration configures a finite bucket. GetLimit initially returns math.MaxUint32. func NewUnlimited(ctx context.Context) *Limiter { - internalctx, cancel := context.WithCancel(context.TODO()) limiter := &Limiter{ - ticker: time.NewTicker(time.Millisecond), - tokens: make(chan struct{}), - ctx: ctx, - cancelFunc: cancel, + unlimited: &unlimitedLimiter{}, + ctx: ctx, + interval: time.Millisecond, } + limiter.maxCount.Store(math.MaxUint32) - limiter.count.Store(math.MaxUint32) - go limiter.run(internalctx) return limiter } -// NewUnlimited create a bucket with approximated unlimited tokens +// NewLeakyBucket creates a limiter that uses golang.org/x/time/rate. func NewLeakyBucket(ctx context.Context, max uint, duration time.Duration) *Limiter { limiter := &Limiter{ strategy: LeakyBucket, leakyBucketLimiter: rate.NewLimiter(rate.Every(duration), int(max)), } + limiter.maxCount.Store(uint32(max)) limiter.interval = duration + return limiter } diff --git a/unlimited_benchmark_test.go b/unlimited_benchmark_test.go new file mode 100644 index 0000000..43dff95 --- /dev/null +++ b/unlimited_benchmark_test.go @@ -0,0 +1,54 @@ +package ratelimit + +import ( + "context" + "testing" +) + +func BenchmarkUnlimitedTake(b *testing.B) { + limiter := NewUnlimited(context.Background()) + defer limiter.Stop() + b.ReportAllocs() + for b.Loop() { + limiter.Take() + } +} + +func BenchmarkUnlimitedTakeParallel(b *testing.B) { + limiter := NewUnlimited(context.Background()) + defer limiter.Stop() + b.ReportAllocs() + b.ResetTimer() + b.RunParallel(func(pb *testing.PB) { + for pb.Next() { + limiter.Take() + } + }) +} + +func BenchmarkUnlimitedLifecycle(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + limiter := NewUnlimited(context.Background()) + limiter.Stop() + // Wait for the old implementation to exit so workers do not accumulate. + if limiter.tokens != nil { + for range limiter.tokens { + } + } + } +} + +func BenchmarkMultiLimiterUnlimitedTake(b *testing.B) { + limiter, err := NewMultiLimiter(context.Background(), &Options{Key: "unlimited", IsUnlimited: true}) + if err != nil { + b.Fatal(err) + } + defer limiter.Stop() + b.ReportAllocs() + for b.Loop() { + if err := limiter.Take("unlimited"); err != nil { + b.Fatal(err) + } + } +} diff --git a/unlimited_idle_linux_test.go b/unlimited_idle_linux_test.go new file mode 100644 index 0000000..4ff57ef --- /dev/null +++ b/unlimited_idle_linux_test.go @@ -0,0 +1,51 @@ +package ratelimit + +import ( + "context" + "fmt" + "runtime" + "syscall" + "testing" + "time" +) + +// Run separately from other benchmarks: CPU usage includes the entire process. +func BenchmarkUnlimitedIdle(b *testing.B) { + for _, n := range []int{0, 30, 1000} { + b.Run(fmt.Sprint(n), func(b *testing.B) { + before := runtime.NumGoroutine() + limiters := make([]*Limiter, n) + for i := range limiters { + limiters[i] = NewUnlimited(context.Background()) + } + defer func() { + for _, limiter := range limiters { + limiter.Stop() + } + for _, limiter := range limiters { + if limiter.tokens != nil { + for range limiter.tokens { + } + } + } + }() + workers := runtime.NumGoroutine() - before + cpuTime := func() time.Duration { + var usage syscall.Rusage + if err := syscall.Getrusage(syscall.RUSAGE_SELF, &usage); err != nil { + b.Fatal(err) + } + return time.Duration(usage.Utime.Nano() + usage.Stime.Nano()) + } + b.ResetTimer() + startCPU, start := cpuTime(), time.Now() + for i := 0; i < b.N; i++ { + time.Sleep(100 * time.Millisecond) + } + elapsed, cpu := time.Since(start), cpuTime()-startCPU + b.StopTimer() + b.ReportMetric(float64(cpu.Nanoseconds())/elapsed.Seconds(), "cpu-ns/s") + b.ReportMetric(float64(workers), "goroutines") + }) + } +} diff --git a/unlimited_test.go b/unlimited_test.go new file mode 100644 index 0000000..5cd2a41 --- /dev/null +++ b/unlimited_test.go @@ -0,0 +1,158 @@ +package ratelimit + +import ( + "context" + "fmt" + "math" + "sync" + "testing" + "testing/synctest" + "time" + + "github.com/stretchr/testify/require" +) + +func TestUnlimitedNoResources(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + limiter := NewUnlimited(context.Background()) + defer limiter.Stop() + require.Nil(t, limiter.ticker, "unlimited mode must not schedule refills") + require.Nil(t, limiter.tokens, "unlimited mode must not exchange tokens") + require.Nil(t, limiter.cancelFunc, "unlimited mode must not start a worker") + require.Equal(t, uint(math.MaxUint32), limiter.GetLimit()) + for i := 0; i < 1000; i++ { + limiter.Take() + require.True(t, limiter.CanTake()) + } + }) +} + +func TestUnlimitedLifecycle(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + limiter := NewUnlimited(ctx) + cancel() + limiter.Stop() + limiter.Stop() + limiter.Take() + require.True(t, limiter.CanTake()) + }) +} + +func TestUnlimitedSetters(t *testing.T) { + for _, durationFirst := range []bool{false, true} { + t.Run(fmt.Sprint(durationFirst), func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + limiter := NewUnlimited(context.Background()) + defer limiter.Stop() + if durationFirst { + limiter.SetDuration(time.Second) + limiter.SetLimit(2) + } else { + limiter.SetLimit(2) + limiter.SetDuration(time.Second) + } + require.Equal(t, uint(2), limiter.GetLimit()) + // SetLimit changes the next refill, not the existing burst. + for i := 0; i < 3; i++ { + limiter.Take() + } + time.Sleep(time.Second) + synctest.Wait() + limiter.Take() + limiter.Take() + synctest.Wait() + require.False(t, limiter.CanTake()) + start := time.Now() + limiter.Take() + require.Equal(t, time.Second, time.Since(start)) + }) + }) + } +} + +func TestUnlimitedSetLimitDefaultDuration(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + limiter := NewUnlimited(context.Background()) + defer limiter.Stop() + t.Cleanup(func() { limiter.Stop(); time.Sleep(time.Millisecond); synctest.Wait() }) + limiter.SetLimit(1) + time.Sleep(time.Millisecond) + synctest.Wait() + limiter.Take() + start := time.Now() + limiter.Take() + require.Equal(t, time.Millisecond, time.Since(start)) + }) +} + +func TestUnlimitedSettersAfterStop(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + limiter := NewUnlimited(context.Background()) + limiter.Stop() + limiter.SetLimit(0) + limiter.SetDuration(time.Hour) + require.Equal(t, uint(0), limiter.GetLimit()) + limiter.Take() + require.Panics(t, func() { limiter.SetDuration(0) }) + require.Panics(t, func() { limiter.SetDuration(-time.Second) }) + }) +} + +func TestUnlimitedConcurrentCalls(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + limiter := NewUnlimited(context.Background()) + var wg sync.WaitGroup + for i := 0; i < 20; i++ { + wg.Go(func() { + limiter.Take() + limiter.CanTake() + limiter.GetLimit() + limiter.SetLimit(math.MaxUint32) + limiter.SetDuration(time.Millisecond) + limiter.Stop() + }) + } + wg.Wait() + limiter.Take() + }) +} + +func TestUnlimitedCanceledBeforeConfiguration(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + limiter := NewUnlimited(ctx) + cancel() + limiter.SetDuration(time.Second) + synctest.Wait() + limiter.Take() + limiter.Stop() + }) +} + +func TestUnlimitedWrappers(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + multi, err := NewMultiLimiter(context.Background(), &Options{Key: "unlimited", IsUnlimited: true}) + require.NoError(t, err) + defer multi.Stop() + auto := NewAutoLimiter(context.Background(), WithUnlimited()) + defer auto.Stop() + for i := 0; i < 100; i++ { + require.NoError(t, multi.Take("unlimited")) + require.True(t, multi.CanTake("unlimited")) + require.NoError(t, auto.Take("unlimited")) + } + direct, err := multi.get("unlimited") + require.NoError(t, err) + automatic, err := auto.get("unlimited") + require.NoError(t, err) + for _, limiter := range []*Limiter{direct, automatic} { + require.Nil(t, limiter.ticker) + require.Nil(t, limiter.tokens) + require.Nil(t, limiter.cancelFunc) + require.Equal(t, uint(math.MaxUint32), limiter.GetLimit()) + } + auto.Stop() + require.NoError(t, auto.Take("unlimited")) + }) +}