diff --git a/limiter.go b/limiter.go index 1c6ab9c..27eb522 100644 --- a/limiter.go +++ b/limiter.go @@ -6,7 +6,6 @@ type Limiter[TInput any, TKey comparable] struct { limits []Limit limitFuncs []LimitFunc[TInput] buckets bucketMap[TKey] - waiters syncMap[TKey, *waiter] } // KeyFunc is a function that takes an input and returns a bucket key. diff --git a/limiter_wait.go b/limiter_wait.go index cdc2d58..8ba10fb 100644 --- a/limiter_wait.go +++ b/limiter_wait.go @@ -2,8 +2,6 @@ package rate import ( "context" - "sync" - "sync/atomic" "time" "github.com/clipperhouse/ntime" @@ -30,12 +28,13 @@ import ( // // ctx := context.WithTimeout(ctx, limit.DurationPerToken()) // -// Wait offers best-effort FIFO ordering of requests. Under sustained -// contention (when multiple requests wait longer than ~1ms), the Go runtime's -// mutex starvation mode ensures strict FIFO ordering. Under light load, -// ordering may be less strict but performance is optimized. -func (r *Limiter[TInput, TKey]) Wait(ctx context.Context, input TInput) bool { - return r.waitN(ctx, input, ntime.Now(), 1) +// The returned error will be non-nil if the context is cancelled. +// +// Wait makes no ordering guarantees. Multiple concurrent calls may +// acquire tokens in any order. +func (r *Limiter[TInput, TKey]) Wait(ctx context.Context, input TInput) (bool, error) { + allow, _, err := r.waitNWithDetails(ctx, input, ntime.Now(), 1) + return allow, err } // WaitN will poll [Limiter.AllowN] for a period of time, @@ -59,114 +58,201 @@ func (r *Limiter[TInput, TKey]) Wait(ctx context.Context, input TInput) bool { // // ctx := context.WithTimeout(ctx, limit.DurationPerToken()) // -// WaitN offers best-effort FIFO ordering of requests. Under sustained -// contention (when multiple requests wait longer than ~1ms), the Go runtime's -// mutex starvation mode ensures strict FIFO ordering. Under light load, -// ordering may be less strict but performance is optimized. -func (r *Limiter[TInput, TKey]) WaitN(ctx context.Context, input TInput, n int64) bool { - return r.waitN(ctx, input, ntime.Now(), n) +// The returned error will be non-nil if the context is cancelled. +// +// WaitN makes no ordering guarantees. Multiple concurrent calls may +// acquire tokens in any order. +func (r *Limiter[TInput, TKey]) WaitN(ctx context.Context, input TInput, n int64) (bool, error) { + allow, _, err := r.waitNWithDetails(ctx, input, ntime.Now(), n) + return allow, err } -func (r *Limiter[TInput, TKey]) waitN(ctx context.Context, input TInput, executionTime ntime.Time, n int64) bool { - return r.waitNWithCancellation( - input, - executionTime, - n, - ctx.Deadline, - ctx.Done, - ) +// WaitWithDetails will poll [Allow] for a period of time, +// until it is cancelled by the passed context. It has the +// effect of adding latency to requests instead of refusing +// them immediately. Consider it graceful degradation. +// +// WaitWithDetails will return true if a token becomes available prior to +// the context cancellation, and will consume a token. It will +// return false if not, and therefore not consume a token. It will +// also return the details of the request. +// +// Take care to create an appropriate context. You almost certainly +// want [context.WithTimeout] or [context.WithDeadline]. +// +// You should be conservative, as WaitWithDetails will introduce +// backpressure on your upstream systems -- connections +// may be held open longer, requests may queue in memory. +// +// A good starting place will be to timeout after waiting +// for one token. For example: +// +// ctx := context.WithTimeout(ctx, limit.DurationPerToken()) +// +// The returned error will be non-nil if the context is cancelled. +// +// WaitWithDetails makes no ordering guarantees. Multiple concurrent calls may +// acquire tokens in any order. +func (r *Limiter[TInput, TKey]) WaitWithDetails(ctx context.Context, input TInput) (bool, Details[TInput, TKey], error) { + return r.WaitNWithDetails(ctx, input, 1) } -func (r *Limiter[TInput, TKey]) waitWithCancellation( - input TInput, - startTime ntime.Time, - deadline func() (time.Time, bool), - done func() <-chan struct{}, -) bool { - return r.waitNWithCancellation(input, startTime, 1, deadline, done) +// WaitNWithDetails will poll [Limiter.AllowN] for a period of time, +// until it is cancelled by the passed context. It has the +// effect of adding latency to requests instead of refusing +// them immediately. Consider it graceful degradation. +// +// WaitNWithDetails will return true if `n` tokens become available prior to +// the context cancellation, and will consume `n` tokens. If not, +// it will return false, and therefore consume no tokens. It will +// also return the details of the request. +// +// Take care to create an appropriate context. You almost certainly +// want [context.WithTimeout] or [context.WithDeadline]. +// +// You should be conservative, as WaitWithDetails will introduce +// backpressure on your upstream systems -- connections +// may be held open longer, requests may queue in memory. +// +// A good starting place will be to timeout after waiting +// for one token. For example: +// +// ctx := context.WithTimeout(ctx, limit.DurationPerToken()) +// +// The returned error will be non-nil if the context is cancelled. +// +// WaitWithDetails makes no ordering guarantees. Multiple concurrent calls may +// acquire tokens in any order. +func (r *Limiter[TInput, TKey]) WaitNWithDetails(ctx context.Context, input TInput, n int64) (bool, Details[TInput, TKey], error) { + return r.waitNWithDetails(ctx, input, ntime.Now(), n) } -// waitWithCancellation is a more testable version of wait that accepts -// deadline and done functions instead of a context, allowing for deterministic testing. -func (r *Limiter[TInput, TKey]) waitNWithCancellation( +// waitNWithDetails is the internal implementation that accepts a context. +// It is designed to be testable by accepting any context implementation, +// including test contexts that provide deterministic behavior. +func (r *Limiter[TInput, TKey]) waitNWithDetails( + ctx context.Context, input TInput, startTime ntime.Time, n int64, - deadline func() (time.Time, bool), - done func() <-chan struct{}, -) bool { - if r.allowN(input, startTime, n) { - return true - } - - // currentTime is an approximation of the real clock moving forward - // it's imprecise because it depends on time.After below. - // For testing purposes, we want startTime (execution time) to - // be a parameter. +) (bool, Details[TInput, TKey], error) { + // currentTime is an approximation of the real clock moving forward. + // It's imprecise because it depends on time.After below. currentTime := startTime - userKey := r.keyFunc(input) - waiter := r.getWaiter(userKey) - - // Ensure cleanup happens when this waiter exits - defer func() { - // If no more waiters for this key, remove the entry to prevent memory leak, - // reference-counted - if waiter.decrement() == 0 { - r.waiters.delete(userKey) - } - }() - for { - // The goroutine at the front of the queue gets to try for a token first. - waiter.mu.Lock() allow, details := r.allowNWithDetails(input, currentTime, n) - waiter.mu.Unlock() if allow { - return true + return allow, details, nil } retryAfter := details.RetryAfter() - // if we can't possibly get a token, fail fast - if deadline, ok := deadline(); ok { - if deadline.Before(currentTime.Add(retryAfter).ToTime()) { - return false - } - } - select { - case <-done(): - return false + case <-ctx.Done(): + // Need to get updated details, since this cancellation + // event might have been a while after the last call. + // We'll choose the semantics of "cancellation always + // means deny". + _, details := r.peekNWithDetails(input, currentTime, n) + return false, details, ctx.Err() case <-time.After(retryAfter): currentTime = currentTime.Add(retryAfter) } } } -// waiter represents a reservation queue for a specific key with reference counting -type waiter struct { - mu sync.Mutex - count int64 // number of active waiters for this key +// WaitWithDebug will poll [Limiter.Allow] for a period of time, +// until it is cancelled by the passed context. It has the +// effect of adding latency to requests instead of refusing +// them immediately. Consider it graceful degradation. +// +// WaitWithDebug will return true if a token becomes available prior to +// the context cancellation, and will consume a token. It will +// return false if not, and therefore not consume a token. It will +// also return the debugs of the request. +// +// Take care to create an appropriate context. You almost certainly +// want [context.WithTimeout] or [context.WithDeadline]. +// +// The returned error will be non-nil if the context is cancelled. +// +// WaitNWithDebug makes no ordering guarantees. Multiple concurrent calls may +// acquire tokens in any order. +func (r *Limiter[TInput, TKey]) WaitWithDebug( + ctx context.Context, + input TInput, +) (bool, []Debug[TInput, TKey], error) { + return r.WaitNWithDebug(ctx, input, 1) } -// increment atomically increments the waiter count and returns the new count -func (w *waiter) increment() int64 { - return atomic.AddInt64(&w.count, 1) +// WaitNWithDebug will poll [Limiter.AllowN] for a period of time, +// until it is cancelled by the passed context. It has the +// effect of adding latency to requests instead of refusing +// them immediately. Consider it graceful degradation. +// +// WaitNWithDebug will return true if `n` tokens become available prior to +// the context cancellation, and will consume `n` tokens. If not, +// it will return false, and therefore consume no tokens. It will +// also return the debugs of the request. +// +// Take care to create an appropriate context. You almost certainly +// want [context.WithTimeout] or [context.WithDeadline]. +// +// The returned error will be non-nil if the context is cancelled. +// +// WaitNWithDebug makes no ordering guarantees. Multiple concurrent calls may +// acquire tokens in any order. +func (r *Limiter[TInput, TKey]) WaitNWithDebug( + ctx context.Context, + input TInput, + n int64, +) (bool, []Debug[TInput, TKey], error) { + return r.waitNWithDebug(ctx, input, ntime.Now(), n) } -// decrement atomically decrements the waiter count and returns the new count -func (w *waiter) decrement() int64 { - return atomic.AddInt64(&w.count, -1) -} +// waitNWithDetails is the internal implementation that accepts a context. +// It is designed to be testable by accepting any context implementation, +// including test contexts that provide deterministic behavior. +func (r *Limiter[TInput, TKey]) waitNWithDebug( + ctx context.Context, + input TInput, + startTime ntime.Time, + n int64, +) (bool, []Debug[TInput, TKey], error) { + // currentTime is an approximation of the real clock moving forward. + // It's imprecise because it depends on time.After below. + currentTime := startTime -func getWaiter() *waiter { - return &waiter{} -} + for { + allow, debugs := r.allowNWithDebug(input, currentTime, n) + if allow { + return allow, debugs, nil + } + + // len(debugs) will be greater than 0 here; + // if there were zero, allow was true + + // Get the max retryAfter + retryAfter := debugs[0].RetryAfter() + for i := 1; i < len(debugs); i++ { + debug := debugs[i] + ra := debug.RetryAfter() + if ra > retryAfter { + retryAfter = ra + } + } -// getWaiter atomically gets or creates a waiter entry and increments its reference count -func (r *Limiter[TInput, TKey]) getWaiter(key TKey) *waiter { - waiter := r.waiters.loadOrStore(key, getWaiter) - waiter.increment() - return waiter + select { + case <-ctx.Done(): + // Need to get updated debugs, since this cancellation + // event might have been a while after the last call. + // We'll choose the semantics of "cancellation always + // means deny". + _, debugs := r.peekNWithDebug(input, currentTime, n) + return false, debugs, ctx.Err() + case <-time.After(retryAfter): + currentTime = currentTime.Add(retryAfter) + } + } } diff --git a/limiter_wait_test.go b/limiter_wait_test.go index e3d260c..f99607d 100644 --- a/limiter_wait_test.go +++ b/limiter_wait_test.go @@ -4,7 +4,6 @@ import ( "context" "fmt" "sync" - "sync/atomic" "testing" "time" @@ -12,1024 +11,557 @@ import ( "github.com/stretchr/testify/require" ) -func TestLimiter_Wait_SingleBucket(t *testing.T) { +func TestLimiter_Wait(t *testing.T) { t.Parallel() - keyFunc := func(input string) string { - return input - } - limit := NewLimit(2, 100*time.Millisecond) - limiter := NewLimiter(keyFunc, limit) - - executionTime := ntime.Now() - - // Consume all tokens - for range limit.count { - ok := limiter.allow("test", executionTime) - require.True(t, ok, "should allow initial tokens") - } - - // Should not allow immediately - ok := limiter.allow("test", executionTime) - require.False(t, ok, "should not allow when tokens exhausted") - - // Test 1: Wait with enough time to acquire a token - { - // Deadline that gives enough time - deadline := func() (time.Time, bool) { - return executionTime.Add(limit.durationPerToken).ToTime(), true - } - - // Done channel that never closes (no cancellation) - done := func() <-chan struct{} { - return make(chan struct{}) // never closes - } - - allow := limiter.waitWithCancellation("test", executionTime, deadline, done) - require.True(t, allow, "should acquire token after waiting") - - // We waited - executionTime = executionTime.Add(limit.durationPerToken) - } - - // Should not allow again immediately after wait - ok = limiter.allow("test", executionTime) - require.False(t, ok, "should not allow again immediately after wait") - - // Test 2: Wait with deadline that expires before token is available - { - // Deadline that expires too soon - deadline := func() (time.Time, bool) { - return executionTime.Add(limit.durationPerToken / 2).ToTime(), true - } - - // Done channel that never closes - done := func() <-chan struct{} { - return make(chan struct{}) // never closes - } - - allow := limiter.waitWithCancellation("test", executionTime, deadline, done) - require.False(t, allow, "should not acquire token if deadline expires before token is available") - } - - // Test 3: Wait with immediate cancellation - { - // Deadline that gives enough time - deadline := func() (time.Time, bool) { - return executionTime.Add(limit.durationPerToken).ToTime(), true - } - - // Done channel that closes immediately - done := func() <-chan struct{} { - ch := make(chan struct{}) - close(ch) // immediately closed - return ch - } - allow := limiter.waitWithCancellation("test", executionTime, deadline, done) - require.False(t, allow, "should not acquire token if context is cancelled immediately") - } - - // Test 4: Wait with no deadline - { - // Deadline that returns no deadline - deadline := func() (time.Time, bool) { - return time.Time{}, false + t.Run("SingleBucket", func(t *testing.T) { + t.Parallel() + keyFunc := func(input string) string { + return input } + limit := NewLimit(2, 100*time.Millisecond) + limiter := NewLimiter(keyFunc, limit) - // Done channel that never closes - done := func() <-chan struct{} { - return make(chan struct{}) // never closes - } + executionTime := ntime.Now() - { - allow := limiter.waitWithCancellation("test", executionTime, deadline, done) - require.True(t, allow, "should acquire token when no deadline is set") - } - - // We waited - executionTime = executionTime.Add(limit.durationPerToken) - - // Should not allow again immediately - { - allow := limiter.allow("test", executionTime) - require.False(t, allow, "should not allow again immediately after wait") + // Consume all tokens + for range limit.count { + ok := limiter.allow("test", executionTime) + require.True(t, ok, "should allow initial tokens") } - } -} - -func TestLimiter_WaitN_SingleBucket(t *testing.T) { - t.Parallel() - keyFunc := func(input string) string { - return input - } - limit := NewLimit(2, 100*time.Millisecond) - limiter := NewLimiter(keyFunc, limit) - executionTime := ntime.Now() - - // Consume all tokens - { - ok := limiter.allowN("test", executionTime, limit.count) - require.True(t, ok, "should allow initial tokens") - } - { // Should not allow immediately - ok := limiter.allowN("test", executionTime, 1) + ok := limiter.allow("test", executionTime) require.False(t, ok, "should not allow when tokens exhausted") - } - - // Test 1: Wait with enough time to acquire tokens - { - wait := time.Duration(limit.count) * limit.durationPerToken - // Deadline that gives enough time - deadline := func() (time.Time, bool) { - return executionTime.Add(wait).ToTime(), true - } - // Done channel that never closes (no cancellation) - done := func() <-chan struct{} { - return make(chan struct{}) // never closes - } + // Test 1: Wait for 1 token with no cancellation + { + ctx := &testContext{ + done: make(chan struct{}), // never closes + } - allow := limiter.waitNWithCancellation("test", executionTime, limit.count, deadline, done) - require.True(t, allow, "should acquire tokens after waiting") + allow, details, err := limiter.waitNWithDetails(ctx, "test", executionTime, 1) + require.True(t, allow, "should acquire 1 token after waiting") + require.NotZero(t, details, "should return details") + require.NoError(t, err, "should not return error") - // We waited - executionTime = executionTime.Add(limit.durationPerToken) - } + // We waited + executionTime = executionTime.Add(limit.durationPerToken) + } - { // Should not allow again immediately after wait - ok := limiter.allow("test", executionTime) + ok = limiter.allow("test", executionTime) require.False(t, ok, "should not allow again immediately after wait") - } - - // Test 2: Wait with deadline that expires before n tokens are available - { - // Deadline that expires too soon - deadline := func() (time.Time, bool) { - return executionTime.Add(limit.durationPerToken).ToTime(), true - } - - // Done channel that never closes - done := func() <-chan struct{} { - return make(chan struct{}) // never closes - } - - allow := limiter.waitNWithCancellation("test", executionTime, limit.count, deadline, done) - require.False(t, allow, "should not acquire tokens if deadline expires before tokens are available") - } - - // Test 3: Wait with immediate cancellation - { - // Deadline that gives enough time - deadline := func() (time.Time, bool) { - return executionTime.Add(limit.durationPerToken * time.Duration(limit.count)).ToTime(), true - } - - // Done channel that closes immediately - done := func() <-chan struct{} { - ch := make(chan struct{}) - close(ch) // immediately closed - return ch - } - - allow := limiter.waitNWithCancellation("test", executionTime, limit.count, deadline, done) - require.False(t, allow, "should not acquire tokens if context is cancelled immediately") - } - - // Test 4: Wait with no deadline - { - // Deadline that returns no deadline - deadline := func() (time.Time, bool) { - return time.Time{}, false - } - - // Done channel that never closes - done := func() <-chan struct{} { - return make(chan struct{}) // never closes - } + // Test 2: Wait for 1 token with immediate cancellation { - allow := limiter.waitNWithCancellation("test", executionTime, limit.count, deadline, done) - require.True(t, allow, "should acquire tokens when no deadline is set") - } - - // We waited - executionTime = executionTime.Add(limit.durationPerToken) - - // Should not allow again immediately - { - allow := limiter.allowN("test", executionTime, limit.count) - require.False(t, allow, "should not allow again immediately after wait") - } - } -} - -func TestLimiter_Wait_MultipleBuckets(t *testing.T) { - t.Parallel() - keyFunc := func(input int) string { - return fmt.Sprintf("test-bucket-%d", input) - } - const buckets = 3 - limit := NewLimit(2, 100*time.Millisecond) - limiter := NewLimiter(keyFunc, limit) - executionTime := ntime.Now() - - // Exhaust tokens for all buckets - for bucketID := range buckets { - for range limit.count { - allow := limiter.allow(bucketID, executionTime) - require.True(t, allow, "should allow initial tokens for bucket %d", bucketID) - } - } - - // Should not allow immediately for any bucket - for bucketID := range buckets { - allow := limiter.allow(bucketID, executionTime) - require.False(t, allow, "should not allow when tokens exhausted for bucket %d", bucketID) - } - - // Wait for a token for each bucket - for bucketID := range buckets { - // Deadline that gives enough time - deadline := func() (time.Time, bool) { - return executionTime.Add(limit.durationPerToken).ToTime(), true - } + // Context that's immediately cancelled (closed) + done := make(chan struct{}) + ctx := &testContext{ + done: done, + } + close(done) - // Done channel that never closes (no cancellation) - done := func() <-chan struct{} { - return make(chan struct{}) // never closes + allow, details, err := limiter.waitNWithDetails(ctx, "test", executionTime, 1) + require.False(t, allow, "should not acquire 1 token if context is cancelled immediately") + require.NotZero(t, details, "should return details even on cancellation") + require.ErrorIs(t, err, context.Canceled, "should return context canceled error") } - allow := limiter.waitWithCancellation(bucketID, executionTime, deadline, done) - require.True(t, allow, "should acquire token after waiting for bucket %d", bucketID) - } - - // We waited - executionTime = executionTime.Add(limit.durationPerToken) - - // Buckets should be empty again - for bucketID := range buckets { - ok := limiter.allow(bucketID, executionTime) - require.False(t, ok, "should not allow again immediately after wait for bucket %d", bucketID) - } -} - -func TestLimiter_Wait_MultipleBuckets_Concurrent(t *testing.T) { - t.Parallel() - keyFunc := func(input int64) string { - return fmt.Sprintf("test-bucket-%d", input) - } - const buckets int64 = 3 - limit := NewLimit(2, 100*time.Millisecond) - limiter := NewLimiter(keyFunc, limit) - executionTime := ntime.Now() - - // Exhaust tokens for all buckets - for bucketID := range buckets { - for range limit.count { - allow := limiter.allow(bucketID, executionTime) - require.True(t, allow, "should allow initial tokens for bucket %d", bucketID) - } - } + // Test 3: Wait for multiple tokens (limit.count) with no cancellation + { + // Wait for bucket to refill enough tokens + executionTime = executionTime.Add(limit.durationPerToken * time.Duration(limit.count)) - // Buckets should be empty - for bucketID := range buckets { - allow := limiter.allow(bucketID, executionTime) - require.False(t, allow, "should not allow when tokens exhausted for bucket %d", bucketID) - } + ctx := &testContext{ + done: make(chan struct{}), // never closes + } - // Test 1: Multiple goroutines competing for tokens with enough time - { - // More goroutines than available tokens to create competition - tokens := buckets * limit.count - concurrency := tokens * 3 // oversubscribe by 3x - results := make([]bool, concurrency) + allow, details, err := limiter.waitNWithDetails(ctx, "test", executionTime, limit.count) + require.True(t, allow, "should acquire %d tokens after waiting", limit.count) + require.NotZero(t, details, "should return details") + require.NoError(t, err, "should not return error") - // Deadline that gives enough time for all tokens to be refilled - deadline := func() (time.Time, bool) { - return executionTime.Add(limit.period).ToTime(), true + // We waited + executionTime = executionTime.Add(limit.durationPerToken) } - // Done channel that never closes - done := func() <-chan struct{} { - return make(chan struct{}) // never closes - } + // Test 4: Wait for multiple tokens with immediate cancellation + { + // Wait for bucket to refill enough tokens + executionTime = executionTime.Add(limit.durationPerToken * time.Duration(limit.count)) - // Start concurrent waits - var wg sync.WaitGroup - for i := range concurrency { - wg.Add(1) - go func(i int64) { - defer wg.Done() - bucketID := i % buckets - results[i] = limiter.waitWithCancellation(bucketID, executionTime, deadline, done) - }(i) - } - wg.Wait() - - // We waited - executionTime = executionTime.Add(limit.period) - - // Count successes and failures - var successes, failures int64 - for _, result := range results { - if result { - successes++ - } else { - failures++ + // Consume all tokens to exhaust the bucket again + for range limit.count { + ok := limiter.allow("test", executionTime) + require.True(t, ok, "should allow tokens to exhaust bucket") } - } - // Exactly limit.count tokens per bucket should succeed, 3 buckets * 2 tokens each - expectedSuccesses := buckets * limit.count - require.Equal(t, expectedSuccesses, successes, "expected exactly %d goroutines to acquire tokens", expectedSuccesses) - - // The rest should fail due to competition - expectedFailures := concurrency - expectedSuccesses - require.Equal(t, expectedFailures, failures, "expected %d goroutines to fail due to competition", expectedFailures) - } - - // Test 2: Multiple goroutines with deadline that expires before tokens are available - { - concurrency := buckets - results := make([]bool, concurrency) + // Context that's immediately cancelled (closed) + done := make(chan struct{}) + ctx := &testContext{ + done: done, + } + close(done) - // Deadline that expires too soon - deadline := func() (time.Time, bool) { - return executionTime.Add(limit.durationPerToken / 2).ToTime(), true + allow, details, err := limiter.waitNWithDetails(ctx, "test", executionTime, limit.count) + require.False(t, allow, "should not acquire %d tokens if context is cancelled immediately", limit.count) + require.NotZero(t, details, "should return details even on cancellation") + require.ErrorIs(t, err, context.Canceled, "should return context canceled error") } + }) - // Done channel that never closes - done := func() <-chan struct{} { - return make(chan struct{}) // never closes - } + t.Run("MultipleBuckets", func(t *testing.T) { + t.Run("Serial", func(t *testing.T) { + t.Parallel() + keyFunc := func(input int) string { + return fmt.Sprintf("test-bucket-%d", input) + } + const buckets = 3 + limit := NewLimit(2, 100*time.Millisecond) + limiter := NewLimiter(keyFunc, limit) + executionTime := ntime.Now() + + // Exhaust tokens for all buckets + for bucketID := range buckets { + for range limit.count { + allow := limiter.allow(bucketID, executionTime) + require.True(t, allow, "should allow initial tokens for bucket %d", bucketID) + } + } - // Start concurrent waits - var wg sync.WaitGroup - for i := range concurrency { - wg.Add(1) - go func(i int64) { - defer wg.Done() - results[i] = limiter.waitWithCancellation(i, executionTime, deadline, done) - }(i) - } + // Should not allow immediately for any bucket + for bucketID := range buckets { + allow := limiter.allow(bucketID, executionTime) + require.False(t, allow, "should not allow when tokens exhausted for bucket %d", bucketID) + } - wg.Wait() + // Wait for a token for each bucket + for bucketID := range buckets { + ctx := &testContext{ + done: make(chan struct{}), // never closes + } - // All should fail because deadline expires before tokens are available - for i, result := range results { - require.False(t, result, "goroutine %d should not acquire token due to early deadline", i) - } - } + allow, details, err := limiter.waitNWithDetails(ctx, bucketID, executionTime, 1) + require.True(t, allow, "should acquire token after waiting for bucket %d", bucketID) + require.NotZero(t, details, "should return details") + require.NoError(t, err, "should not return error") + } - // Test 3: Multiple goroutines with immediate cancellation - { - concurrency := buckets - results := make([]bool, concurrency) + // We waited + executionTime = executionTime.Add(limit.durationPerToken) - // Deadline that gives enough time - deadline := func() (time.Time, bool) { - return executionTime.Add(limit.period).ToTime(), true - } + // Buckets should be empty again + for bucketID := range buckets { + ok := limiter.allow(bucketID, executionTime) + require.False(t, ok, "should not allow again immediately after wait for bucket %d", bucketID) + } + }) - // Done channel that closes immediately - done := func() <-chan struct{} { - ch := make(chan struct{}) - close(ch) // immediately closed - return ch - } + t.Run("Concurrent", func(t *testing.T) { + t.Parallel() + keyFunc := func(input int64) string { + return fmt.Sprintf("test-bucket-%d", input) + } + const buckets int64 = 3 + limit := NewLimit(2, 100*time.Millisecond) + limiter := NewLimiter(keyFunc, limit) + executionTime := ntime.Now() + + // Exhaust tokens for all buckets + for bucketID := range buckets { + for range limit.count { + allow := limiter.allow(bucketID, executionTime) + require.True(t, allow, "should allow initial tokens for bucket %d", bucketID) + } + } - // Start concurrent waits - var wg sync.WaitGroup - for i := range concurrency { - wg.Add(1) - go func(i int64) { - defer wg.Done() - results[i] = limiter.waitWithCancellation(i, executionTime, deadline, done) - }(i) - } + // Buckets should be empty + for bucketID := range buckets { + allow := limiter.allow(bucketID, executionTime) + require.False(t, allow, "should not allow when tokens exhausted for bucket %d", bucketID) + } - wg.Wait() + // Test 1: Multiple goroutines with immediate cancellation + { + concurrency := buckets + results := make([]bool, concurrency) + + // Context that's immediately cancelled (closed) + done := make(chan struct{}) + ctx := &testContext{ + done: done, + } + close(done) + + // Start concurrent waits + var wg sync.WaitGroup + for i := range concurrency { + wg.Add(1) + go func(i int64) { + defer wg.Done() + allow, _, err := limiter.waitNWithDetails(ctx, i, executionTime, 1) + require.False(t, allow, "should not acquire token due to immediate cancellation") + require.ErrorIs(t, err, context.Canceled, "should return context canceled error") + results[i] = allow + }(i) + } + wg.Wait() + + // All should fail because context is cancelled immediately + for i, result := range results { + require.False(t, result, "goroutine %d should not acquire token due to immediate cancellation", i) + } + } - // All should fail because context is cancelled immediately - for i, result := range results { - require.False(t, result, "goroutine %d should not acquire token due to immediate cancellation", i) - } - } + // Test 2: Multiple goroutines with delayed cancellation + { + keyFunc := func(input int64) int64 { + return input + } + + limit := NewLimit(1, 200*time.Millisecond) + limiter := NewLimiter(keyFunc, limit) + + // Exhaust tokens + for bucketID := range buckets { + for range limit.count { + allow := limiter.allow(bucketID, executionTime) + require.True(t, allow, "should allow initial tokens for bucket %d", bucketID) + } + } + + // We're trying to emulate a wait call + // that starts waiting with a retry, but + // is cancelled in the meantime. + + const buckets int64 = 3 + results := make([]bool, buckets) + + var wg sync.WaitGroup + for bucketID := range buckets { + wg.Add(1) + + done := make(chan struct{}) + ctx := &testContext{ + done: done, + } + + go func(bucketID int64) { + defer wg.Done() + // Should start waiting, since the bucket is empty. + // retryAfter should be ~200ms + allow, _, err := limiter.waitNWithDetails(ctx, bucketID, executionTime, 1) + require.False(t, allow, "should not acquire token due to cancellation") + require.ErrorIs(t, err, context.Canceled, "should return context canceled error") + results[bucketID] = allow + }(bucketID) + + // Cancel context after goroutine has started. + // Originally, this had a delay, for a better test, + // but it is flaky on GitHub Actions, presumably + // because the runner is resource constrained. + // Delay works fine locally on Mac M2. + close(done) + } + wg.Wait() + + // All should fail because context is cancelled + for i, result := range results { + require.False(t, result, "goroutine %d should not acquire token due to cancellation", i) + } + } + }) + }) } -func TestLimiter_WaitN_ConsumesCorrectTokens(t *testing.T) { +func TestLimiter_WaitWithDebug(t *testing.T) { t.Parallel() - keyFunc := func(input string) string { - return input - } - limit := NewLimit(10, 100*time.Millisecond) - limiter := NewLimiter(keyFunc, limit) - - executionTime := ntime.Now() - - // Test 1: Wait should consume exactly 1 token - t.Run("Wait_Consumes_One_Token", func(t *testing.T) { - // Verify initial state - _, initialDetails := limiter.peekWithDebug("test-wait-1", executionTime) - require.Equal(t, limit.count, initialDetails[0].TokensRemaining(), "should start with all tokens") - // Deadline that gives enough time - deadline := func() (time.Time, bool) { - return executionTime.Add(time.Second).ToTime(), true - } - done := func() <-chan struct{} { - return make(chan struct{}) // never closes - } - - // Wait should succeed and consume exactly 1 token - allowed := limiter.waitWithCancellation("test-wait-1", executionTime, deadline, done) - require.True(t, allowed, "wait should succeed") - - // Verify exactly 1 token was consumed - _, finalDetails := limiter.peekWithDebug("test-wait-1", executionTime) - require.Equal(t, limit.count-1, finalDetails[0].TokensRemaining(), "should have consumed exactly 1 token") - }) - - // Test 2: WaitN should consume exactly n tokens - t.Run("WaitN_Consumes_N_Tokens", func(t *testing.T) { - const tokensToWait = 3 - - // Verify initial state - _, initialDetails := limiter.peekWithDebug("test-waitn-3", executionTime) - require.Equal(t, limit.count, initialDetails[0].TokensRemaining(), "should start with all tokens") - - // Deadline that gives enough time - deadline := func() (time.Time, bool) { - return executionTime.Add(time.Second).ToTime(), true - } - done := func() <-chan struct{} { - return make(chan struct{}) // never closes + t.Run("SingleBucket", func(t *testing.T) { + t.Parallel() + keyFunc := func(input string) string { + return input } + limit := NewLimit(2, 100*time.Millisecond) + limiter := NewLimiter(keyFunc, limit) - // WaitN should succeed and consume exactly tokensToWait tokens - allowed := limiter.waitNWithCancellation("test-waitn-3", executionTime, tokensToWait, deadline, done) - require.True(t, allowed, "waitN should succeed") + executionTime := ntime.Now() - // Verify exactly tokensToWait tokens were consumed - _, details := limiter.peekWithDebug("test-waitn-3", executionTime) - require.Equal(t, limit.count-tokensToWait, details[0].TokensRemaining(), "should have consumed exactly %d tokens", tokensToWait) - }) - - // Test 3: WaitN with multiple limits should consume n tokens from all buckets - t.Run("WaitN_MultipleLimits_Consumes_N_From_All", func(t *testing.T) { - perSecond := NewLimit(5, time.Second) - perMinute := NewLimit(20, time.Minute) - limiter := NewLimiter(keyFunc, perSecond, perMinute) - const tokensToWait = 2 - - // Verify initial state - _, initialDetails := limiter.peekWithDebug("test-multi-waitn", executionTime) - require.Len(t, initialDetails, 2, "should have details for both limits") - require.Equal(t, perSecond.count, initialDetails[0].TokensRemaining(), "per-second should start with all tokens") - require.Equal(t, perMinute.count, initialDetails[1].TokensRemaining(), "per-minute should start with all tokens") - - // Deadline that gives enough time - deadline := func() (time.Time, bool) { - return executionTime.Add(time.Second).ToTime(), true - } - done := func() <-chan struct{} { - return make(chan struct{}) // never closes - } - - // WaitN should succeed and consume exactly tokensToWait tokens from both limits - allowed := limiter.waitNWithCancellation("test-multi-waitn", executionTime, tokensToWait, deadline, done) - require.True(t, allowed, "waitN should succeed with multiple limits") - - // Verify exactly tokensToWait tokens were consumed from both buckets - _, finalDetails := limiter.peekWithDebug("test-multi-waitn", executionTime) - require.Len(t, finalDetails, 2, "should have details for both limits") - require.Equal(t, perSecond.count-tokensToWait, finalDetails[0].TokensRemaining(), "per-second should have consumed exactly %d tokens", tokensToWait) - require.Equal(t, perMinute.count-tokensToWait, finalDetails[1].TokensRemaining(), "per-minute should have consumed exactly %d tokens", tokensToWait) - }) - - // Test 4: WaitN that fails should consume zero tokens - t.Run("WaitN_Fails_Consumes_Zero_Tokens", func(t *testing.T) { - // First, exhaust the bucket + // Consume all tokens for range limit.count { - limiter.allow("test-fail-waitn", executionTime) + ok := limiter.allow("test", executionTime) + require.True(t, ok, "should allow initial tokens") } - // Verify bucket is exhausted - _, exhaustedDetails := limiter.peekWithDebug("test-fail-waitn", executionTime) - require.Equal(t, int64(0), exhaustedDetails[0].TokensRemaining(), "bucket should be exhausted") - - // Deadline that expires immediately (no time to refill) - deadline := func() (time.Time, bool) { - return executionTime.ToTime(), true // expires immediately - } - done := func() <-chan struct{} { - return make(chan struct{}) // never closes - } - - // WaitN should fail and consume zero tokens - allowed := limiter.waitNWithCancellation("test-fail-waitn", executionTime, 1, deadline, done) - require.False(t, allowed, "waitN should fail when deadline expires before tokens available") - - // Verify no tokens were consumed - _, finalDetails := limiter.peekWithDebug("test-fail-waitn", executionTime) - require.Equal(t, int64(0), finalDetails[0].TokensRemaining(), "should still have 0 tokens after failed wait") - }) + // Should not allow immediately + ok := limiter.allow("test", executionTime) + require.False(t, ok, "should not allow when tokens exhausted") - // Test 5: Verify WaitN with high token count - t.Run("WaitN_High_Token_Count", func(t *testing.T) { - bigLimit := NewLimit(50, time.Second) - bigLimiter := NewLimiter(keyFunc, bigLimit) - const tokensToWait = 25 + // Test 1: Wait for 1 token with no cancellation + { + ctx := &testContext{ + done: make(chan struct{}), // never closes + } - // Verify initial state - _, initialDetails := bigLimiter.peekWithDebug("test-big-waitn", executionTime) - require.Equal(t, bigLimit.count, initialDetails[0].TokensRemaining(), "should start with all tokens") + allow, debugs, err := limiter.waitNWithDebug(ctx, "test", executionTime, 1) + require.True(t, allow, "should acquire 1 token after waiting") + require.NotNil(t, debugs, "should return debugs") + require.NoError(t, err, "should not return error") - // Deadline that gives enough time - deadline := func() (time.Time, bool) { - return executionTime.Add(time.Second).ToTime(), true + // We waited + executionTime = executionTime.Add(limit.durationPerToken) } - done := func() <-chan struct{} { - return make(chan struct{}) // never closes - } - - // WaitN should succeed and consume exactly tokensToWait tokens - allowed := bigLimiter.waitNWithCancellation("test-big-waitn", executionTime, tokensToWait, deadline, done) - require.True(t, allowed, "waitN should succeed with high token count") - // Verify exactly tokensToWait tokens were consumed - _, finalDetails := bigLimiter.peekWithDebug("test-big-waitn", executionTime) - require.Equal(t, bigLimit.count-tokensToWait, finalDetails[0].TokensRemaining(), "should have consumed exactly %d tokens", tokensToWait) - }) - - // Test 6: Concurrent WaitN should consume correct total tokens - t.Run("WaitN_Concurrent_Token_Consumption", func(t *testing.T) { - concurrentLimit := NewLimit(20, time.Second) - concurrentLimiter := NewLimiter(keyFunc, concurrentLimit) - const tokensPerWait = 2 - const numGoroutines = 5 // Will try to consume 10 total tokens - const expectedSuccesses = int64(10) // 20 / 2 = 10 successful waits possible + // Should not allow again immediately after wait + ok = limiter.allow("test", executionTime) + require.False(t, ok, "should not allow again immediately after wait") - results := make([]bool, numGoroutines) + // Test 2: Wait for 1 token with immediate cancellation + { + // Context that's immediately cancelled (closed) + done := make(chan struct{}) + ctx := &testContext{ + done: done, + } + close(done) - // Deadline that gives enough time - deadline := func() (time.Time, bool) { - return executionTime.Add(time.Second).ToTime(), true - } - done := func() <-chan struct{} { - return make(chan struct{}) // never closes + allow, debugs, err := limiter.waitNWithDebug(ctx, "test", executionTime, 1) + require.False(t, allow, "should not acquire 1 token if context is cancelled immediately") + require.NotNil(t, debugs, "should return debugs even on cancellation") + require.ErrorIs(t, err, context.Canceled, "should return context canceled error") } - // Start concurrent waits - var wg sync.WaitGroup - for i := range numGoroutines { - wg.Add(1) - go func(i int) { - defer wg.Done() - results[i] = concurrentLimiter.waitNWithCancellation("test-concurrent-waitn", executionTime, tokensPerWait, deadline, done) - }(i) - } - wg.Wait() + // Test 3: Wait for multiple tokens (limit.count) with no cancellation + { + // Wait for bucket to refill enough tokens + executionTime = executionTime.Add(limit.durationPerToken * time.Duration(limit.count)) - // Count successes - var successes int64 - for _, result := range results { - if result { - successes++ + ctx := &testContext{ + done: make(chan struct{}), // never closes } - } - - require.Equal(t, expectedSuccesses/tokensPerWait, successes, "expected exactly %d successful waits", expectedSuccesses/tokensPerWait) - - // Verify total tokens consumed - _, finalDetails := concurrentLimiter.peekWithDebug("test-concurrent-waitn", executionTime) - expectedRemaining := concurrentLimit.count - (successes * tokensPerWait) - require.Equal(t, expectedRemaining, finalDetails[0].TokensRemaining(), "should have consumed exactly %d tokens total", successes*tokensPerWait) - }) -} - -func TestLimiter_Wait_FIFO_Ordering_SingleBucket_Flaky(t *testing.T) { - t.Skip("this test is flaky because the implementation is not deterministic") - t.Parallel() - keyFunc := func(input string) string { - return "test-bucket" - } - // 1 token per 50ms - limit := NewLimit(1, 50*time.Millisecond) - limiter := NewLimiter(keyFunc, limit) - executionTime := ntime.Now() + allow, debugs, err := limiter.waitNWithDebug(ctx, "test", executionTime, limit.count) + require.True(t, allow, "should acquire %d tokens after waiting", limit.count) + require.NotNil(t, debugs, "should return debugs") + require.NoError(t, err, "should not return error") - // Exhaust the single token - require.True(t, limiter.allow("key", executionTime), "should allow initial token") - require.False(t, limiter.allow("key", executionTime), "should not allow second token") - - const concurrency = 5 - var wg sync.WaitGroup - wg.Add(concurrency) - - startOrder := make(chan int, concurrency) - successOrder := make(chan int, concurrency) - - // Deadline that gives enough time for all tokens to be refilled - deadline := func() (time.Time, bool) { - return executionTime.Add(time.Duration(concurrency) * limit.durationPerToken).ToTime(), true - } - - // Done channel that never closes - done := func() <-chan struct{} { - return make(chan struct{}) // never closes - } - - for i := range concurrency { - go func(id int) { - defer wg.Done() + // We waited + executionTime = executionTime.Add(limit.durationPerToken) + } - // Signal that this goroutine is starting its wait - startOrder <- id + // Test 4: Wait for multiple tokens with immediate cancellation + { + // Wait for bucket to refill enough tokens + executionTime = executionTime.Add(limit.durationPerToken * time.Duration(limit.count)) - // Wait for a token - if limiter.waitWithCancellation("key", executionTime, deadline, done) { - // Signal that this goroutine successfully acquired a token - successOrder <- id + // Consume all tokens to exhaust the bucket again + for range limit.count { + ok := limiter.allow("test", executionTime) + require.True(t, ok, "should allow tokens to exhaust bucket") } - }(i) - // A small, non-deterministic delay to encourage goroutines to queue up in order. - // This helps simulate a real-world scenario where requests arrive sequentially. - time.Sleep(5 * time.Millisecond) - } - - wg.Wait() - - close(startOrder) - close(successOrder) - - var starts []int - for id := range startOrder { - starts = append(starts, id) - } - var successes []int - for id := range successOrder { - successes = append(successes, id) - } - - require.Equal(t, concurrency, len(successes), "all goroutines should have acquired a token") - require.Equal(t, starts, successes, "success order should match start order, proving FIFO") -} - -func TestLimiter_Wait_FIFO_Ordering_MultipleBuckets_Flaky(t *testing.T) { - t.Skip("this test is flaky because the implementation is not deterministic") - t.Parallel() - - // The FIFO behavior is best-effort, this test is known-flaky - // as a result - - const buckets = 3 - const concurrencyPerBucket = 5 - const concurrency = buckets * concurrencyPerBucket - - keyFunc := func(input int) string { - return fmt.Sprintf("test-bucket-%d", input) - } - // 1 token per 50ms, to make the test run reasonably fast - limit := NewLimit(1, 50*time.Millisecond) - limiter := NewLimiter(keyFunc, limit) - executionTime := ntime.Now() - - // Exhaust the single token for each bucket - for i := range buckets { - require.True(t, limiter.allow(i, executionTime), "should allow initial token for bucket %d", i) - require.False(t, limiter.allow(i, executionTime), "should not allow second token for bucket %d", i) - } - - // Create maps to hold order channels for each bucket - startOrders := make(map[int]chan int, buckets) - successOrders := make(map[int]chan int, buckets) - for i := range buckets { - startOrders[i] = make(chan int, concurrencyPerBucket) - successOrders[i] = make(chan int, concurrencyPerBucket) - } - - // Deadline that gives enough time for all tokens to be refilled for all goroutines - deadline := func() (time.Time, bool) { - return executionTime.Add(time.Duration(concurrencyPerBucket) * limit.durationPerToken).ToTime(), true - } - - // Done channel that never closes - done := func() <-chan struct{} { - return make(chan struct{}) // never closes - } - - var wg sync.WaitGroup - wg.Add(concurrency) - for i := range concurrency { - go func(id int) { - defer wg.Done() - bucketID := id % buckets - - // Announce that this goroutine is starting its wait for its bucket - startOrders[bucketID] <- id - - // Wait for a token - if limiter.waitWithCancellation(bucketID, executionTime, deadline, done) { - // Announce that this goroutine successfully acquired a token - successOrders[bucketID] <- id + // Context that's immediately cancelled (closed) + done := make(chan struct{}) + ctx := &testContext{ + done: done, } - }(i) - // A small, non-deterministic delay to encourage goroutines to queue up in order. - time.Sleep(5 * time.Millisecond) - } - - wg.Wait() - - // Close all channels - for bucketID := range buckets { - close(startOrders[bucketID]) - close(successOrders[bucketID]) - } + close(done) - // Verify FIFO order for each bucket - for bucketID := range buckets { - var starts []int - for id := range startOrders[bucketID] { - starts = append(starts, id) + allow, debugs, err := limiter.waitNWithDebug(ctx, "test", executionTime, limit.count) + require.False(t, allow, "should not acquire %d tokens if context is cancelled immediately", limit.count) + require.NotNil(t, debugs, "should return debugs even on cancellation") + require.ErrorIs(t, err, context.Canceled, "should return context canceled error") } + }) - var successes []int - for id := range successOrders[bucketID] { - successes = append(successes, id) - } - - require.Equal(t, concurrencyPerBucket, len(successes), "all goroutines for bucket %d should have acquired a token", bucketID) - require.Equal(t, starts, successes, "success order should match start order for bucket %d, proving FIFO", bucketID) - } -} - -func TestLimiter_WaitersCleanup_Basic(t *testing.T) { - t.Parallel() - - keyFunc := func(input string) string { - return input - } - limit := NewLimit(1, 100*time.Millisecond) - limiter := NewLimiter(keyFunc, limit) - - // Initial waiters count should be 0 - require.Equal(t, 0, limiter.waiters.count(), "initial waiters count should be 0") + t.Run("MultipleBuckets", func(t *testing.T) { + t.Run("Serial", func(t *testing.T) { + t.Parallel() + keyFunc := func(input int) string { + return fmt.Sprintf("test-bucket-%d", input) + } + const buckets = 3 + limit := NewLimit(2, 100*time.Millisecond) + limiter := NewLimiter(keyFunc, limit) + executionTime := ntime.Now() + + // Exhaust tokens for all buckets + for bucketID := range buckets { + for range limit.count { + allow := limiter.allow(bucketID, executionTime) + require.True(t, allow, "should allow initial tokens for bucket %d", bucketID) + } + } - executionTime := ntime.Now() + // Should not allow immediately for any bucket + for bucketID := range buckets { + allow := limiter.allow(bucketID, executionTime) + require.False(t, allow, "should not allow when tokens exhausted for bucket %d", bucketID) + } - // Exhaust tokens for multiple keys - require.True(t, limiter.allow("key1", executionTime), "should allow initial token for key1") - require.True(t, limiter.allow("key2", executionTime), "should allow initial token for key2") + // Wait for a token for each bucket + for bucketID := range buckets { + ctx := &testContext{ + done: make(chan struct{}), // never closes + } - // Create deadline that will timeout immediately - deadline := func() (time.Time, bool) { - return executionTime.ToTime(), true // immediate timeout - } - done := func() <-chan struct{} { - ch := make(chan struct{}) - close(ch) // immediately cancelled - return ch - } + allow, debugs, err := limiter.waitNWithDebug(ctx, bucketID, executionTime, 1) + require.True(t, allow, "should acquire token after waiting for bucket %d", bucketID) + require.NotNil(t, debugs, "should return debugs") + require.NoError(t, err, "should not return error") + } - // Try to wait - this should create a waiter entry that gets cleaned up - allow := limiter.waitWithCancellation("key1", executionTime, deadline, done) - require.False(t, allow, "should timeout immediately") + // We waited + executionTime = executionTime.Add(limit.durationPerToken) - // Check if waiter was cleaned up - require.Equal(t, 0, limiter.waiters.count(), "waiters should be cleaned up after timeout") + // Buckets should be empty again + for bucketID := range buckets { + ok := limiter.allow(bucketID, executionTime) + require.False(t, ok, "should not allow again immediately after wait for bucket %d", bucketID) + } + }) - // Try again with a different key - allow = limiter.waitWithCancellation("key2", executionTime, deadline, done) - require.False(t, allow, "should timeout immediately") + t.Run("Concurrent", func(t *testing.T) { + t.Parallel() + keyFunc := func(input int64) string { + return fmt.Sprintf("test-bucket-%d", input) + } + const buckets int64 = 3 + limit := NewLimit(2, 100*time.Millisecond) + limiter := NewLimiter(keyFunc, limit) + executionTime := ntime.Now() + + // Exhaust tokens for all buckets + for bucketID := range buckets { + for range limit.count { + allow := limiter.allow(bucketID, executionTime) + require.True(t, allow, "should allow initial tokens for bucket %d", bucketID) + } + } - // Check waiters count again - require.Equal(t, 0, limiter.waiters.count(), "waiters should be cleaned up after timeout") + // Buckets should be empty + for bucketID := range buckets { + allow := limiter.allow(bucketID, executionTime) + require.False(t, allow, "should not allow when tokens exhausted for bucket %d", bucketID) + } - // Try the same key again - allow = limiter.waitWithCancellation("key1", executionTime, deadline, done) - require.False(t, allow, "should timeout immediately") + // Test 1: Multiple goroutines with immediate cancellation + { + concurrency := buckets + results := make([]bool, concurrency) + + // Context that's immediately cancelled (closed) + done := make(chan struct{}) + ctx := &testContext{ + done: done, + } + close(done) + + // Start concurrent waits + var wg sync.WaitGroup + for i := range concurrency { + wg.Add(1) + go func(i int64) { + defer wg.Done() + allow, _, err := limiter.waitNWithDebug(ctx, i, executionTime, 1) + require.False(t, allow, "should not acquire token due to immediate cancellation") + require.ErrorIs(t, err, context.Canceled, "should return context canceled error") + results[i] = allow + }(i) + } + wg.Wait() + + // All should fail because context is cancelled immediately + for i, result := range results { + require.False(t, result, "goroutine %d should not acquire token due to immediate cancellation", i) + } + } - // Check waiters count - require.Equal(t, 0, limiter.waiters.count(), "waiters should be cleaned up after timeout") + // Test 2: Multiple goroutines with delayed cancellation + { + keyFunc := func(input int64) int64 { + return input + } + + limit := NewLimit(1, 200*time.Millisecond) + limiter := NewLimiter(keyFunc, limit) + + // Exhaust tokens + for bucketID := range buckets { + for range limit.count { + allow := limiter.allow(bucketID, executionTime) + require.True(t, allow, "should allow initial tokens for bucket %d", bucketID) + } + } + + // We're trying to emulate a wait call + // that starts waiting with a retry, but + // is cancelled in the meantime. + + const buckets int64 = 3 + results := make([]bool, buckets) + + var wg sync.WaitGroup + for bucketID := range buckets { + wg.Add(1) + + done := make(chan struct{}) + ctx := &testContext{ + done: done, + } + + go func(bucketID int64) { + defer wg.Done() + // Should start waiting, since the bucket is empty. + // retryAfter should be ~200ms + allow, _, err := limiter.waitNWithDebug(ctx, bucketID, executionTime, 1) + require.False(t, allow, "should not acquire token due to cancellation") + require.ErrorIs(t, err, context.Canceled, "should return context canceled error") + results[bucketID] = allow + }(bucketID) + + // Cancel context after goroutine has started. + // Originally, this had a delay, for a better test, + // but it is flaky on GitHub Actions, presumably + // because the runner is resource constrained. + // Delay works fine locally on Mac M2. + close(done) + } + wg.Wait() + + // All should fail because context is cancelled + for i, result := range results { + require.False(t, result, "goroutine %d should not acquire token due to cancellation", i) + } + } + }) + }) } -func TestLimiter_WaitersCleanup_MemoryLeak_Prevention(t *testing.T) { - t.Parallel() - - keyFunc := func(input int) string { - return fmt.Sprintf("key-%d", input) - } - limit := NewLimit(1, 100*time.Millisecond) - limiter := NewLimiter(keyFunc, limit) - - executionTime := ntime.Now() - - // Create deadline that will timeout immediately - deadline := func() (time.Time, bool) { - return executionTime.ToTime(), true // immediate timeout - } - done := func() <-chan struct{} { - ch := make(chan struct{}) - close(ch) // immediately cancelled - return ch - } - - const numKeys = 1000 - - // Simulate many different keys trying to wait - for i := range numKeys { - // First exhaust the token for this key - limiter.allow(i, executionTime) - - // Then try to wait (which will timeout immediately) - result := limiter.waitWithCancellation(i, executionTime, deadline, done) - require.False(t, result, "should timeout immediately for key %d", i) - } +var _ context.Context = &testContext{} - require.Equal(t, 0, limiter.waiters.count(), "waiters should be cleaned up after use") +type testContext struct { + done <-chan struct{} } -func TestLimiter_WaitersCleanup_Concurrent(t *testing.T) { - t.Parallel() - - keyFunc := func(input int) string { - return fmt.Sprintf("key-%d", input) - } - limit := NewLimit(1, 100*time.Millisecond) - limiter := NewLimiter(keyFunc, limit) - - executionTime := ntime.Now() - - // Create deadline that will timeout immediately - deadline := func() (time.Time, bool) { - return executionTime.ToTime(), true // immediate timeout - } - done := func() <-chan struct{} { - ch := make(chan struct{}) - close(ch) // immediately cancelled - return ch - } - - const concurrency = 100 - const keysPerGoroutine = 10 - - var wg sync.WaitGroup - wg.Add(concurrency) - for i := range concurrency { - go func(i int) { - defer wg.Done() - - for k := range keysPerGoroutine { - keyID := i*keysPerGoroutine + k - - // First exhaust the token for this key - limiter.allow(keyID, executionTime) - - // Then try to wait (which will timeout immediately) - limiter.waitWithCancellation(keyID, executionTime, deadline, done) - } - }(i) - } - wg.Wait() - - require.Equal(t, 0, limiter.waiters.count(), "waiters should be cleaned up after use") +func (t *testContext) Deadline() (deadline time.Time, ok bool) { + return time.Time{}, false } -func TestLimiter_WaitersCleanup_WithSuccessfulWaits(t *testing.T) { - t.Parallel() - - keyFunc := func(input string) string { - return input - } - limit := NewLimit(1, 100*time.Millisecond) - limiter := NewLimiter(keyFunc, limit) - - executionTime := ntime.Now() - - // Exhaust the token - require.True(t, limiter.allow("key", executionTime), "should allow initial token") - - const concurrency = 20 - results := make([]bool, concurrency) - - // Create deadline that gives enough time for tokens to be refilled - deadline := func() (time.Time, bool) { - return executionTime.Add(time.Duration(concurrency) * limit.durationPerToken).ToTime(), true - } - - done := func() <-chan struct{} { - return make(chan struct{}) // never closes - } - - var wg sync.WaitGroup - wg.Add(concurrency) - for i := range concurrency { - go func(i int) { - defer wg.Done() - // stagger the start times - time.Sleep(time.Duration(i) * time.Millisecond) - results[i] = limiter.waitWithCancellation("key", executionTime, deadline, done) - }(i) - } - wg.Wait() - - // All waiters should have eventually succeeded - successes := 0 - for _, result := range results { - if result { - successes++ - } - } - - require.Equal(t, concurrency, successes, "all waiters should eventually succeed") - require.Equal(t, 0, limiter.waiters.count(), "all waiters should be cleaned up after completion") +func (t *testContext) Done() <-chan struct{} { + return t.done } -func TestLimiter_Wait_FIFOOrdering_HighContention(t *testing.T) { - t.Parallel() - - keyFunc := func(input int) int { - return input - } - - limiter := NewLimiter(keyFunc, NewLimit(1, 50*time.Millisecond)) - - // Exhaust the bucket - require.True(t, limiter.Allow(1)) - - // Test that multiple waiters can successfully acquire tokens - const waiters = 5 - results := make(chan bool, waiters) - - for range waiters { - go func() { - ctx, cancel := context.WithTimeout(context.Background(), 1*time.Second) - defer cancel() - results <- limiter.Wait(ctx, 1) - }() - } - - // Collect results - successes := 0 - for range waiters { - select { - case success := <-results: - if success { - successes++ - } - case <-time.After(2 * time.Second): - t.Fatal("Test timed out") - } +func (t *testContext) Err() error { + select { + case <-t.done: + return context.Canceled + default: + return nil } - - // Should get at least some successes (the fix should prevent token starvation) - require.Greater(t, successes, 0, "Should have some successful token acquisitions") - require.Equal(t, 0, limiter.waiters.count(), "All waiters should be cleaned up") } -func TestLimiter_Wait_RaceCondition_Prevention(t *testing.T) { - t.Parallel() - - keyFunc := func(input int) int { - return input - } - // Use a very restrictive limit to force contention - limiter := NewLimiter(keyFunc, NewLimit(1, 100*time.Millisecond)) - - // Exhaust the bucket - require.True(t, limiter.Allow(1)) - - const concurrency = 50 - successes := int64(0) - - var wg sync.WaitGroup - wg.Add(concurrency) - // Start many goroutines that will compete for tokens - for range concurrency { - go func() { - defer wg.Done() - ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) - defer cancel() - - if limiter.Wait(ctx, 1) { - atomic.AddInt64(&successes, 1) - } - }() - } - wg.Wait() - - // Should have some successes but not more than the tokens available - // in the timeout period (roughly 20 tokens in 2 seconds at 100ms per token) - actual := atomic.LoadInt64(&successes) - require.Greater(t, actual, int64(0), "should have some successful acquisitions") - require.LessOrEqual(t, actual, int64(25), "should not exceed reasonable token availability") - - // Verify no memory leaks - require.Equal(t, 0, limiter.waiters.count(), "all waiters should be cleaned up") +func (t *testContext) Value(key any) any { + return nil } diff --git a/limiters_wait.go b/limiters_wait.go new file mode 100644 index 0000000..4ba431b --- /dev/null +++ b/limiters_wait.go @@ -0,0 +1,251 @@ +package rate + +import ( + "context" + "time" + + "github.com/clipperhouse/ntime" +) + +// Wait will poll [Limiters.Allow] for a period of time, +// until it is cancelled by the passed context. It has the +// effect of adding latency to requests instead of refusing +// them immediately. Consider it graceful degradation. +// +// Wait will return true if a token becomes available prior to +// the context cancellation, and will consume a token. It will +// return false if not, and therefore not consume a token. +// +// Take care to create an appropriate context. You almost certainly +// want [context.WithTimeout] or [context.WithDeadline]. +// +// You should be conservative, as Wait will introduce +// backpressure on your upstream systems -- connections +// may be held open longer, requests may queue in memory. +// +// A good starting place will be to timeout after waiting +// for one token. For example: +// +// ctx := context.WithTimeout(ctx, limit.DurationPerToken()) +// +// The returned error will be non-nil if the context is cancelled. +// +// Wait makes no ordering guarantees. Multiple concurrent calls may +// acquire tokens in any order. +func (r *Limiters[TInput, TKey]) Wait(ctx context.Context, input TInput) (bool, error) { + return r.WaitN(ctx, input, 1) +} + +// WaitN will poll [Limiters.AllowN] for a period of time, +// until it is cancelled by the passed context. It has the +// effect of adding latency to requests instead of refusing +// them immediately. Consider it graceful degradation. +// +// WaitN will return true if `n` tokens become available prior to +// the context cancellation, and will consume `n` tokens. If not, +// it will return false, and therefore consume no tokens. +// +// Take care to create an appropriate context. You almost certainly +// want [context.WithTimeout] or [context.WithDeadline]. +// +// You should be conservative, as Wait will introduce +// backpressure on your upstream systems -- connections +// may be held open longer, requests may queue in memory. +// +// A good starting place will be to timeout after waiting +// for one token. For example: +// +// ctx := context.WithTimeout(ctx, limit.DurationPerToken()) +// +// The returned error will be non-nil if the context is cancelled. +// +// WaitN makes no ordering guarantees. Multiple concurrent calls may +// acquire tokens in any order. +func (r *Limiters[TInput, TKey]) WaitN(ctx context.Context, input TInput, n int64) (bool, error) { + allow, _, err := r.WaitNWithDetails(ctx, input, n) + return allow, err +} + +// WaitWithDetails will poll [Limiters.Allow] for a period of time, +// until it is cancelled by the passed context. It has the +// effect of adding latency to requests instead of refusing +// them immediately. Consider it graceful degradation. +// +// WaitWithDetails will return true if a token becomes available prior to +// the context cancellation, and will consume a token. It will +// return false if not, and therefore not consume a token. It will +// also return the details of the request. +// +// Take care to create an appropriate context. You almost certainly +// want [context.WithTimeout] or [context.WithDeadline]. +// +// You should be conservative, as WaitWithDetails will introduce +// backpressure on your upstream systems -- connections +// may be held open longer, requests may queue in memory. +// +// A good starting place will be to timeout after waiting +// for one token. For example: +// +// ctx := context.WithTimeout(ctx, limit.DurationPerToken()) +// +// The returned error will be non-nil if the context is cancelled. +// +// WaitWithDetails makes no ordering guarantees. Multiple concurrent calls may +// acquire tokens in any order. +func (r *Limiters[TInput, TKey]) WaitWithDetails(ctx context.Context, input TInput) (bool, Details[TInput, TKey], error) { + return r.WaitNWithDetails(ctx, input, 1) +} + +// WaitNWithDetails will poll [Limiters.AllowN] for a period of time, +// until it is cancelled by the passed context. It has the +// effect of adding latency to requests instead of refusing +// them immediately. Consider it graceful degradation. +// +// WaitNWithDetails will return true if `n` tokens become available prior to +// the context cancellation, and will consume `n` tokens. If not, +// it will return false, and therefore consume no tokens. It will +// also return the details of the request. +// +// Take care to create an appropriate context. You almost certainly +// want [context.WithTimeout] or [context.WithDeadline]. +// +// You should be conservative, as WaitWithDetails will introduce +// backpressure on your upstream systems -- connections +// may be held open longer, requests may queue in memory. +// +// A good starting place will be to timeout after waiting +// for one token. For example: +// +// ctx := context.WithTimeout(ctx, limit.DurationPerToken()) +// +// The returned error will be non-nil if the context is cancelled. +// +// WaitWithDetails makes no ordering guarantees. Multiple concurrent calls may +// acquire tokens in any order. +func (r *Limiters[TInput, TKey]) WaitNWithDetails(ctx context.Context, input TInput, n int64) (bool, Details[TInput, TKey], error) { + return r.waitNWithDetails(ctx, input, ntime.Now(), n) +} + +func (r *Limiters[TInput, TKey]) waitNWithDetails( + ctx context.Context, + input TInput, + startTime ntime.Time, + n int64, +) (bool, Details[TInput, TKey], error) { + // currentTime is an approximation of the real clock moving forward. + // It's imprecise because it depends on time.After below. + currentTime := startTime + + for { + allow, details := r.allowNWithDetails(input, currentTime, n) + if allow { + return allow, details, nil + } + + retryAfter := details.RetryAfter() + + select { + case <-ctx.Done(): + // Need to get updated details, since this cancellation + // event might have been a while after the last call. + // We'll choose the semantics of "cancellation always + // means deny". + _, details := r.peekNWithDetails(input, currentTime, n) + return false, details, ctx.Err() + case <-time.After(retryAfter): + currentTime = currentTime.Add(retryAfter) + } + } +} + +// WaitWithDebug will poll [Limiters.Allow] for a period of time, +// until it is cancelled by the passed context. It has the +// effect of adding latency to requests instead of refusing +// them immediately. Consider it graceful degradation. +// +// WaitWithDebug will return true if a token becomes available prior to +// the context cancellation, and will consume a token. It will +// return false if not, and therefore not consume a token. It will +// also return the debugs of the request. +// +// Take care to create an appropriate context. You almost certainly +// want [context.WithTimeout] or [context.WithDeadline]. +// +// The returned error will be non-nil if the context is cancelled. +// +// WaitNWithDebug makes no ordering guarantees. Multiple concurrent calls may +// acquire tokens in any order. +func (r *Limiters[TInput, TKey]) WaitWithDebug( + ctx context.Context, + input TInput, +) (bool, []Debug[TInput, TKey], error) { + return r.WaitNWithDebug(ctx, input, 1) +} + +// WaitNWithDebug will poll [Limiters.AllowN] for a period of time, +// until it is cancelled by the passed context. It has the +// effect of adding latency to requests instead of refusing +// them immediately. Consider it graceful degradation. +// +// WaitNWithDebug will return true if `n` tokens become available prior to +// the context cancellation, and will consume `n` tokens. If not, +// it will return false, and therefore consume no tokens. It will +// also return the debugs of the request. +// +// Take care to create an appropriate context. You almost certainly +// want [context.WithTimeout] or [context.WithDeadline]. +// +// The returned error will be non-nil if the context is cancelled. +// +// WaitNWithDebug makes no ordering guarantees. Multiple concurrent calls may +// acquire tokens in any order. +func (r *Limiters[TInput, TKey]) WaitNWithDebug( + ctx context.Context, + input TInput, + n int64, +) (bool, []Debug[TInput, TKey], error) { + return r.waitNWithDebug(ctx, input, ntime.Now(), n) +} + +func (r *Limiters[TInput, TKey]) waitNWithDebug( + ctx context.Context, + input TInput, + startTime ntime.Time, + n int64, +) (bool, []Debug[TInput, TKey], error) { + // currentTime is an approximation of the real clock moving forward. + // It's imprecise because it depends on time.After below. + currentTime := startTime + + for { + allow, debugs := r.allowNWithDebug(input, currentTime, n) + if allow { + return allow, debugs, nil + } + + // len(debugs) will be greater than 0 here; + // if there were zero, allow was true + + // Get the max retryAfter + retryAfter := debugs[0].RetryAfter() + for i := 1; i < len(debugs); i++ { + debug := debugs[i] + ra := debug.RetryAfter() + if ra > retryAfter { + retryAfter = ra + } + } + + select { + case <-ctx.Done(): + // Need to get updated debugs, since this cancellation + // event might have been a while after the last call. + // We'll choose the semantics of "cancellation always + // means deny". + _, debugs := r.peekNWithDebug(input, currentTime, n) + return false, debugs, ctx.Err() + case <-time.After(retryAfter): + currentTime = currentTime.Add(retryAfter) + } + } +} diff --git a/limiters_wait_test.go b/limiters_wait_test.go new file mode 100644 index 0000000..8d29d99 --- /dev/null +++ b/limiters_wait_test.go @@ -0,0 +1,692 @@ +package rate + +import ( + "context" + "fmt" + "sync" + "testing" + "time" + + "github.com/clipperhouse/ntime" + "github.com/stretchr/testify/require" +) + +func TestLimiters_Wait(t *testing.T) { + t.Parallel() + + t.Run("SingleBucket", func(t *testing.T) { + t.Parallel() + keyFunc := func(input string) string { + return input + } + // Create two limiters with different limits + limit1 := NewLimit(2, 100*time.Millisecond) // 2 per 100ms (more restrictive) + limit2 := NewLimit(5, 100*time.Millisecond) // 5 per 100ms + limiter1 := NewLimiter(keyFunc, limit1) + limiter2 := NewLimiter(keyFunc, limit2) + limiters := Combine(limiter1, limiter2) + + executionTime := ntime.Now() + + // Consume all tokens from the more restrictive limiter + for range limit1.count { + ok := limiters.allowN("test", executionTime, 1) + require.True(t, ok, "should allow initial tokens") + } + + // Should not allow immediately (limiter1 exhausted) + ok := limiters.allowN("test", executionTime, 1) + require.False(t, ok, "should not allow when limiter1 tokens exhausted") + + // Test 1: Wait for 1 token with no cancellation + { + ctx := &testContext{ + done: make(chan struct{}), // never closes + } + + allow, details, err := limiters.waitNWithDetails(ctx, "test", executionTime, 1) + require.True(t, allow, "should acquire 1 token after waiting") + require.NotZero(t, details, "should return details") + require.NoError(t, err, "should not return error") + + // We waited + executionTime = executionTime.Add(limit1.durationPerToken) + } + + // Should not allow again immediately after wait + ok = limiters.allowN("test", executionTime, 1) + require.False(t, ok, "should not allow again immediately after wait") + + // Test 2: Wait for 1 token with immediate cancellation + { + // Context that's immediately cancelled (closed) + done := make(chan struct{}) + ctx := &testContext{ + done: done, + } + close(done) + + allow, details, err := limiters.waitNWithDetails(ctx, "test", executionTime, 1) + require.False(t, allow, "should not acquire 1 token if context is cancelled immediately") + require.NotZero(t, details, "should return details even on cancellation") + require.ErrorIs(t, err, context.Canceled, "should return context canceled error") + } + + // Test 3: Wait for multiple tokens (limit1.count) with no cancellation + { + // Wait for bucket to refill enough tokens + executionTime = executionTime.Add(limit1.period) + + ctx := &testContext{ + done: make(chan struct{}), // never closes + } + + allow, details, err := limiters.waitNWithDetails(ctx, "test", executionTime, limit1.count) + require.True(t, allow, "should acquire %d tokens after waiting", limit1.count) + require.NotZero(t, details, "should return details") + require.NoError(t, err, "should not return error") + } + + // Test 4: Wait for multiple tokens with immediate cancellation + { + // Wait for bucket to refill enough tokens + executionTime = executionTime.Add(limit1.durationPerToken * time.Duration(limit1.count)) + + // Consume all tokens to exhaust the bucket again + for range limit1.count { + ok := limiters.allowN("test", executionTime, 1) + require.True(t, ok, "should allow tokens to exhaust bucket") + } + + // Context that's immediately cancelled (closed) + done := make(chan struct{}) + ctx := &testContext{ + done: done, + } + close(done) + + allow, details, err := limiters.waitNWithDetails(ctx, "test", executionTime, limit1.count) + require.False(t, allow, "should not acquire %d tokens if context is cancelled immediately", limit1.count) + require.NotZero(t, details, "should return details even on cancellation") + require.ErrorIs(t, err, context.Canceled, "should return context canceled error") + } + }) + + t.Run("MultipleBuckets", func(t *testing.T) { + t.Run("Serial", func(t *testing.T) { + t.Parallel() + keyFunc := func(input int) string { + return fmt.Sprintf("test-bucket-%d", input) + } + const buckets = 3 + // Create two limiters with different limits + limit1 := NewLimit(2, 100*time.Millisecond) // 2 per 100ms (more restrictive) + limit2 := NewLimit(4, 100*time.Millisecond) // 4 per 100ms + limiter1 := NewLimiter(keyFunc, limit1) + limiter2 := NewLimiter(keyFunc, limit2) + limiters := Combine(limiter1, limiter2) + executionTime := ntime.Now() + + // Exhaust tokens for all buckets from the more restrictive limiter + for bucketID := range buckets { + for range limit1.count { + allow := limiters.allowN(bucketID, executionTime, 1) + require.True(t, allow, "should allow initial tokens for bucket %d", bucketID) + } + } + + // Should not allow immediately for any bucket (limiter1 exhausted) + for bucketID := range buckets { + allow := limiters.allowN(bucketID, executionTime, 1) + require.False(t, allow, "should not allow when limiter1 tokens exhausted for bucket %d", bucketID) + } + + // Wait for a token for each bucket + for bucketID := range buckets { + ctx := &testContext{ + done: make(chan struct{}), // never closes + } + + allow, details, err := limiters.waitNWithDetails(ctx, bucketID, executionTime, 1) + require.True(t, allow, "should acquire token after waiting for bucket %d", bucketID) + require.NotZero(t, details, "should return details") + require.NoError(t, err, "should not return error") + } + + // We waited + executionTime = executionTime.Add(limit1.durationPerToken) + + // Buckets should be empty again + for bucketID := range buckets { + ok := limiters.allowN(bucketID, executionTime, 1) + require.False(t, ok, "should not allow again immediately after wait for bucket %d", bucketID) + } + }) + + t.Run("Concurrent", func(t *testing.T) { + // t.Parallel() + // GitHub Actions are pretty resource constrained, + // and this test is sensitive to timing. + keyFunc := func(input int64) string { + return fmt.Sprintf("test-bucket-%d", input) + } + const buckets int64 = 3 + // Create two limiters with different limits + limit1 := NewLimit(2, 100*time.Millisecond) // 2 per 100ms (more restrictive) + limit2 := NewLimit(4, 100*time.Millisecond) // 4 per 100ms + limiter1 := NewLimiter(keyFunc, limit1) + limiter2 := NewLimiter(keyFunc, limit2) + limiters := Combine(limiter1, limiter2) + executionTime := ntime.Now() + + // Exhaust tokens for all buckets from the more restrictive limiter + for bucketID := range buckets { + for range limit1.count { + allow := limiters.allowN(bucketID, executionTime, 1) + require.True(t, allow, "should allow initial tokens for bucket %d", bucketID) + } + } + + // Buckets should be empty (limiter1 exhausted) + for bucketID := range buckets { + allow := limiters.allowN(bucketID, executionTime, 1) + require.False(t, allow, "should not allow when limiter1 tokens exhausted for bucket %d", bucketID) + } + + // Test 1: Multiple goroutines with immediate cancellation + { + concurrency := buckets + results := make([]bool, concurrency) + + // Context that's immediately cancelled (closed) + done := make(chan struct{}) + ctx := &testContext{ + done: done, + } + close(done) + + // Start concurrent waits + var wg sync.WaitGroup + for i := range concurrency { + wg.Add(1) + go func(i int64) { + defer wg.Done() + allow, _, err := limiters.waitNWithDetails(ctx, i, executionTime, 1) + require.False(t, allow, "should not acquire token due to immediate cancellation") + require.ErrorIs(t, err, context.Canceled, "should return context canceled error") + results[i] = allow + }(i) + } + wg.Wait() + + // All should fail because context is cancelled immediately + for i, result := range results { + require.False(t, result, "goroutine %d should not acquire token due to immediate cancellation", i) + } + } + + // Test 2: Multiple goroutines with delayed cancellation + { + keyFunc := func(input int64) int64 { + return input + } + + // Create two limiters with different limits + limit1 := NewLimit(1, 200*time.Millisecond) // 1 per 200ms (more restrictive) + limit2 := NewLimit(3, 200*time.Millisecond) // 3 per 200ms + limiter1 := NewLimiter(keyFunc, limit1) + limiter2 := NewLimiter(keyFunc, limit2) + limiters := Combine(limiter1, limiter2) + + // Exhaust tokens from the more restrictive limiter + for bucketID := range buckets { + for range limit1.count { + allow := limiters.allowN(bucketID, executionTime, 1) + require.True(t, allow, "should allow initial tokens for bucket %d", bucketID) + } + } + + // We're trying to emulate a wait call + // that starts waiting with a retry, but + // is cancelled in the meantime. + + const buckets int64 = 3 + results := make([]bool, buckets) + + var wg sync.WaitGroup + for bucketID := range buckets { + wg.Add(1) + + done := make(chan struct{}) + ctx := &testContext{ + done: done, + } + + go func(bucketID int64) { + defer wg.Done() + // Should start waiting, since the bucket is empty. + // retryAfter should be ~200ms + allow, _, err := limiters.waitNWithDetails(ctx, bucketID, executionTime, 1) + require.False(t, allow, "should not acquire token due to cancellation") + require.ErrorIs(t, err, context.Canceled, "should return context canceled error") + results[bucketID] = allow + }(bucketID) + + // Cancel context after goroutine has started. + // Originally, this had a delay, for a better test, + // but it is flaky on GitHub Actions, presumably + // because the runner is resource constrained. + // Delay works fine locally on Mac M2. + close(done) + } + wg.Wait() + + // All should fail because context is cancelled + for i, result := range results { + require.False(t, result, "goroutine %d should not acquire token due to cancellation", i) + } + } + }) + }) + + t.Run("MultipleLimiters", func(t *testing.T) { + t.Parallel() + t.Run("SameKeyer", func(t *testing.T) { + t.Parallel() + + keyFunc := func(input string) string { + return fmt.Sprintf("bucket-%s", input) + } + + // Create two limiters with different limits + limit1 := NewLimit(3, 100*time.Millisecond) // 3 per 100ms (more restrictive) + limit2 := NewLimit(5, 100*time.Millisecond) // 5 per 100ms + limiter1 := NewLimiter(keyFunc, limit1) + limiter2 := NewLimiter(keyFunc, limit2) + + // Create Limiters with both limiters + limiters := Combine(limiter1, limiter2) + + executionTime := ntime.Now() + + // Consume all tokens from the more restrictive limiter + for range limit1.count { + ok := limiters.allowN("test", executionTime, 1) + require.True(t, ok, "should allow initial tokens") + } + + // Should not allow immediately (limiter1 exhausted) + ok := limiters.allowN("test", executionTime, 1) + require.False(t, ok, "should not allow when limiter1 tokens exhausted") + + // Test 1: Wait for 1 token with no cancellation + { + ctx := &testContext{ + done: make(chan struct{}), // never closes + } + + allow, details, err := limiters.waitNWithDetails(ctx, "test", executionTime, 1) + require.True(t, allow, "should acquire 1 token after waiting") + require.NotZero(t, details, "should return details") + require.NoError(t, err, "should not return error") + + // We waited + executionTime = executionTime.Add(limit1.durationPerToken) + } + + // Should not allow again immediately after wait + ok = limiters.allowN("test", executionTime, 1) + require.False(t, ok, "should not allow again immediately after wait") + + // Test 2: Wait for 1 token with immediate cancellation + { + // Context that's immediately cancelled (closed) + done := make(chan struct{}) + ctx := &testContext{ + done: done, + } + close(done) + + allow, details, err := limiters.waitNWithDetails(ctx, "test", executionTime, 1) + require.False(t, allow, "should not acquire 1 token if context is cancelled immediately") + require.NotZero(t, details, "should return details even on cancellation") + require.ErrorIs(t, err, context.Canceled, "should return context canceled error") + } + }) + + t.Run("DifferentKeyers", func(t *testing.T) { + t.Parallel() + + // Create limiters with different key functions + keyFunc1 := func(input string) string { return input + "-1" } + keyFunc2 := func(input string) string { return input + "-2" } + + limit1 := NewLimit(2, 100*time.Millisecond) + limit2 := NewLimit(3, 100*time.Millisecond) + limiter1 := NewLimiter(keyFunc1, limit1) + limiter2 := NewLimiter(keyFunc2, limit2) + + limiters := Combine(limiter1, limiter2) + executionTime := ntime.Now() + + // Test that we can consume tokens from both keyers + // The Combine function treats all limiters as a single unit + // We should be able to consume up to the most restrictive limit (2) from each keyer + for i := range 2 { + ok := limiters.allowN("test", executionTime, 1) + require.True(t, ok, "should allow token %d from keyer1", i+1) + } + + // Should not allow more from keyer1 (limit1 exhausted) + ok := limiters.allowN("test", executionTime, 1) + require.False(t, ok, "should not allow when keyer1 limit exhausted") + + // Test: Wait for 1 token with no cancellation + { + ctx := &testContext{ + done: make(chan struct{}), // never closes + } + + allow, details, err := limiters.waitNWithDetails(ctx, "test", executionTime, 1) + require.True(t, allow, "should acquire 1 token after waiting") + require.NotZero(t, details, "should return details") + require.NoError(t, err, "should not return error") + + // We waited + executionTime = executionTime.Add(limit1.durationPerToken) + } + + // Should not allow again immediately after wait + ok = limiters.allowN("test", executionTime, 1) + require.False(t, ok, "should not allow again immediately after wait") + }) + }) +} + +func TestLimiters_WaitWithDebug(t *testing.T) { + t.Parallel() + + t.Run("SingleBucket", func(t *testing.T) { + t.Parallel() + keyFunc := func(input string) string { + return input + } + // Create two limiters with different limits + limit1 := NewLimit(2, 100*time.Millisecond) // 2 per 100ms (more restrictive) + limit2 := NewLimit(5, 100*time.Millisecond) // 5 per 100ms + limiter1 := NewLimiter(keyFunc, limit1) + limiter2 := NewLimiter(keyFunc, limit2) + limiters := Combine(limiter1, limiter2) + + executionTime := ntime.Now() + + // Consume all tokens from the more restrictive limiter + for range limit1.count { + ok := limiters.allowN("test", executionTime, 1) + require.True(t, ok, "should allow initial tokens") + } + + // Should not allow immediately (limiter1 exhausted) + ok := limiters.allowN("test", executionTime, 1) + require.False(t, ok, "should not allow when limiter1 tokens exhausted") + + // Test 1: Wait for 1 token with no cancellation + { + ctx := &testContext{ + done: make(chan struct{}), // never closes + } + + allow, debugs, err := limiters.waitNWithDebug(ctx, "test", executionTime, 1) + require.True(t, allow, "should acquire 1 token after waiting") + require.NotNil(t, debugs, "should return debugs") + require.NoError(t, err, "should not return error") + + // We waited + executionTime = executionTime.Add(limit1.durationPerToken) + } + + // Should not allow again immediately after wait + ok = limiters.allowN("test", executionTime, 1) + require.False(t, ok, "should not allow again immediately after wait") + + // Test 2: Wait for 1 token with immediate cancellation + { + // Context that's immediately cancelled (closed) + done := make(chan struct{}) + ctx := &testContext{ + done: done, + } + close(done) + + allow, debugs, err := limiters.waitNWithDebug(ctx, "test", executionTime, 1) + require.False(t, allow, "should not acquire 1 token if context is cancelled immediately") + require.NotNil(t, debugs, "should return debugs even on cancellation") + require.ErrorIs(t, err, context.Canceled, "should return context canceled error") + } + + // Test 3: Wait for multiple tokens (limit1.count) with no cancellation + { + // Wait for bucket to refill enough tokens + executionTime = executionTime.Add(limit1.durationPerToken * time.Duration(limit1.count)) + + ctx := &testContext{ + done: make(chan struct{}), // never closes + } + + allow, debugs, err := limiters.waitNWithDebug(ctx, "test", executionTime, limit1.count) + require.True(t, allow, "should acquire %d tokens after waiting", limit1.count) + require.NotNil(t, debugs, "should return debugs") + require.NoError(t, err, "should not return error") + + // We waited + executionTime = executionTime.Add(limit1.durationPerToken) + } + + // Test 4: Wait for multiple tokens with immediate cancellation + { + // Wait for bucket to refill enough tokens + executionTime = executionTime.Add(limit1.durationPerToken * time.Duration(limit1.count)) + + // Consume all tokens to exhaust the bucket again + for range limit1.count { + ok := limiters.allowN("test", executionTime, 1) + require.True(t, ok, "should allow tokens to exhaust bucket") + } + + // Context that's immediately cancelled (closed) + done := make(chan struct{}) + ctx := &testContext{ + done: done, + } + close(done) + + allow, debugs, err := limiters.waitNWithDebug(ctx, "test", executionTime, limit1.count) + require.False(t, allow, "should not acquire %d tokens if context is cancelled immediately", limit1.count) + require.NotNil(t, debugs, "should return debugs even on cancellation") + require.ErrorIs(t, err, context.Canceled, "should return context canceled error") + } + }) + + t.Run("MultipleBuckets", func(t *testing.T) { + t.Run("Serial", func(t *testing.T) { + t.Parallel() + keyFunc := func(input int) string { + return fmt.Sprintf("test-bucket-%d", input) + } + const buckets = 3 + // Create two limiters with different limits + limit1 := NewLimit(2, 100*time.Millisecond) // 2 per 100ms (more restrictive) + limit2 := NewLimit(4, 100*time.Millisecond) // 4 per 100ms + limiter1 := NewLimiter(keyFunc, limit1) + limiter2 := NewLimiter(keyFunc, limit2) + limiters := Combine(limiter1, limiter2) + executionTime := ntime.Now() + + // Exhaust tokens for all buckets from the more restrictive limiter + for bucketID := range buckets { + for range limit1.count { + allow := limiters.allowN(bucketID, executionTime, 1) + require.True(t, allow, "should allow initial tokens for bucket %d", bucketID) + } + } + + // Should not allow immediately for any bucket (limiter1 exhausted) + for bucketID := range buckets { + allow := limiters.allowN(bucketID, executionTime, 1) + require.False(t, allow, "should not allow when limiter1 tokens exhausted for bucket %d", bucketID) + } + + // Wait for a token for each bucket + for bucketID := range buckets { + ctx := &testContext{ + done: make(chan struct{}), // never closes + } + + allow, debugs, err := limiters.waitNWithDebug(ctx, bucketID, executionTime, 1) + require.True(t, allow, "should acquire token after waiting for bucket %d", bucketID) + require.NotNil(t, debugs, "should return debugs") + require.NoError(t, err, "should not return error") + } + + // We waited + executionTime = executionTime.Add(limit1.durationPerToken) + + // Buckets should be empty again + for bucketID := range buckets { + ok := limiters.allowN(bucketID, executionTime, 1) + require.False(t, ok, "should not allow again immediately after wait for bucket %d", bucketID) + } + }) + + t.Run("Concurrent", func(t *testing.T) { + // t.Parallel() + // GitHub Actions are pretty resource constrained, + // and this test is sensitive to timing. + keyFunc := func(input int64) string { + return fmt.Sprintf("test-bucket-%d", input) + } + const buckets int64 = 3 + // Create two limiters with different limits + limit1 := NewLimit(2, 100*time.Millisecond) // 2 per 100ms (more restrictive) + limit2 := NewLimit(4, 100*time.Millisecond) // 4 per 100ms + limiter1 := NewLimiter(keyFunc, limit1) + limiter2 := NewLimiter(keyFunc, limit2) + limiters := Combine(limiter1, limiter2) + executionTime := ntime.Now() + + // Exhaust tokens for all buckets from the more restrictive limiter + for bucketID := range buckets { + for range limit1.count { + allow := limiters.allowN(bucketID, executionTime, 1) + require.True(t, allow, "should allow initial tokens for bucket %d", bucketID) + } + } + + // Buckets should be empty (limiter1 exhausted) + for bucketID := range buckets { + allow := limiters.allowN(bucketID, executionTime, 1) + require.False(t, allow, "should not allow when limiter1 tokens exhausted for bucket %d", bucketID) + } + + // Test 1: Multiple goroutines with immediate cancellation + { + concurrency := buckets + results := make([]bool, concurrency) + + // Context that's immediately cancelled (closed) + done := make(chan struct{}) + ctx := &testContext{ + done: done, + } + close(done) + + // Start concurrent waits + var wg sync.WaitGroup + for i := range concurrency { + wg.Add(1) + go func(i int64) { + defer wg.Done() + allow, _, err := limiters.waitNWithDebug(ctx, i, executionTime, 1) + require.False(t, allow, "should not acquire token due to immediate cancellation") + require.ErrorIs(t, err, context.Canceled, "should return context canceled error") + results[i] = allow + }(i) + } + wg.Wait() + + // All should fail because context is cancelled immediately + for i, result := range results { + require.False(t, result, "goroutine %d should not acquire token due to immediate cancellation", i) + } + } + + // Test 2: Multiple goroutines with delayed cancellation + { + keyFunc := func(input int64) int64 { + return input + } + + limit1 := NewLimit(1, 200*time.Millisecond) + limit2 := NewLimit(3, 200*time.Millisecond) + limiter1 := NewLimiter(keyFunc, limit1) + limiter2 := NewLimiter(keyFunc, limit2) + limiters := Combine(limiter1, limiter2) + + // Exhaust tokens + for bucketID := range buckets { + // limit 1 being exhausted means the overall combined limiters are exhausted + for range limit1.count { + allow := limiters.allowN(bucketID, executionTime, 1) + require.True(t, allow, "should allow initial tokens for bucket %d", bucketID) + } + } + + // Confirm exhausted + for bucketID := range buckets { + allow := limiters.allowN(bucketID, executionTime, 1) + require.False(t, allow, "should not allow when limiter1 tokens exhausted for bucket %d", bucketID) + } + + // We're trying to emulate a wait call + // that starts waiting with a retry, but + // is cancelled in the meantime. + + const buckets int64 = 3 + results := make([]bool, buckets) + + var wg sync.WaitGroup + for bucketID := range buckets { + wg.Add(1) + + done := make(chan struct{}) + ctx := &testContext{ + done: done, + } + + go func(bucketID int64) { + defer wg.Done() + // Should start waiting, since the bucket is empty. + // retryAfter should be ~200ms + allow, _, err := limiters.waitNWithDebug(ctx, bucketID, executionTime, 1) + require.False(t, allow, "should not acquire token due to cancellation, bucketID: %d", bucketID) + require.ErrorIs(t, err, context.Canceled, "should return context canceled error") + results[bucketID] = allow + }(bucketID) + + // Cancel context after goroutine has started. + // Originally, this had a delay, for a better test, + // but it is flaky on GitHub Actions, presumably + // because the runner is resource constrained. + // Delay works fine locally on Mac M2. + close(done) + } + wg.Wait() + + // All should fail because context is cancelled + for i, result := range results { + require.False(t, result, "goroutine %d should not acquire token due to cancellation", i) + } + } + }) + }) +} diff --git a/syncmap.go b/syncmap.go index bf1ae3e..9f9c240 100644 --- a/syncmap.go +++ b/syncmap.go @@ -6,39 +6,7 @@ import ( "github.com/clipperhouse/ntime" ) -// syncMap is a typed wrapper around sync.Map for our specific use case -type syncMap[K comparable, V any] struct { - m sync.Map -} - -// loadOrStore returns the existing value for the key if present. -// Otherwise, it calls the factory function to create a new value, stores it, and returns it. -// This avoids creating the value unless it's actually needed. -func (sm *syncMap[K, V]) loadOrStore(key K, getter func() V) V { - if loaded, ok := sm.m.Load(key); ok { - return loaded.(V) - } - // Only create the value if we didn't find an existing one - value := getter() - actual, _ := sm.m.LoadOrStore(key, value) - return actual.(V) -} - -func (sm *syncMap[K, V]) count() int { - count := 0 - sm.m.Range(func(_, _ any) bool { - count++ - return true - }) - return count -} - -// delete removes a key from the map -func (sm *syncMap[K, V]) delete(key K) { - sm.m.Delete(key) -} - -// bucketMap is a specialized sync.Map for storing buckets to avoid allocations +// bucketMap is a specialized sync.Map for storing buckets type bucketMap[TKey comparable] struct { m sync.Map } @@ -47,13 +15,12 @@ type bucketMap[TKey comparable] struct { // It is a composite key to ensure that each bucket is unique for a given limit and user key. type bucketSpec[TKey comparable] struct { limit Limit - // userKey is the result of calling the user-defined Keyer + // userKey is the result of calling the user-defined KeyFunc userKey TKey } // loadOrStore returns the existing bucket for the key if present. // Otherwise, it creates a new bucket, stores it, and returns it. -// This is specialized to avoid a closure allocation for the getter. func (bm *bucketMap[TKey]) loadOrStore(userKey TKey, executionTime ntime.Time, limit Limit) *bucket { spec := bucketSpec[TKey]{ limit: limit, @@ -63,27 +30,12 @@ func (bm *bucketMap[TKey]) loadOrStore(userKey TKey, executionTime ntime.Time, l if loaded, ok := bm.m.Load(spec); ok { return loaded.(*bucket) } - // Only create the b if we didn't find an existing one + // Only create the bucket if we didn't find an existing one b := newBucket(executionTime, limit) actual, _ := bm.m.LoadOrStore(spec, &b) return actual.(*bucket) } -// loadOrGet returns the existing value for the key if present. -// Otherwise, it returns a new (temporary) value. -func (bm *bucketMap[TKey]) loadOrGet(userKey TKey, executionTime ntime.Time, limit Limit) *bucket { - spec := bucketSpec[TKey]{ - limit: limit, - userKey: userKey, - } - loaded, ok := bm.m.Load(spec) - if ok { - return loaded.(*bucket) - } - b := newBucket(executionTime, limit) - return &b -} - func (bm *bucketMap[TKey]) load(userKey TKey, limit Limit) (*bucket, bool) { spec := bucketSpec[TKey]{ limit: limit, diff --git a/syncmap_test.go b/syncmap_test.go index 51b37ff..5441a44 100644 --- a/syncmap_test.go +++ b/syncmap_test.go @@ -47,52 +47,6 @@ func TestSyncMap_ConcurrentAccess(t *testing.T) { // If we get here without panic or race conditions, the test passes } -func TestSyncMap_LoadOrStoreFunc(t *testing.T) { - t.Parallel() - - sm := &syncMap[string, int]{} - - factoryCalls := 0 - factory := func() int { - factoryCalls++ - return 42 - } - - // First call should call factory and store the value - result1 := sm.loadOrStore("key1", factory) - require.Equal(t, 42, result1, "should return factory value") - require.Equal(t, 1, factoryCalls, "factory should be called once") - - // Second call with same key should NOT call factory - result2 := sm.loadOrStore("key1", factory) - require.Equal(t, 42, result2, "should return existing value") - require.Equal(t, 1, factoryCalls, "factory should not be called again for existing key") - - // Third call with different key should call factory again - result3 := sm.loadOrStore("key2", factory) - require.Equal(t, 42, result3, "should return factory value for new key") - require.Equal(t, 2, factoryCalls, "factory should be called for new key") -} - -func TestSyncMap_Count(t *testing.T) { - t.Parallel() - - sm := &syncMap[string, int]{} - - for range 2 { - // should only be stored once despite multiple calls - for key := range 101 { - sm.loadOrStore(fmt.Sprint(key), func() int { - return key - }) - } - } - - expected := 101 - actual := sm.count() - require.Equal(t, expected, actual, "expected Count() to be accurate") -} - func TestBucketMap_LoadOrStore(t *testing.T) { t.Parallel() @@ -118,31 +72,6 @@ func TestBucketMap_LoadOrStore(t *testing.T) { require.False(t, bucket1 == bucket4, "should return different bucket for different key") } -func TestBucketMap_LoadOrGet(t *testing.T) { - t.Parallel() - - var bm bucketMap[string] - limit := NewLimit(100, time.Second) - executionTime := ntime.Now() - - // First call should return a temporary bucket (not stored) - bucket1 := bm.loadOrGet("key1", executionTime, limit) - require.NotNil(t, bucket1, "should return a bucket") - - // Second call should return another temporary bucket (different instance) - // Use slightly different execution time to ensure different bucket instances - bucket2 := bm.loadOrGet("key1", executionTime.Add(time.Nanosecond), limit) - require.NotNil(t, bucket2, "should return a bucket") - require.False(t, bucket1 == bucket2, "should return different temporary buckets when none stored") - - // Store a bucket first - storedBucket := bm.loadOrStore("key1", executionTime, limit) - - // Now loadOrGet should return the stored bucket - bucket3 := bm.loadOrGet("key1", executionTime, limit) - require.Equal(t, storedBucket, bucket3, "should return stored bucket when available") -} - func TestBucketMap_Load(t *testing.T) { t.Parallel() @@ -227,9 +156,6 @@ func TestBucketMap_ConcurrentAccess(t *testing.T) { // Test loadOrStore method bm.loadOrStore(key, executionTime, limit) - // Test loadOrGet method - bm.loadOrGet(key, executionTime, limit) - // Test load method bm.load(key, limit)