feat: 建立可靠通知、吊销与跨实例状态
This commit is contained in:
@@ -20,16 +20,16 @@ type JWTClaims struct {
|
||||
|
||||
// 白名单路径,无需鉴权
|
||||
var whitelist = map[string]bool{
|
||||
"/api/v1/health": true,
|
||||
"/health": true,
|
||||
"/api/health": true,
|
||||
"/api/v1/auth/login": true,
|
||||
"/auth/login": true,
|
||||
"/api/v1/auth/register": true,
|
||||
"/auth/register": true,
|
||||
"/api/v1/auth/refresh": true,
|
||||
"/auth/refresh": true,
|
||||
"/api/v1/video/clips/internal": true,
|
||||
"/api/v1/health": true,
|
||||
"/health": true,
|
||||
"/api/health": true,
|
||||
"/api/v1/auth/login": true,
|
||||
"/auth/login": true,
|
||||
"/api/v1/auth/register": true,
|
||||
"/auth/register": true,
|
||||
"/api/v1/auth/refresh": true,
|
||||
"/auth/refresh": true,
|
||||
"/api/v1/video/clips/internal": true,
|
||||
"/api/v1/video/recordings/internal/end": true,
|
||||
}
|
||||
|
||||
@@ -83,7 +83,12 @@ func Auth(cfg *config.Config) gin.HandlerFunc {
|
||||
}
|
||||
|
||||
// 校验令牌是否已被登出吊销
|
||||
if IsRevoked(claims) {
|
||||
revoked, err := IsRevoked(claims)
|
||||
if err != nil {
|
||||
c.AbortWithStatusJSON(http.StatusServiceUnavailable, gin.H{"error": "认证状态服务不可用"})
|
||||
return
|
||||
}
|
||||
if revoked {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "令牌已注销,请重新登录"})
|
||||
return
|
||||
}
|
||||
|
||||
@@ -1,110 +1,41 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// loginAttempt 登录失败计数(按 IP + 用户名维度)
|
||||
type loginAttempt struct {
|
||||
failures int
|
||||
lockUntil time.Time
|
||||
lastFail time.Time
|
||||
}
|
||||
|
||||
type loginLimiter struct {
|
||||
mu sync.Mutex
|
||||
seen map[string]*loginAttempt
|
||||
}
|
||||
|
||||
const (
|
||||
maxFailures = 5 // 连续失败 5 次后锁定
|
||||
lockDuration = 15 * time.Minute
|
||||
failureWindow = 10 * time.Minute // 失败计数窗口
|
||||
cleanupInterval = 5 * time.Minute
|
||||
maxFailures = 5 // 连续失败 5 次后锁定
|
||||
lockDuration = 15 * time.Minute
|
||||
failureWindow = 10 * time.Minute // 失败计数窗口
|
||||
)
|
||||
|
||||
var defaultLoginLimiter = newLoginLimiter()
|
||||
|
||||
func newLoginLimiter() *loginLimiter {
|
||||
l := &loginLimiter{seen: make(map[string]*loginAttempt)}
|
||||
go l.cleanupLoop()
|
||||
return l
|
||||
}
|
||||
|
||||
func (l *loginLimiter) cleanupLoop() {
|
||||
t := time.NewTicker(cleanupInterval)
|
||||
defer t.Stop()
|
||||
for range t.C {
|
||||
l.mu.Lock()
|
||||
now := time.Now()
|
||||
for k, v := range l.seen {
|
||||
if now.After(v.lockUntil) && now.Sub(v.lastFail) > failureWindow {
|
||||
delete(l.seen, k)
|
||||
}
|
||||
}
|
||||
l.mu.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
// key = ip + "|" + username(小写)
|
||||
func limiterKey(c *gin.Context, username string) string {
|
||||
return c.ClientIP() + "|" + strings.ToLower(strings.TrimSpace(username))
|
||||
}
|
||||
|
||||
// checkLock 返回是否被锁定及剩余锁定时间
|
||||
func (l *loginLimiter) checkLock(key string) (bool, time.Duration) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
a, ok := l.seen[key]
|
||||
if !ok {
|
||||
return false, 0
|
||||
}
|
||||
if time.Now().Before(a.lockUntil) {
|
||||
return true, time.Until(a.lockUntil)
|
||||
}
|
||||
return false, 0
|
||||
}
|
||||
|
||||
// recordFailure 记录一次失败,达到阈值则锁定
|
||||
func (l *loginLimiter) recordFailure(key string) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
a, ok := l.seen[key]
|
||||
if !ok {
|
||||
a = &loginAttempt{}
|
||||
l.seen[key] = a
|
||||
}
|
||||
now := time.Now()
|
||||
// 窗口外重置
|
||||
if now.Sub(a.lastFail) > failureWindow {
|
||||
a.failures = 0
|
||||
}
|
||||
a.failures++
|
||||
a.lastFail = now
|
||||
if a.failures >= maxFailures {
|
||||
a.lockUntil = now.Add(lockDuration)
|
||||
}
|
||||
}
|
||||
|
||||
// recordSuccess 登录成功后清空计数
|
||||
func (l *loginLimiter) recordSuccess(key string) {
|
||||
l.mu.Lock()
|
||||
delete(l.seen, key)
|
||||
l.mu.Unlock()
|
||||
}
|
||||
|
||||
// CheckLoginLock 检查是否被锁定,被锁定则写 429 并返回 true(在 handler 解析 body 后调用)
|
||||
func CheckLoginLock(c *gin.Context, username string) bool {
|
||||
if authState == nil {
|
||||
c.AbortWithStatusJSON(http.StatusServiceUnavailable, gin.H{"error": stateUnavailable("执行登录限流").Error()})
|
||||
return true
|
||||
}
|
||||
key := limiterKey(c, username)
|
||||
if locked, remain := defaultLoginLimiter.checkLock(key); locked {
|
||||
locked, remain, err := authState.CheckLoginLock(context.Background(), key)
|
||||
if err != nil {
|
||||
c.AbortWithStatusJSON(http.StatusServiceUnavailable, gin.H{"error": "登录状态服务不可用"})
|
||||
return true
|
||||
}
|
||||
if locked {
|
||||
c.AbortWithStatusJSON(http.StatusTooManyRequests, gin.H{
|
||||
"error": "登录尝试过多,已锁定,请稍后再试",
|
||||
"retry": int(remain.Minutes()) + 1,
|
||||
"error": "登录尝试过多,已锁定,请稍后再试",
|
||||
"retry": int(remain.Minutes()) + 1,
|
||||
})
|
||||
return true
|
||||
}
|
||||
@@ -112,11 +43,17 @@ func CheckLoginLock(c *gin.Context, username string) bool {
|
||||
}
|
||||
|
||||
// RecordLoginFail 记录登录失败
|
||||
func RecordLoginFail(c *gin.Context, username string) {
|
||||
defaultLoginLimiter.recordFailure(limiterKey(c, username))
|
||||
func RecordLoginFail(c *gin.Context, username string) error {
|
||||
if authState == nil {
|
||||
return stateUnavailable("记录登录失败")
|
||||
}
|
||||
return authState.RecordLoginFailure(context.Background(), limiterKey(c, username))
|
||||
}
|
||||
|
||||
// RecordLoginSuccess 登录成功后清空计数
|
||||
func RecordLoginSuccess(c *gin.Context, username string) {
|
||||
defaultLoginLimiter.recordSuccess(limiterKey(c, username))
|
||||
func RecordLoginSuccess(c *gin.Context, username string) error {
|
||||
if authState == nil {
|
||||
return stateUnavailable("清空登录失败计数")
|
||||
}
|
||||
return authState.RecordLoginSuccess(context.Background(), limiterKey(c, username))
|
||||
}
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
const statePrefix = "silk:auth:"
|
||||
|
||||
// StateStore 跨实例认证状态存储。
|
||||
type StateStore interface {
|
||||
RevokeToken(ctx context.Context, id string, exp time.Time) error
|
||||
IsTokenRevoked(ctx context.Context, id string) (bool, error)
|
||||
CheckLoginLock(ctx context.Context, key string) (bool, time.Duration, error)
|
||||
RecordLoginFailure(ctx context.Context, key string) error
|
||||
RecordLoginSuccess(ctx context.Context, key string) error
|
||||
}
|
||||
|
||||
var (
|
||||
authState StateStore
|
||||
appEnv string
|
||||
)
|
||||
|
||||
// InitState 设置认证状态存储;store 为 nil 时认证相关接口保守失败。
|
||||
func InitState(store StateStore, env string) {
|
||||
authState = store
|
||||
appEnv = env
|
||||
}
|
||||
|
||||
// RedisState Redis 实现。
|
||||
type RedisState struct {
|
||||
rdb *redis.Client
|
||||
prefix string
|
||||
}
|
||||
|
||||
// NewRedisState 创建 Redis 状态存储。
|
||||
func NewRedisState(rdb *redis.Client) *RedisState {
|
||||
return &RedisState{rdb: rdb, prefix: statePrefix}
|
||||
}
|
||||
|
||||
func (s *RedisState) RevokeToken(ctx context.Context, id string, exp time.Time) error {
|
||||
ttl := time.Until(exp)
|
||||
if ttl <= 0 {
|
||||
return nil
|
||||
}
|
||||
return s.rdb.Set(ctx, s.prefix+"revoked:"+id, "1", ttl).Err()
|
||||
}
|
||||
|
||||
func (s *RedisState) IsTokenRevoked(ctx context.Context, id string) (bool, error) {
|
||||
count, err := s.rdb.Exists(ctx, s.prefix+"revoked:"+id).Result()
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
func (s *RedisState) CheckLoginLock(ctx context.Context, key string) (bool, time.Duration, error) {
|
||||
lockKey := s.prefix + "login-lock:" + key
|
||||
if _, err := s.rdb.Get(ctx, lockKey).Result(); err == redis.Nil {
|
||||
return false, 0, nil
|
||||
} else if err != nil {
|
||||
return false, 0, err
|
||||
}
|
||||
ttl, err := s.rdb.TTL(ctx, lockKey).Result()
|
||||
if err != nil {
|
||||
return false, 0, err
|
||||
}
|
||||
return true, ttl, nil
|
||||
}
|
||||
|
||||
var loginFailureScript = redis.NewScript(`
|
||||
local count = redis.call('INCR', KEYS[1])
|
||||
redis.call('EXPIRE', KEYS[1], ARGV[1])
|
||||
if tonumber(count) >= tonumber(ARGV[2]) then
|
||||
redis.call('SET', KEYS[2], '1', 'PX', ARGV[3])
|
||||
end
|
||||
return count
|
||||
`)
|
||||
|
||||
func (s *RedisState) RecordLoginFailure(ctx context.Context, key string) error {
|
||||
return loginFailureScript.Run(ctx, s.rdb,
|
||||
[]string{s.prefix + "login-failures:" + key, s.prefix + "login-lock:" + key},
|
||||
int(failureWindow.Seconds()), maxFailures, int(lockDuration.Milliseconds()),
|
||||
).Err()
|
||||
}
|
||||
|
||||
func (s *RedisState) RecordLoginSuccess(ctx context.Context, key string) error {
|
||||
pipe := s.rdb.Pipeline()
|
||||
pipe.Del(ctx, s.prefix+"login-failures:"+key, s.prefix+"login-lock:"+key)
|
||||
_, err := pipe.Exec(ctx)
|
||||
return err
|
||||
}
|
||||
|
||||
func stateUnavailable(operation string) error {
|
||||
msg := "Redis 状态服务不可用,无法" + operation
|
||||
if appEnv == "production" {
|
||||
slog.Error(msg)
|
||||
}
|
||||
return fmt.Errorf("%s", msg)
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type memoryState struct {
|
||||
mu sync.Mutex
|
||||
revoked map[string]time.Time
|
||||
failures map[string]int
|
||||
locked map[string]time.Time
|
||||
}
|
||||
|
||||
func newMemoryState() *memoryState {
|
||||
return &memoryState{
|
||||
revoked: map[string]time.Time{},
|
||||
failures: map[string]int{},
|
||||
locked: map[string]time.Time{},
|
||||
}
|
||||
}
|
||||
|
||||
func (m *memoryState) RevokeToken(_ context.Context, id string, exp time.Time) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.revoked[id] = exp
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *memoryState) IsTokenRevoked(_ context.Context, id string) (bool, error) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
exp, ok := m.revoked[id]
|
||||
if !ok {
|
||||
return false, nil
|
||||
}
|
||||
if time.Now().After(exp) {
|
||||
delete(m.revoked, id)
|
||||
return false, nil
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (m *memoryState) CheckLoginLock(_ context.Context, key string) (bool, time.Duration, error) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
until, ok := m.locked[key]
|
||||
if !ok {
|
||||
return false, 0, nil
|
||||
}
|
||||
if time.Now().After(until) {
|
||||
delete(m.locked, key)
|
||||
return false, 0, nil
|
||||
}
|
||||
return true, time.Until(until), nil
|
||||
}
|
||||
|
||||
func (m *memoryState) RecordLoginFailure(_ context.Context, key string) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.failures[key]++
|
||||
if m.failures[key] >= maxFailures {
|
||||
m.locked[key] = time.Now().Add(lockDuration)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *memoryState) RecordLoginSuccess(_ context.Context, key string) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
delete(m.failures, key)
|
||||
delete(m.locked, key)
|
||||
return nil
|
||||
}
|
||||
|
||||
func withMemoryState(t *testing.T) *memoryState {
|
||||
t.Helper()
|
||||
store := newMemoryState()
|
||||
InitState(store, "test")
|
||||
t.Cleanup(func() { InitState(nil, "test") })
|
||||
return store
|
||||
}
|
||||
|
||||
func TestRevokeTokenIsCrossInstanceState(t *testing.T) {
|
||||
withMemoryState(t)
|
||||
claims := &JWTClaims{}
|
||||
claims.ID = "token-1"
|
||||
if err := RevokeToken(claims, "token", time.Now().Add(time.Hour)); err != nil {
|
||||
t.Fatalf("RevokeToken failed: %v", err)
|
||||
}
|
||||
revoked, err := IsRevoked(claims)
|
||||
if err != nil {
|
||||
t.Fatalf("IsRevoked failed: %v", err)
|
||||
}
|
||||
if !revoked {
|
||||
t.Fatal("token should be revoked")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRevokeTokenUnavailableFailsClosed(t *testing.T) {
|
||||
InitState(nil, "production")
|
||||
t.Cleanup(func() { InitState(nil, "test") })
|
||||
claims := &JWTClaims{}
|
||||
claims.ID = "x"
|
||||
if err := RevokeToken(claims, "token", time.Now().Add(time.Hour)); err == nil {
|
||||
t.Fatal("state unavailable should fail closed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginLockStateSemantics(t *testing.T) {
|
||||
store := withMemoryState(t)
|
||||
for i := 0; i < maxFailures; i++ {
|
||||
if err := store.RecordLoginFailure(context.Background(), "1.1.1.1|user"); err != nil {
|
||||
t.Fatalf("record failure failed: %v", err)
|
||||
}
|
||||
}
|
||||
locked, _, err := store.CheckLoginLock(context.Background(), "1.1.1.1|user")
|
||||
if err != nil {
|
||||
t.Fatalf("check lock failed: %v", err)
|
||||
}
|
||||
if !locked {
|
||||
t.Fatal("expected login lock after max failures")
|
||||
}
|
||||
if err := store.RecordLoginSuccess(context.Background(), "1.1.1.1|user"); err != nil {
|
||||
t.Fatalf("record success failed: %v", err)
|
||||
}
|
||||
locked, _, _ = store.CheckLoginLock(context.Background(), "1.1.1.1|user")
|
||||
if locked {
|
||||
t.Fatal("login lock should be cleared after success")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckLoginLockUnavailableReturns503(t *testing.T) {
|
||||
InitState(nil, "test")
|
||||
t.Cleanup(func() { InitState(nil, "test") })
|
||||
gin.SetMode(gin.TestMode)
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/auth/login", nil)
|
||||
if !CheckLoginLock(c, "user") {
|
||||
t.Fatal("state unavailable should abort login")
|
||||
}
|
||||
if rec.Code != http.StatusServiceUnavailable {
|
||||
t.Fatalf("status = %d, want 503", rec.Code)
|
||||
}
|
||||
}
|
||||
@@ -1,45 +1,26 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
)
|
||||
|
||||
// tokenBlacklist 登出令牌黑名单(内存版,进程重启后失效,令牌自然过期兜底)
|
||||
type tokenBlacklist struct {
|
||||
mu sync.RWMutex
|
||||
revoked map[string]time.Time // tokenID(jti) -> 过期时间
|
||||
}
|
||||
|
||||
var defaultBlacklist = &tokenBlacklist{revoked: make(map[string]time.Time)}
|
||||
|
||||
// RevokeToken 将令牌加入黑名单(按 jti,若无 jti 则按 subject+签发时间)
|
||||
func RevokeToken(claims *JWTClaims, tokenStr string, exp time.Time) {
|
||||
id := tokenIdentifier(claims)
|
||||
defaultBlacklist.mu.Lock()
|
||||
defaultBlacklist.revoked[id] = exp
|
||||
defaultBlacklist.mu.Unlock()
|
||||
func RevokeToken(claims *JWTClaims, tokenStr string, exp time.Time) error {
|
||||
if authState == nil {
|
||||
return stateUnavailable("吊销令牌")
|
||||
}
|
||||
return authState.RevokeToken(context.Background(), tokenIdentifier(claims), exp)
|
||||
}
|
||||
|
||||
// IsRevoked 判断令牌是否已被吊销
|
||||
func IsRevoked(claims *JWTClaims) bool {
|
||||
id := tokenIdentifier(claims)
|
||||
defaultBlacklist.mu.RLock()
|
||||
exp, ok := defaultBlacklist.revoked[id]
|
||||
defaultBlacklist.mu.RUnlock()
|
||||
if !ok {
|
||||
return false
|
||||
func IsRevoked(claims *JWTClaims) (bool, error) {
|
||||
if authState == nil {
|
||||
return false, stateUnavailable("校验令牌吊销状态")
|
||||
}
|
||||
// 已过期的黑名单项自动清理
|
||||
if time.Now().After(exp) {
|
||||
defaultBlacklist.mu.Lock()
|
||||
delete(defaultBlacklist.revoked, id)
|
||||
defaultBlacklist.mu.Unlock()
|
||||
return false
|
||||
}
|
||||
return true
|
||||
return authState.IsTokenRevoked(context.Background(), tokenIdentifier(claims))
|
||||
}
|
||||
|
||||
// ExtractClaims 从 token 字符串解析 claims(供 logout handler 使用)
|
||||
|
||||
Reference in New Issue
Block a user