diff --git a/backend/cmd/server/main.go b/backend/cmd/server/main.go index 3efbd6d..84bd165 100644 --- a/backend/cmd/server/main.go +++ b/backend/cmd/server/main.go @@ -86,9 +86,10 @@ func main() { } // L2: Redis(热数据分布式会话层) + var rdb *redis.Client var redisMgr *session.RedisManager if cfg.Storage.Redis.Enabled { - rdb := redis.NewClient(&redis.Options{ + rdb = redis.NewClient(&redis.Options{ Addr: cfg.Redis.Addr, Password: cfg.Redis.Password, DB: cfg.Redis.DB, @@ -102,9 +103,12 @@ func main() { time.Duration(cfg.Session.TTL)*time.Minute, cfg.Session.MaxHistory, ) + // 包装 userRepo 为带 Redis 缓存的版本(refresh token 二级缓存) + userRepo = store.NewCachedUserRepository(userRepo, rdb, time.Duration(cfg.Auth.RefreshTTL)*time.Minute) logger.Log.Infow("L2 Redis storage initialized", "addr", cfg.Redis.Addr, - "db", cfg.Redis.DB) + "db", cfg.Redis.DB, + "cached_user_repo", true) } // 初始化 Session Manager(三级存储) diff --git a/backend/internal/auth/jwt.go b/backend/internal/auth/jwt.go index 1da1ac0..01121ff 100644 --- a/backend/internal/auth/jwt.go +++ b/backend/internal/auth/jwt.go @@ -15,10 +15,17 @@ var ( ErrInvalidToken = errors.New("invalid or expired token") ) +// 令牌类型常量。 +const ( + TokenTypeAccess = "access" + TokenTypeRefresh = "refresh" +) + // Claims JWT 声明。 type Claims struct { - UserID string `json:"user_id"` - Username string `json:"username"` + UserID string `json:"user_id"` + Username string `json:"username"` + TokenType string `json:"token_type"` jwt.RegisteredClaims } @@ -45,8 +52,9 @@ func (tm *TokenManager) GeneratePair(userID, username string) (access, refresh s // access token accessClaims := &Claims{ - UserID: userID, - Username: username, + UserID: userID, + Username: username, + TokenType: TokenTypeAccess, RegisteredClaims: jwt.RegisteredClaims{ ExpiresAt: jwt.NewNumericDate(now.Add(tm.accessTTL)), IssuedAt: jwt.NewNumericDate(now), @@ -62,8 +70,9 @@ func (tm *TokenManager) GeneratePair(userID, username string) (access, refresh s // refresh token(含唯一 token_id 用于 DB 关联) tokenID := uuid.New().String() refreshClaims := &Claims{ - UserID: userID, - Username: username, + UserID: userID, + Username: username, + TokenType: TokenTypeRefresh, RegisteredClaims: jwt.RegisteredClaims{ ID: tokenID, ExpiresAt: jwt.NewNumericDate(now.Add(tm.refreshTTL)), @@ -78,12 +87,26 @@ func (tm *TokenManager) GeneratePair(userID, username string) (access, refresh s // ValidateAccess 校验 access token 并返回 Claims。 func (tm *TokenManager) ValidateAccess(tokenStr string) (*Claims, error) { - return tm.validate(tokenStr) + claims, err := tm.validate(tokenStr) + if err != nil { + return nil, err + } + if claims.TokenType != TokenTypeAccess { + return nil, ErrInvalidToken + } + return claims, nil } // ValidateRefresh 校验 refresh token 并返回 Claims。 func (tm *TokenManager) ValidateRefresh(tokenStr string) (*Claims, error) { - return tm.validate(tokenStr) + claims, err := tm.validate(tokenStr) + if err != nil { + return nil, err + } + if claims.TokenType != TokenTypeRefresh { + return nil, ErrInvalidToken + } + return claims, nil } // validate 解析并校验 JWT。 diff --git a/backend/internal/auth/jwt_test.go b/backend/internal/auth/jwt_test.go index 35b12fa..7178bae 100644 --- a/backend/internal/auth/jwt_test.go +++ b/backend/internal/auth/jwt_test.go @@ -103,6 +103,44 @@ func TestValidateRefresh_ExpiredToken(t *testing.T) { assert.ErrorIs(t, err, ErrInvalidToken) } +func TestValidateAccess_RejectsRefreshToken(t *testing.T) { + tm := NewTokenManager("test-secret-key", 15*time.Minute, 7*24*time.Hour) + + _, refresh, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + // refresh token 不能通过 access 校验 + _, err = tm.ValidateAccess(refresh) + assert.ErrorIs(t, err, ErrInvalidToken) +} + +func TestValidateRefresh_RejectsAccessToken(t *testing.T) { + tm := NewTokenManager("test-secret-key", 15*time.Minute, 7*24*time.Hour) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + // access token 不能通过 refresh 校验 + _, err = tm.ValidateRefresh(access) + assert.ErrorIs(t, err, ErrInvalidToken) +} + +func TestGeneratePair_TokenTypesAreCorrect(t *testing.T) { + tm := NewTokenManager("test-secret-key", 15*time.Minute, 7*24*time.Hour) + + access, refresh, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + // 通过 validate(不做类型检查)验证 token_type 字段 + accessClaims, err := tm.validate(access) + require.NoError(t, err) + assert.Equal(t, TokenTypeAccess, accessClaims.TokenType) + + refreshClaims, err := tm.validate(refresh) + require.NoError(t, err) + assert.Equal(t, TokenTypeRefresh, refreshClaims.TokenType) +} + func TestGeneratePair_TokenClaimsContainCorrectExpiry(t *testing.T) { accessTTL := 15 * time.Minute refreshTTL := 7 * 24 * time.Hour diff --git a/backend/internal/auth/service.go b/backend/internal/auth/service.go index 4f18447..0cd3f3f 100644 --- a/backend/internal/auth/service.go +++ b/backend/internal/auth/service.go @@ -166,6 +166,9 @@ func (s *authService) Refresh(ctx context.Context, req RefreshRequest) (*AuthRes userID, err := s.userRepo.FindRefreshToken(ctx, tokenHash) if err != nil { if errors.Is(err, store.ErrRefreshTokenNotFound) { + // JWT 校验已通过但 DB 中不存在 → token 已被 rotation 删除,属于复用行为 + // 吊销该用户全部 refresh token,强制所有设备重新登录 + _ = s.userRepo.DeleteUserRefreshTokens(ctx, claims.UserID) return nil, ErrRefreshTokenUsed } return nil, err diff --git a/backend/internal/auth/service_test.go b/backend/internal/auth/service_test.go index 9fc91a5..ab14554 100644 --- a/backend/internal/auth/service_test.go +++ b/backend/internal/auth/service_test.go @@ -188,3 +188,48 @@ func TestLogout_Success(t *testing.T) { }) assert.ErrorIs(t, err, auth.ErrRefreshTokenUsed) } + +// --- Refresh Token 复用检测 --- + +func TestRefresh_ReuseDetectedRevokesAllTokens(t *testing.T) { + svc, repo := newTestService(t) + ctx := context.Background() + + // 注册,获得令牌对 A + regResp, err := svc.Register(ctx, auth.RegisterRequest{ + Username: "eve", + Password: "password123", + }) + require.NoError(t, err) + tokenPairA_refresh := regResp.RefreshToken + + // 再次登录,获得令牌对 B + loginResp, err := svc.Login(ctx, auth.LoginRequest{ + Username: "eve", + Password: "password123", + }) + require.NoError(t, err) + tokenPairB_refresh := loginResp.RefreshToken + + // 用令牌对 A 的 refresh token 正常刷新 → 成功 + refreshResp, err := svc.Refresh(ctx, auth.RefreshRequest{ + RefreshToken: tokenPairA_refresh, + }) + require.NoError(t, err) + assert.NotEmpty(t, refreshResp.AccessToken) + + // 用令牌对 A 的旧 refresh token 再次刷新 → 复用检测,应失败 + _, err = svc.Refresh(ctx, auth.RefreshRequest{ + RefreshToken: tokenPairA_refresh, + }) + assert.ErrorIs(t, err, auth.ErrRefreshTokenUsed) + + // 令牌对 B 的 refresh token 也应被吊销(全量吊销) + _, err = svc.Refresh(ctx, auth.RefreshRequest{ + RefreshToken: tokenPairB_refresh, + }) + assert.ErrorIs(t, err, auth.ErrRefreshTokenUsed) + + // 确认 DB 中该用户已无 refresh token + _ = repo // repo 用于确认,但 MemUserRepository 无直接查询方法,通过 Refresh 失败已间接验证 +} diff --git a/backend/internal/store/cached_user.go b/backend/internal/store/cached_user.go new file mode 100644 index 0000000..dfc2e9e --- /dev/null +++ b/backend/internal/store/cached_user.go @@ -0,0 +1,168 @@ +package store + +import ( + "context" + "time" + + "github.com/redis/go-redis/v9" + + "github.com/hhs/camtalk/internal/logger" +) + +// Redis key 前缀。 +const ( + refreshTokenPrefix = "auth:refresh:" // auth:refresh:{token_hash} → user_id + userRefreshPrefix = "auth:user_refresh:" // auth:user_refresh:{user_id} → Set of token_hash +) + +// CachedUserRepository 装饰器,为 UserRepository 的 refresh token 操作增加 Redis 缓存。 +// 读路径:Redis miss → DB → 回填 Redis。 +// 写路径:同步双写 Redis + DB。 +// 删路径:同步双删 Redis + DB。 +// Redis 操作失败时降级到纯 DB,不阻断主流程。 +type CachedUserRepository struct { + inner UserRepository + rdb *redis.Client + backfillTTL time.Duration // DB 回填 Redis 时使用的默认 TTL +} + +// NewCachedUserRepository 创建带 Redis 缓存的 UserRepository 装饰器。 +// backfillTTL: 从 DB 回填 Redis 时使用的 TTL(因 DB 接口不返回 expiresAt)。 +func NewCachedUserRepository(inner UserRepository, rdb *redis.Client, backfillTTL time.Duration) *CachedUserRepository { + if backfillTTL <= 0 { + backfillTTL = 24 * time.Hour + } + return &CachedUserRepository{ + inner: inner, + rdb: rdb, + backfillTTL: backfillTTL, + } +} + +// refreshTokenKey 生成 refresh token 的 Redis key。 +func refreshTokenKey(tokenHash string) string { + return refreshTokenPrefix + tokenHash +} + +// userRefreshKey 生成用户 refresh token 集合的 Redis key。 +func userRefreshKey(userID string) string { + return userRefreshPrefix + userID +} + +// --- 委托方法(不做缓存) --- + +func (r *CachedUserRepository) Create(ctx context.Context, username, passwordHash string) (string, error) { + return r.inner.Create(ctx, username, passwordHash) +} + +func (r *CachedUserRepository) FindByUsername(ctx context.Context, username string) (*User, error) { + return r.inner.FindByUsername(ctx, username) +} + +func (r *CachedUserRepository) FindByID(ctx context.Context, id string) (*User, error) { + return r.inner.FindByID(ctx, id) +} + +// --- 缓存方法 --- + +// SaveRefreshToken Write-Through:先写 DB,再写 Redis。 +func (r *CachedUserRepository) SaveRefreshToken(ctx context.Context, userID, tokenHash string, expiresAt time.Time) error { + // 先写 DB + if err := r.inner.SaveRefreshToken(ctx, userID, tokenHash, expiresAt); err != nil { + return err + } + + // 写 Redis(SET + SADD),设置 TTL 为 token 剩余有效期 + ttl := time.Until(expiresAt) + if ttl <= 0 { + return nil + } + + key := refreshTokenKey(tokenHash) + pipe := r.rdb.Pipeline() + pipe.Set(ctx, key, userID, ttl) + pipe.SAdd(ctx, userRefreshKey(userID), tokenHash) + if _, err := pipe.Exec(ctx); err != nil { + logger.Log.Warnw("Redis cache write failed for refresh token", "error", err) + // 降级:DB 已写入成功,Redis 失败不影响正确性 + } + return nil +} + +// FindRefreshToken Read-Through:先查 Redis,miss 时查 DB 并回填。 +func (r *CachedUserRepository) FindRefreshToken(ctx context.Context, tokenHash string) (string, error) { + key := refreshTokenKey(tokenHash) + + // 查 Redis + userID, err := r.rdb.Get(ctx, key).Result() + if err == nil { + return userID, nil + } + // redis.Nil 表示 key 不存在,其他错误记录日志后降级到 DB + if err != redis.Nil { + logger.Log.Warnw("Redis cache read failed for refresh token", "error", err) + } + + // 降级到 DB + userID, err = r.inner.FindRefreshToken(ctx, tokenHash) + if err != nil { + return "", err + } + + // 回填 Redis(SET + SADD),TTL 使用保守默认值 + go func() { + bgCtx := context.Background() + pipe := r.rdb.Pipeline() + pipe.Set(bgCtx, key, userID, r.backfillTTL) + pipe.SAdd(bgCtx, userRefreshKey(userID), tokenHash) + _, _ = pipe.Exec(bgCtx) + }() + + return userID, nil +} + +// DeleteRefreshToken 双删:先删 DB,再删 Redis。 +func (r *CachedUserRepository) DeleteRefreshToken(ctx context.Context, tokenHash string) error { + // 先从 Redis 获取 user_id(用于从集合中移除) + userID, _ := r.rdb.Get(ctx, refreshTokenKey(tokenHash)).Result() + + // 删 DB + if err := r.inner.DeleteRefreshToken(ctx, tokenHash); err != nil { + return err + } + + // 删 Redis + key := refreshTokenKey(tokenHash) + pipe := r.rdb.Pipeline() + pipe.Del(ctx, key) + if userID != "" { + pipe.SRem(ctx, userRefreshKey(userID), tokenHash) + } + if _, err := pipe.Exec(ctx); err != nil { + logger.Log.Warnw("Redis cache delete failed for refresh token", "error", err) + } + return nil +} + +// DeleteUserRefreshTokens 批量清理:先从 Redis 获取集合,逐个删缓存,再删 DB。 +func (r *CachedUserRepository) DeleteUserRefreshTokens(ctx context.Context, userID string) error { + userKey := userRefreshKey(userID) + + // 从 Redis 获取该用户所有 token hash + hashes, _ := r.rdb.SMembers(ctx, userKey).Result() + + // 批量删除 Redis 缓存 + if len(hashes) > 0 { + keys := make([]string, 0, len(hashes)+1) + for _, h := range hashes { + keys = append(keys, refreshTokenKey(h)) + } + keys = append(keys, userKey) + if err := r.rdb.Del(ctx, keys...).Err(); err != nil { + logger.Log.Warnw("Redis cache batch delete failed for user refresh tokens", "error", err, "userID", userID) + } + } + + // 删 DB(无论 Redis 是否成功都执行) + return r.inner.DeleteUserRefreshTokens(ctx, userID) +} diff --git a/frontend/src/lib/api.ts b/frontend/src/lib/api.ts index ea11424..960b076 100644 --- a/frontend/src/lib/api.ts +++ b/frontend/src/lib/api.ts @@ -1,6 +1,7 @@ // ============================================================ // HTTP API 客户端 // 职责:封装 REST API 请求(auth、conversations 等) +// 内置 401 拦截 + 自动刷新 + 重试机制 // ============================================================ const API_BASE = "/api"; @@ -11,9 +12,69 @@ interface ApiResponse { status: number; } +// ---- 认证回调(由 AuthProvider 注入,避免循环依赖) ---- + +interface AuthCallbacks { + getAccessToken: () => string | null; + getRefreshToken: () => string | null; + onRefreshSuccess: (user: AuthUser, accessToken: string, refreshToken: string) => void; + onRefreshFailed: () => void; +} + +let authCallbacks: AuthCallbacks | null = null; +let refreshPromise: Promise | null = null; + +/** 由 AuthProvider 在初始化时调用,注入认证回调。 */ +export function setAuthCallbacks(callbacks: AuthCallbacks): void { + authCallbacks = callbacks; +} + +/** 不需要认证的公开路径。 */ +const PUBLIC_PATHS = new Set([ + "/auth/register", + "/auth/login", + "/auth/refresh", +]); + +function isPublicPath(path: string): boolean { + return PUBLIC_PATHS.has(path); +} + +/** 尝试用 refresh token 换取新的 access token。 */ +async function doRefresh(): Promise { + const rt = authCallbacks?.getRefreshToken(); + if (!rt) return false; + + const res = await refreshTokenDirect(rt); + if (res.data) { + authCallbacks?.onRefreshSuccess( + res.data.user, + res.data.access_token, + res.data.refresh_token + ); + return true; + } + + authCallbacks?.onRefreshFailed(); + return false; +} + +/** 带并发保护的刷新:多个 401 只触发一次 refresh。 */ +async function refreshWithLock(): Promise { + if (!refreshPromise) { + refreshPromise = doRefresh().finally(() => { + refreshPromise = null; + }); + } + return refreshPromise; +} + +// ---- 核心请求函数 ---- + async function request( path: string, - options: RequestInit = {} + options: RequestInit = {}, + _retry = false ): Promise> { const url = `${API_BASE}${path}`; const headers: Record = { @@ -21,6 +82,14 @@ async function request( ...(options.headers as Record), }; + // 对非公开路径自动附加 access token + if (!isPublicPath(path) && !headers["Authorization"]) { + const token = authCallbacks?.getAccessToken(); + if (token) { + headers["Authorization"] = `Bearer ${token}`; + } + } + try { const res = await fetch(url, { ...options, headers }); const status = res.status; @@ -31,6 +100,14 @@ async function request( const body = await res.json(); + // 401 拦截:尝试刷新 token 后重试(仅重试一次) + if (res.status === 401 && !_retry && !isPublicPath(path) && authCallbacks) { + const refreshed = await refreshWithLock(); + if (refreshed) { + return request(path, options, true); + } + } + if (!res.ok) { return { error: { code: body.error || "UNKNOWN", message: body.message || "请求失败" }, @@ -39,7 +116,7 @@ async function request( } return { data: body as T, status }; - } catch (err) { + } catch { return { error: { code: "NETWORK_ERROR", message: "网络连接失败,请检查网络" }, status: 0, @@ -85,6 +162,33 @@ export async function login( }); } +/** 内部用的 refresh 请求,不经过 401 拦截(避免递归)。 */ +async function refreshTokenDirect( + refresh_token: string +): Promise> { + const url = `${API_BASE}/auth/refresh`; + try { + const res = await fetch(url, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ refresh_token }), + }); + const body = await res.json(); + if (!res.ok) { + return { + error: { code: body.error || "UNKNOWN", message: body.message || "请求失败" }, + status: res.status, + }; + } + return { data: body as AuthResponse, status: res.status }; + } catch { + return { + error: { code: "NETWORK_ERROR", message: "网络连接失败" }, + status: 0, + }; + } +} + export async function refreshToken( refresh_token: string ): Promise> { @@ -96,11 +200,11 @@ export async function refreshToken( export async function logout( accessToken: string, - refreshToken: string + refreshTokenStr: string ): Promise> { return request<{ message: string }>("/auth/logout", { method: "POST", headers: authHeaders(accessToken), - body: JSON.stringify({ refresh_token: refreshToken }), + body: JSON.stringify({ refresh_token: refreshTokenStr }), }); } diff --git a/frontend/src/lib/auth.tsx b/frontend/src/lib/auth.tsx index b270f6d..733175c 100644 --- a/frontend/src/lib/auth.tsx +++ b/frontend/src/lib/auth.tsx @@ -15,6 +15,7 @@ import { } from "react"; import * as api from "./api"; import type { AuthUser } from "./api"; +import { setAuthCallbacks } from "./api"; import { clearAuth, loadAccessToken, @@ -108,6 +109,23 @@ export function AuthProvider({ children }: { children: ReactNode }) { [clearRefreshTimer, persistAuth] ); + // 注册 API 层认证回调(用于 401 拦截器) + useEffect(() => { + setAuthCallbacks({ + getAccessToken: () => loadAccessToken(), + getRefreshToken: () => loadRefreshToken(), + onRefreshSuccess: (u, at, rt) => { + persistAuth(u, at, rt); + scheduleRefresh(at); + }, + onRefreshFailed: () => { + clearAuth(); + setUser(null); + setAccessToken(null); + }, + }); + }, [persistAuth, scheduleRefresh]); + // 初始化:检查已有 token 并尝试刷新 useEffect(() => { const init = async () => {