fix: 锁定计数表过期清理并限制总量

This commit is contained in:
Nixevol
2026-09-30 16:21:51 +08:00
parent 8c20b76df5
commit 05a3741fde
3 changed files with 190 additions and 8 deletions
+121 -8
View File
@@ -21,23 +21,34 @@ func policyFor(kind LockKind) lockPolicy {
}
}
const (
defaultMaxLockEntries = 65536
lockSweepEveryOps = 1024
lockSweepInterval = time.Minute
)
type lockEntry struct {
fails []time.Time
lockedUntil time.Time
lastFail time.Time
}
// MemoryLocks 是内存锁定计数器(重启清零)。
type MemoryLocks struct {
mu sync.Mutex
entries map[string]*lockEntry
now func() time.Time
mu sync.Mutex
entries map[string]*lockEntry
now func() time.Time
ops int
lastSweep time.Time
maxEntries int
}
// NewLoginLocks 创建默认锁定计数器。
func NewLoginLocks() *MemoryLocks {
return &MemoryLocks{
entries: make(map[string]*lockEntry),
now: time.Now,
entries: make(map[string]*lockEntry),
now: time.Now,
maxEntries: defaultMaxLockEntries,
}
}
@@ -56,23 +67,121 @@ func lockMapKey(key LockKey) string {
return string(key.Kind) + "|" + key.EndpointID + "|" + key.PeerID + "|" + key.IP
}
func (l *MemoryLocks) maxCap() int {
if l.maxEntries <= 0 {
return defaultMaxLockEntries
}
return l.maxEntries
}
func (e *lockEntry) stale(now time.Time, window time.Duration) bool {
if e == nil {
return true
}
if e.lockedUntil.After(now) {
return false
}
if e.lastFail.IsZero() {
return true
}
return !e.lastFail.After(now.Add(-window))
}
func (l *MemoryLocks) maybeSweepLocked(now time.Time) {
max := l.maxCap()
if l.lastSweep.IsZero() {
l.lastSweep = now
}
l.ops++
if l.ops < lockSweepEveryOps && now.Sub(l.lastSweep) < lockSweepInterval && len(l.entries) <= max {
return
}
l.ops = 0
l.lastSweep = now
l.sweepExpiredLocked(now)
l.enforceCapLocked(now)
}
func (l *MemoryLocks) sweepExpiredLocked(now time.Time) {
for k, e := range l.entries {
window := policyFor("").window
parts := splitLockKey(k)
if len(parts) == 4 {
window = policyFor(LockKind(parts[0])).window
}
if e.stale(now, window) {
delete(l.entries, k)
}
}
}
func (l *MemoryLocks) enforceCapLocked(now time.Time) {
max := l.maxCap()
for len(l.entries) > max {
var (
victim string
victimLocked bool
victimTime time.Time
found bool
)
for k, e := range l.entries {
locked := e.lockedUntil.After(now)
t := e.lastFail
if locked {
t = e.lockedUntil
}
better := !found
if found {
if victimLocked && !locked {
better = true
} else if victimLocked == locked && (t.Before(victimTime) || (t.Equal(victimTime) && k < victim)) {
better = true
}
}
if better {
found = true
victim, victimLocked, victimTime = k, locked, t
}
}
if !found {
return
}
delete(l.entries, victim)
}
}
func (l *MemoryLocks) entryCount() int {
l.mu.Lock()
defer l.mu.Unlock()
return len(l.entries)
}
// Check 若当前已锁定返回 locked=true 与剩余时间。
func (l *MemoryLocks) Check(key LockKey) (bool, time.Duration) {
l.mu.Lock()
defer l.mu.Unlock()
now := l.now()
e := l.entries[lockMapKey(key)]
k := lockMapKey(key)
e := l.entries[k]
if e == nil {
l.maybeSweepLocked(now)
return false, 0
}
if e.lockedUntil.After(now) {
l.maybeSweepLocked(now)
return true, e.lockedUntil.Sub(now)
}
// 到期自动解除:清空锁定与窗口内失败(保留结构以便后续 Fail)。
if !e.lockedUntil.IsZero() && !e.lockedUntil.After(now) {
if e.stale(now, policyFor(key.Kind).window) {
delete(l.entries, k)
l.maybeSweepLocked(now)
return false, 0
}
// 到期自动解除:清空锁定与窗口内失败(保留仍在窗口内的失败计数)。
if !e.lockedUntil.IsZero() {
e.lockedUntil = time.Time{}
e.fails = nil
}
l.maybeSweepLocked(now)
return false, 0
}
@@ -88,6 +197,7 @@ func (l *MemoryLocks) Fail(key LockKey) (bool, time.Duration) {
l.entries[k] = e
}
if e.lockedUntil.After(now) {
l.maybeSweepLocked(now)
return true, e.lockedUntil.Sub(now)
}
if !e.lockedUntil.IsZero() {
@@ -103,11 +213,14 @@ func (l *MemoryLocks) Fail(key LockKey) (bool, time.Duration) {
}
}
e.fails = append(kept, now)
e.lastFail = now
if len(e.fails) >= pol.maxFails {
e.lockedUntil = now.Add(pol.lockFor)
e.fails = nil
l.maybeSweepLocked(now)
return true, pol.lockFor
}
l.maybeSweepLocked(now)
return false, 0
}
+60
View File
@@ -0,0 +1,60 @@
package auth
import (
"fmt"
"testing"
"time"
)
func TestLoginLocksSweepExpiredKeepsLocked(t *testing.T) {
locks := NewLoginLocks()
now := time.Date(2026, 9, 30, 12, 0, 0, 0, time.UTC)
locks.SetClock(func() time.Time { return now })
for i := 0; i < 10000; i++ {
if locked, _ := locks.Fail(LockKey{Kind: LockRegisterIP, IP: fmt.Sprintf("2001:db8::%d", i)}); locked {
t.Fatalf("unexpected lock at i=%d", i)
}
}
if n := locks.entryCount(); n < 10000 {
t.Fatalf("want 10000 entries before sweep, got %d", n)
}
now = now.Add(4 * time.Minute)
keep := LockKey{Kind: LockRegisterIP, IP: "keep-locked"}
for i := 0; i < 10; i++ {
locks.Fail(keep)
}
if locked, _ := locks.Check(keep); !locked {
t.Fatal("keep-locked should be locked")
}
now = now.Add(2 * time.Minute) // 10k 已过 5 分钟窗口;keep 仍锁定至 +5min
if locked, _ := locks.Check(LockKey{Kind: LockRegisterIP, IP: "trigger-sweep"}); locked {
t.Fatal("trigger must not lock")
}
if n := locks.entryCount(); n != 1 {
t.Fatalf("want only locked entry after sweep, got %d", n)
}
if locked, _ := locks.Check(keep); !locked {
t.Fatal("locked entry must remain")
}
}
func TestLoginLocksCapDoesNotGrow(t *testing.T) {
locks := NewLoginLocks()
locks.maxEntries = 64
now := time.Date(2026, 9, 30, 15, 0, 0, 0, time.UTC)
locks.SetClock(func() time.Time { return now })
for i := 0; i < 200; i++ {
locks.Fail(LockKey{Kind: LockRegisterIP, IP: fmt.Sprintf("ip-%d", i)})
now = now.Add(time.Millisecond)
if n := locks.entryCount(); n > 64 {
t.Fatalf("entries=%d exceeded cap at i=%d", n, i)
}
}
if n := locks.entryCount(); n != 64 {
t.Fatalf("entries=%d want 64", n)
}
}