package ratelimit import ( "context" "sync" "testing" "time" "github.com/hhs/camtalk/internal/config" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) func TestTokenBucket_Allow_FirstRequest(t *testing.T) { bucket := newTokenBucket(5, 0.2) allowed, retryAfter := bucket.allow() assert.True(t, allowed) assert.Equal(t, time.Duration(0), retryAfter) } func TestTokenBucket_Allow_ConsumeUntilEmpty(t *testing.T) { bucket := newTokenBucket(3, 0.2) // 连续消耗 3 个令牌 for i := 0; i < 3; i++ { allowed, _ := bucket.allow() assert.True(t, allowed, "request %d should be allowed", i+1) } // 第 4 个请求应被拒绝 allowed, retryAfter := bucket.allow() assert.False(t, allowed) assert.Greater(t, retryAfter, time.Duration(0)) } func TestTokenBucket_Allow_RetryAfterCorrect(t *testing.T) { bucket := newTokenBucket(1, 1.0) // 每秒 1 个令牌 // 消耗唯一的令牌 allowed, _ := bucket.allow() require.True(t, allowed) // 立即再次请求应被拒绝 allowed, retryAfter := bucket.allow() assert.False(t, allowed) // retryAfter 应约为 1 秒(允许一定误差) assert.InDelta(t, 1000, retryAfter.Milliseconds(), 100) } func TestTokenBucket_Allow_RefillAfterWait(t *testing.T) { bucket := newTokenBucket(2, 10.0) // 每秒 10 个令牌(每 100ms 一个) // 消耗 2 个令牌 bucket.allow() bucket.allow() // 等待 150ms,应补充至少 1 个令牌 time.Sleep(150 * time.Millisecond) allowed, _ := bucket.allow() assert.True(t, allowed) } func TestTokenBucket_Allow_CapacityLimit(t *testing.T) { bucket := newTokenBucket(3, 1.0) // 等待足够长时间让桶"溢出" time.Sleep(100 * time.Millisecond) // 但最多只能消耗 capacity 个令牌 for i := 0; i < 3; i++ { allowed, _ := bucket.allow() assert.True(t, allowed, "request %d should be allowed", i+1) } // 第 4 个应被拒绝 allowed, _ := bucket.allow() assert.False(t, allowed) } func TestTokenBucket_Allow_ConcurrentSafe(t *testing.T) { bucket := newTokenBucket(100, 10.0) var wg sync.WaitGroup successCount := 0 var mu sync.Mutex // 100 个并发请求 for i := 0; i < 100; i++ { wg.Add(1) go func() { defer wg.Done() allowed, _ := bucket.allow() if allowed { mu.Lock() successCount++ mu.Unlock() } }() } wg.Wait() // 应该正好 100 个成功(桶容量为 100) assert.Equal(t, 100, successCount) } func TestTokenBucket_Allow_ZeroCapacity(t *testing.T) { bucket := newTokenBucket(0, 1.0) allowed, retryAfter := bucket.allow() assert.False(t, allowed) assert.Greater(t, retryAfter, time.Duration(0)) } func TestTokenBucket_Allow_ZeroRate(t *testing.T) { bucket := newTokenBucket(1, 0.0) // 第一个通过 allowed, _ := bucket.allow() assert.True(t, allowed) // 第二个被拒绝,且 retryAfter 应为无限大(实际上会很大) allowed, retryAfter := bucket.allow() assert.False(t, allowed) // rate=0 时,retryAfter 理论上无限大,实际会是一个很大的值 assert.Greater(t, retryAfter, 1*time.Hour) } func TestMemoryLimiter_Allow_DifferentKeys(t *testing.T) { cfg := config.RateLimitConfig{ Enabled: true, Query: config.BucketConfig{Capacity: 2, Rate: 1.0}, } limiter := NewMemoryLimiter(cfg) defer limiter.Stop() ctx := context.Background() // user1 消耗 2 个令牌 allowed, _ := limiter.Allow(ctx, "user1:query") assert.True(t, allowed) allowed, _ = limiter.Allow(ctx, "user1:query") assert.True(t, allowed) // user1 第 3 个被拒绝 allowed, _ = limiter.Allow(ctx, "user1:query") assert.False(t, allowed) // user2 应该不受影响 allowed, _ = limiter.Allow(ctx, "user2:query") assert.True(t, allowed) } func TestMemoryLimiter_Cleanup(t *testing.T) { cfg := config.RateLimitConfig{ Enabled: true, Query: config.BucketConfig{Capacity: 1, Rate: 1.0}, } limiter := NewMemoryLimiter(cfg) defer limiter.Stop() ctx := context.Background() // 创建一个桶 limiter.Allow(ctx, "user1:query") // 验证桶已创建 limiter.mu.RLock() initialCount := len(limiter.buckets) limiter.mu.RUnlock() assert.Equal(t, 1, initialCount) // 手动触发清理(模拟 10 分钟后) limiter.mu.Lock() for _, bucket := range limiter.buckets { bucket.mu.Lock() bucket.lastRefill = time.Now().Add(-11 * time.Minute) bucket.mu.Unlock() } limiter.mu.Unlock() limiter.removeInactiveBuckets() // 验证桶已被清理 limiter.mu.RLock() finalCount := len(limiter.buckets) limiter.mu.RUnlock() assert.Equal(t, 0, finalCount) } func TestMemoryLimiter_Stop(t *testing.T) { cfg := config.RateLimitConfig{ Enabled: true, Query: config.BucketConfig{Capacity: 1, Rate: 1.0}, } limiter := NewMemoryLimiter(cfg) // 多次调用 Stop 不应 panic limiter.Stop() limiter.Stop() }