feat: 建立可靠通知、吊销与跨实例状态

This commit is contained in:
weijuesen
2026-08-14 01:41:14 +08:00
parent a6a996a518
commit a25bc7abc6
22 changed files with 1017 additions and 245 deletions
+16 -11
View File
@@ -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
}
+26 -89
View File
@@ -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))
}
+104
View File
@@ -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)
}
+152
View File
@@ -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 使用)