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)