Files
NixMsg/internal/auth/locks.go

156 lines
3.4 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}
}
}
type lockEntry struct {
fails []time.Time
lockedUntil time.Time
}
// MemoryLocks 是内存锁定计数器(重启清零)。
type MemoryLocks struct {
mu sync.Mutex
entries map[string]*lockEntry
now func() time.Time
}
// NewLoginLocks 创建默认锁定计数器。
func NewLoginLocks() *MemoryLocks {
return &MemoryLocks{
entries: make(map[string]*lockEntry),
now: time.Now,
}
}
// 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
}
// 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)]
if e == nil {
return false, 0
}
if e.lockedUntil.After(now) {
return true, e.lockedUntil.Sub(now)
}
// 到期自动解除:清空锁定与窗口内失败(保留结构以便后续 Fail)。
if !e.lockedUntil.IsZero() && !e.lockedUntil.After(now) {
e.lockedUntil = time.Time{}
e.fails = nil
}
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) {
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)
if len(e.fails) >= pol.maxFails {
e.lockedUntil = now.Add(pol.lockFor)
e.fails = nil
return true, pol.lockFor
}
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
}
}
}
// 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)