package ratelimit import ( "sync" "time" "golang.org/x/time/rate" ) // 内存实现:自 middleware.keyedLimiter 与 service.loginLimiter 迁入(2026-07, // 状态外置改造),行为与迁入前一致。单实例进程内状态,重启即清零。 // 所有 map 均带 janitor 定期淘汰,保证不随攻击者构造的随机 key 无限增长。 const ( memSweepEvery = 5 * time.Minute memLimiterIdleTTL = 10 * time.Minute memLockerIdleTTL = 30 * time.Minute // 内存版给固定退避提示(redis 版有精确值)。 memRetryAfter = 60 * time.Second ) type memoryStore struct { counter *memCounter locker *memFailLocker } func newMemoryStore() *memoryStore { return &memoryStore{ counter: newMemCounter(), locker: newMemFailLocker(), } } func (s *memoryStore) NewLimiter(r rate.Limit, burst int) Limiter { return newMemLimiter(r, burst) } func (s *memoryStore) Counter() Counter { return s.counter } func (s *memoryStore) FailLocker() FailLocker { return s.locker } // ── Limiter:按 key 的令牌桶(原 keyedLimiter)────────────────────────────── type memLimiter struct { mu sync.Mutex entries map[string]*memBucket r rate.Limit burst int } type memBucket struct { lim *rate.Limiter lastSeen time.Time } func newMemLimiter(r rate.Limit, burst int) *memLimiter { l := &memLimiter{entries: map[string]*memBucket{}, r: r, burst: burst} go l.janitor() return l } func (l *memLimiter) Allow(key string) Result { l.mu.Lock() b := l.entries[key] if b == nil { b = &memBucket{lim: rate.NewLimiter(l.r, l.burst)} l.entries[key] = b } b.lastSeen = time.Now() l.mu.Unlock() return Result{Allowed: b.lim.Allow(), RetryAfter: memRetryAfter} } func (l *memLimiter) janitor() { t := time.NewTicker(memSweepEvery) defer t.Stop() for range t.C { l.sweep(memLimiterIdleTTL) } } // sweep 淘汰超过 ttl 未活动的 key。拆出便于测试。 func (l *memLimiter) sweep(ttl time.Duration) { now := time.Now() l.mu.Lock() defer l.mu.Unlock() for k, b := range l.entries { if now.Sub(b.lastSeen) > ttl { delete(l.entries, k) } } } // ── Counter:带 TTL 的累加计数(日配额用)──────────────────────────────────── type memCounter struct { mu sync.Mutex entries map[string]*memCount janitorOnce sync.Once } type memCount struct { val int64 expireAt time.Time } func newMemCounter() *memCounter { return &memCounter{entries: map[string]*memCount{}} } func (c *memCounter) Incr(key string, ttl time.Duration) (int64, error) { c.startJanitor() now := time.Now() c.mu.Lock() defer c.mu.Unlock() e := c.entries[key] if e == nil || now.After(e.expireAt) { e = &memCount{expireAt: now.Add(ttl)} c.entries[key] = e } e.val++ return e.val, nil } func (c *memCounter) startJanitor() { c.janitorOnce.Do(func() { go func() { t := time.NewTicker(memSweepEvery) defer t.Stop() for range t.C { now := time.Now() c.mu.Lock() for k, e := range c.entries { if now.After(e.expireAt) { delete(c.entries, k) } } c.mu.Unlock() } }() }) } // ── FailLocker:登录失败计数与锁(原 service.loginLimiter)─────────────────── type memFailLocker struct { mu sync.Mutex entries map[string]*memFailEntry janitorOnce sync.Once } type memFailEntry struct { failures int lockedTill time.Time lastSeen time.Time } func newMemFailLocker() *memFailLocker { return &memFailLocker{entries: map[string]*memFailEntry{}} } func (l *memFailLocker) Locked(key string) bool { l.mu.Lock() defer l.mu.Unlock() e := l.entries[key] if e == nil { return false } e.lastSeen = time.Now() return time.Now().Before(e.lockedTill) } func (l *memFailLocker) RecordFailure(key string, max int, lockFor time.Duration) { l.startJanitor() l.mu.Lock() defer l.mu.Unlock() e := l.entries[key] if e == nil { e = &memFailEntry{} l.entries[key] = e } e.lastSeen = time.Now() e.failures++ if max > 0 && e.failures >= max { e.lockedTill = time.Now().Add(lockFor) e.failures = 0 } } func (l *memFailLocker) Reset(key string) { l.mu.Lock() defer l.mu.Unlock() delete(l.entries, key) } // startJanitor 惰性启动后台清理:淘汰「未锁定且空闲超 TTL」的 entry。 func (l *memFailLocker) startJanitor() { l.janitorOnce.Do(func() { go func() { t := time.NewTicker(memSweepEvery) defer t.Stop() for range t.C { now := time.Now() l.mu.Lock() for k, e := range l.entries { if now.After(e.lockedTill) && now.Sub(e.lastSeen) > memLockerIdleTTL { delete(l.entries, k) } } l.mu.Unlock() } }() }) }