- Limiter 接口定义(Allow + Stop 方法) - TokenBucket 实现(容量、填充速率、并发安全) - MemoryLimiter 管理多用户令牌桶 - 后台 goroutine 定期清理不活跃桶(10 分钟) - 完整单元测试覆盖(11 个测试用例,全部通过) - 边界情况处理(rate=0、capacity=0、并发安全)
173 lines
3.7 KiB
Go
173 lines
3.7 KiB
Go
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)
|