diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 5b3292a..40dd7e0 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -1367,3 +1367,12 @@ issue #3 未关闭,`feat/fix-3-downlink-deadlock` 未合入 `main`。下面是 - 原因:新库无码行或空码时 `constantTimeEqual("", "")` 为真,只开开关即可裸注册。 - 备选方案:仅拦管理 PUT、不拦已处于「开启+空码」的旧库(否决,缺少纵深防御)。 - 影响:原先「先开开关再设码」的两步会 400;须先设码或一次提交开启与码。 + +### 复审修复 U-04 + +1. **锁定计数表过期清理与总量上限** + - 原条款:DEVELOPMENT 第 5 节锁定计数;issue #42。 + - 实际做法:`internal/auth/locks.go` 的 Fail/Check 顺手删除已过期且最近失败在窗口外的条目;每 1024 次或每分钟全表扫描;默认上限 65536,超出时优先淘汰最旧的非锁定条目。生产接线使用 `auth.NewLoginLocks()`。 + - 原因:注册安全码错误与错误 API 令牌按 IP 建条目,轮换地址会使 map 只增不减。 + - 备选方案:一并改 `internal/admin/memlock.go`(否决,本波只改 locks.go;admin 测试用内存锁若仍独立注入需后续对齐)。未改对话密码锁键(U-03)。 + - 影响:过期未锁定条目会被回收;极端并发失败时最早的非锁定计数可能被挤出。 diff --git a/internal/auth/locks.go b/internal/auth/locks.go index 2c97074..b2a4dd9 100644 --- a/internal/auth/locks.go +++ b/internal/auth/locks.go @@ -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 } diff --git a/internal/auth/locks_test.go b/internal/auth/locks_test.go new file mode 100644 index 0000000..7be4ce9 --- /dev/null +++ b/internal/auth/locks_test.go @@ -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) + } +}