3c687e5e0b
新包 internal/ratelimit:内存/redis(GCRA+Lua) 双实现 + 出错逐调用降级内存 (fail-open 到内存不 fail-closed)。REDIS_ADDR 空=内存模式,行为与既往一致; 配置后跨重启保状态、支持多实例。miniredis 全覆盖测试,零真实外部依赖。 Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
213 lines
4.7 KiB
Go
213 lines
4.7 KiB
Go
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()
|
||
}
|
||
}()
|
||
})
|
||
}
|