Files
wangjia 3c687e5e0b feat(backend): 限流/登录失败锁状态外置 Redis(todo #2)
新包 internal/ratelimit:内存/redis(GCRA+Lua) 双实现 + 出错逐调用降级内存
(fail-open 到内存不 fail-closed)。REDIS_ADDR 空=内存模式,行为与既往一致;
配置后跨重启保状态、支持多实例。miniredis 全覆盖测试,零真实外部依赖。

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-05 10:42:23 +08:00

213 lines
4.7 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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()
}
}()
})
}