287 lines
6.0 KiB
Go
287 lines
6.0 KiB
Go
package auth
|
|
|
|
import (
|
|
"sync"
|
|
"time"
|
|
)
|
|
|
|
type lockPolicy struct {
|
|
window time.Duration
|
|
maxFails int
|
|
lockFor time.Duration
|
|
}
|
|
|
|
func policyFor(kind LockKind) lockPolicy {
|
|
switch kind {
|
|
case LockLoginEndpoint, LockTalkTarget:
|
|
return lockPolicy{window: time.Hour, maxFails: 50, lockFor: time.Hour}
|
|
default:
|
|
// LockLoginEndpointIP、LockTalkPair、LockAdminIP、LockRegisterIP
|
|
return lockPolicy{window: 5 * time.Minute, maxFails: 10, lockFor: 5 * time.Minute}
|
|
}
|
|
}
|
|
|
|
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
|
|
ops int
|
|
lastSweep time.Time
|
|
maxEntries int
|
|
}
|
|
|
|
// NewLoginLocks 创建默认锁定计数器。
|
|
func NewLoginLocks() *MemoryLocks {
|
|
return &MemoryLocks{
|
|
entries: make(map[string]*lockEntry),
|
|
now: time.Now,
|
|
maxEntries: defaultMaxLockEntries,
|
|
}
|
|
}
|
|
|
|
// SetClock 注入时钟(测试到期解除)。
|
|
func (l *MemoryLocks) SetClock(now func() time.Time) {
|
|
l.mu.Lock()
|
|
defer l.mu.Unlock()
|
|
if now == nil {
|
|
l.now = time.Now
|
|
return
|
|
}
|
|
l.now = now
|
|
}
|
|
|
|
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()
|
|
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)
|
|
}
|
|
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
|
|
}
|
|
|
|
// Fail 记录一次失败;若因此触发锁定,返回 locked=true。
|
|
func (l *MemoryLocks) Fail(key LockKey) (bool, time.Duration) {
|
|
l.mu.Lock()
|
|
defer l.mu.Unlock()
|
|
now := l.now()
|
|
k := lockMapKey(key)
|
|
e := l.entries[k]
|
|
if e == nil {
|
|
e = &lockEntry{}
|
|
l.entries[k] = e
|
|
}
|
|
if e.lockedUntil.After(now) {
|
|
l.maybeSweepLocked(now)
|
|
return true, e.lockedUntil.Sub(now)
|
|
}
|
|
if !e.lockedUntil.IsZero() {
|
|
e.lockedUntil = time.Time{}
|
|
e.fails = nil
|
|
}
|
|
pol := policyFor(key.Kind)
|
|
cutoff := now.Add(-pol.window)
|
|
kept := e.fails[:0]
|
|
for _, t := range e.fails {
|
|
if t.After(cutoff) {
|
|
kept = append(kept, t)
|
|
}
|
|
}
|
|
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
|
|
}
|
|
|
|
// ClearEndpoint 清除某端编号相关的登录锁定(两种都清)。
|
|
func (l *MemoryLocks) ClearEndpoint(endpointID string) {
|
|
l.mu.Lock()
|
|
defer l.mu.Unlock()
|
|
for k, e := range l.entries {
|
|
// kind|endpoint|peer|ip
|
|
parts := splitLockKey(k)
|
|
if len(parts) != 4 {
|
|
continue
|
|
}
|
|
kind, ep := LockKind(parts[0]), parts[1]
|
|
if ep != endpointID {
|
|
continue
|
|
}
|
|
if kind == LockLoginEndpointIP || kind == LockLoginEndpoint {
|
|
delete(l.entries, k)
|
|
_ = e
|
|
}
|
|
}
|
|
}
|
|
|
|
// ClearAllForEndpoint 清除该编号作为 EndpointID 或 PeerID 出现的全部锁定。
|
|
func (l *MemoryLocks) ClearAllForEndpoint(endpointID string) {
|
|
l.mu.Lock()
|
|
defer l.mu.Unlock()
|
|
if endpointID == "" {
|
|
return
|
|
}
|
|
for k := range l.entries {
|
|
parts := splitLockKey(k)
|
|
if len(parts) != 4 {
|
|
continue
|
|
}
|
|
if parts[1] == endpointID || parts[2] == endpointID {
|
|
delete(l.entries, k)
|
|
}
|
|
}
|
|
}
|
|
|
|
// Clear 清除精确键。
|
|
func (l *MemoryLocks) Clear(key LockKey) {
|
|
l.mu.Lock()
|
|
defer l.mu.Unlock()
|
|
delete(l.entries, lockMapKey(key))
|
|
}
|
|
|
|
func splitLockKey(k string) []string {
|
|
out := make([]string, 0, 4)
|
|
start := 0
|
|
for i := 0; i < len(k); i++ {
|
|
if k[i] == '|' {
|
|
out = append(out, k[start:i])
|
|
start = i + 1
|
|
}
|
|
}
|
|
out = append(out, k[start:])
|
|
return out
|
|
}
|
|
|
|
var _ LoginLocks = (*MemoryLocks)(nil)
|