Files
mce/internal/middleware/ratelimit.go

150 lines
3.9 KiB
Go
Raw 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 middleware
import (
"fmt"
"sync"
"time"
)
// RateLimiter 双维度滑动窗口登录限流
// 维度一(账户):同邮箱 1 分钟内失败 3 次 → 锁定 1 分钟
// 维度二IP同 IP 1 分钟内失败 10 次 → 锁定 1 分钟
type RateLimiter struct {
mu sync.RWMutex
// accountFailures key = 邮箱(小写)
accountFailures map[string]*windowState
// ipFailures key = IP
ipFailures map[string]*windowState
}
// windowState 单个维度的限流状态
type windowState struct {
count int
windowStart time.Time
blockedUntil time.Time // 零值为未封锁
}
// RateLimitResult 限流检查结果
type RateLimitResult struct {
Blocked bool
RetryAfter int // 剩余封锁秒数
Message string
}
const (
accountWindow = 1 * time.Minute
accountMaxFails = 3
accountBlockDur = 1 * time.Minute
ipWindow = 1 * time.Minute
ipMaxFails = 10
ipBlockDur = 1 * time.Minute
cleanupInterval = 2 * time.Minute
maxEntries = 10000 // 单维度最大条目数(防内存耗尽)
)
// NewRateLimiter 创建限流器并启动后台清理
func NewRateLimiter() *RateLimiter {
rl := &RateLimiter{
accountFailures: make(map[string]*windowState),
ipFailures: make(map[string]*windowState),
}
go rl.cleanupLoop()
return rl
}
// AllowAccount 检查账户维度是否允许登录尝试
// 返回 (result, 登录失败时应调用的记录函数)
func (rl *RateLimiter) AllowAccount(email string) (RateLimitResult, func()) {
return rl.check(rl.accountFailures, email, accountWindow, accountBlockDur, accountMaxFails, true)
}
// AllowIP 检查 IP 维度是否允许登录尝试
func (rl *RateLimiter) AllowIP(ip string) (RateLimitResult, func()) {
return rl.check(rl.ipFailures, ip, ipWindow, ipBlockDur, ipMaxFails, false)
}
// check 核心检查逻辑
// lowKey: 是否需要脱敏日志true = 暗示账户存在,仅泄漏给已知该邮箱的人)
func (rl *RateLimiter) check(m map[string]*windowState, key string, window, blockDur time.Duration, maxFails int, lowKey bool) (RateLimitResult, func()) {
rl.mu.Lock()
defer rl.mu.Unlock()
now := time.Now()
state, exists := m[key]
if !exists {
state = &windowState{}
if len(m) < maxEntries {
m[key] = state
}
}
// 检查是否处于封锁期
if !state.blockedUntil.IsZero() && now.Before(state.blockedUntil) {
retry := int(state.blockedUntil.Sub(now).Seconds()) + 1
msg := "请求过于频繁,请稍后重试"
if lowKey {
msg = fmt.Sprintf("该账号登录尝试过于频繁,请 %d 秒后重试", retry)
}
return RateLimitResult{Blocked: true, RetryAfter: retry, Message: msg}, nil
}
// 窗口过期 → 重置
if now.Sub(state.windowStart) > window {
state.count = 0
state.windowStart = now
state.blockedUntil = time.Time{}
}
// 记录失败(由调用方在登录失败时调用)的闭包
recordFail := func() {
rl.mu.Lock()
defer rl.mu.Unlock()
s := m[key]
if s == nil {
return
}
if now.Sub(s.windowStart) > window {
s.count = 1
s.windowStart = now
return
}
s.count++
if s.count >= maxFails {
s.blockedUntil = now.Add(blockDur)
}
}
return RateLimitResult{Blocked: false}, recordFail
}
// Clear 登录成功后清除该 email 和 IP 的失败计数
func (rl *RateLimiter) Clear(email, ip string) {
rl.mu.Lock()
defer rl.mu.Unlock()
delete(rl.accountFailures, email)
delete(rl.ipFailures, ip)
}
// cleanupLoop 定期清理过期条目
func (rl *RateLimiter) cleanupLoop() {
ticker := time.NewTicker(cleanupInterval)
defer ticker.Stop()
for range ticker.C {
rl.mu.Lock()
now := time.Now()
clean := func(m map[string]*windowState) {
for k, v := range m {
// 封锁期已过且窗口已过期 → 删除
if (v.blockedUntil.IsZero() || now.After(v.blockedUntil)) &&
now.Sub(v.windowStart) > accountWindow+ipBlockDur {
delete(m, k)
}
}
}
clean(rl.accountFailures)
clean(rl.ipFailures)
rl.mu.Unlock()
}
}