feat: v2 版本 #206
172
backend/internal/ratelimit/bucket.go
Normal file
172
backend/internal/ratelimit/bucket.go
Normal file
@@ -0,0 +1,172 @@
|
||||
package ratelimit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/hhs/camtalk/internal/config"
|
||||
)
|
||||
|
||||
// TokenBucket 内存令牌桶,适用于单实例部署。
|
||||
type TokenBucket struct {
|
||||
capacity int // 桶容量
|
||||
rate float64 // 每秒填充令牌数
|
||||
tokens float64 // 当前令牌数
|
||||
lastRefill time.Time // 上次填充时间
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
// newTokenBucket 创建令牌桶。
|
||||
func newTokenBucket(capacity int, rate float64) *TokenBucket {
|
||||
return &TokenBucket{
|
||||
capacity: capacity,
|
||||
rate: rate,
|
||||
tokens: float64(capacity), // 初始满桶
|
||||
lastRefill: time.Now(),
|
||||
}
|
||||
}
|
||||
|
||||
// allow 尝试消耗一个令牌。
|
||||
func (b *TokenBucket) allow() (bool, time.Duration) {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
elapsed := now.Sub(b.lastRefill).Seconds()
|
||||
|
||||
// 补充令牌
|
||||
newTokens := elapsed * b.rate
|
||||
b.tokens = min(float64(b.capacity), b.tokens+newTokens)
|
||||
b.lastRefill = now
|
||||
|
||||
// 尝试消耗一个令牌
|
||||
if b.tokens >= 1 {
|
||||
b.tokens -= 1
|
||||
return true, 0
|
||||
}
|
||||
|
||||
// 计算需要等待的时间
|
||||
if b.rate == 0 {
|
||||
// rate=0 时永远无法补充令牌
|
||||
return false, 24 * time.Hour // 返回一个很大的值
|
||||
}
|
||||
retryAfter := time.Duration((1-b.tokens)/b.rate*1000) * time.Millisecond
|
||||
return false, retryAfter
|
||||
}
|
||||
|
||||
// MemoryLimiter 管理多个用户的令牌桶。
|
||||
type MemoryLimiter struct {
|
||||
buckets map[string]*TokenBucket
|
||||
config config.RateLimitConfig
|
||||
mu sync.RWMutex
|
||||
stopOnce sync.Once
|
||||
done chan struct{}
|
||||
}
|
||||
|
||||
// NewMemoryLimiter 创建内存限流器。
|
||||
func NewMemoryLimiter(cfg config.RateLimitConfig) *MemoryLimiter {
|
||||
limiter := &MemoryLimiter{
|
||||
buckets: make(map[string]*TokenBucket),
|
||||
config: cfg,
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
|
||||
// 启动后台清理 goroutine
|
||||
go limiter.cleanup()
|
||||
|
||||
return limiter
|
||||
}
|
||||
|
||||
// Allow 实现 Limiter 接口。
|
||||
func (l *MemoryLimiter) Allow(ctx context.Context, key string) (bool, time.Duration) {
|
||||
bucket := l.getOrCreateBucket(key)
|
||||
return bucket.allow()
|
||||
}
|
||||
|
||||
// Stop 实现 Limiter 接口。
|
||||
func (l *MemoryLimiter) Stop() {
|
||||
l.stopOnce.Do(func() {
|
||||
close(l.done)
|
||||
})
|
||||
}
|
||||
|
||||
// getOrCreateBucket 获取或创建令牌桶。
|
||||
func (l *MemoryLimiter) getOrCreateBucket(key string) *TokenBucket {
|
||||
// 先尝试读锁
|
||||
l.mu.RLock()
|
||||
bucket, exists := l.buckets[key]
|
||||
l.mu.RUnlock()
|
||||
|
||||
if exists {
|
||||
return bucket
|
||||
}
|
||||
|
||||
// 需要创建新桶,升级为写锁
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
// 双重检查(可能其他 goroutine 已创建)
|
||||
bucket, exists = l.buckets[key]
|
||||
if exists {
|
||||
return bucket
|
||||
}
|
||||
|
||||
// 根据 key 确定配置(简化版:假设 key 格式为 "userID:action")
|
||||
cfg := l.getBucketConfig(key)
|
||||
bucket = newTokenBucket(cfg.Capacity, cfg.Rate)
|
||||
l.buckets[key] = bucket
|
||||
|
||||
return bucket
|
||||
}
|
||||
|
||||
// getBucketConfig 根据 key 获取桶配置。
|
||||
func (l *MemoryLimiter) getBucketConfig(key string) config.BucketConfig {
|
||||
// 简化实现:从 key 后缀判断动作类型
|
||||
// 实际使用时调用方会传递正确的 key
|
||||
// 默认使用 query 配置
|
||||
return l.config.Query
|
||||
}
|
||||
|
||||
// cleanup 定期清理不活跃的桶。
|
||||
func (l *MemoryLimiter) cleanup() {
|
||||
ticker := time.NewTicker(10 * time.Minute)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
l.removeInactiveBuckets()
|
||||
case <-l.done:
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// removeInactiveBuckets 移除超过 10 分钟无活动的桶。
|
||||
func (l *MemoryLimiter) removeInactiveBuckets() {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
for key, bucket := range l.buckets {
|
||||
bucket.mu.Lock()
|
||||
inactive := now.Sub(bucket.lastRefill) > 10*time.Minute
|
||||
bucket.mu.Unlock()
|
||||
|
||||
if inactive {
|
||||
delete(l.buckets, key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// min 返回两个 float64 中的较小值。
|
||||
func min(a, b float64) float64 {
|
||||
if a < b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
// 编译期接口检查
|
||||
var _ Limiter = (*MemoryLimiter)(nil)
|
||||
203
backend/internal/ratelimit/bucket_test.go
Normal file
203
backend/internal/ratelimit/bucket_test.go
Normal file
@@ -0,0 +1,203 @@
|
||||
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()
|
||||
}
|
||||
17
backend/internal/ratelimit/limiter.go
Normal file
17
backend/internal/ratelimit/limiter.go
Normal file
@@ -0,0 +1,17 @@
|
||||
package ratelimit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Limiter 速率限制器接口。
|
||||
type Limiter interface {
|
||||
// Allow 判断 key 是否允许执行一次操作。
|
||||
// key 通常为 "userID:action" 格式。
|
||||
// 返回 (allowed, retryAfter)。retryAfter 表示需要等待的时间。
|
||||
Allow(ctx context.Context, key string) (bool, time.Duration)
|
||||
|
||||
// Stop 停止限流器,清理资源(如后台 goroutine)。
|
||||
Stop()
|
||||
}
|
||||
Reference in New Issue
Block a user