Merge pull request 'feat: 完善鉴权模块' (#142) from fix/auth into develop

Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/142
This commit was merged in pull request #142.
This commit is contained in:
2026-06-20 16:34:47 +08:00
8 changed files with 417 additions and 14 deletions

View File

@@ -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三级存储

View File

@@ -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。

View File

@@ -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

View File

@@ -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

View File

@@ -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 失败已间接验证
}

View File

@@ -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
}
// 写 RedisSET + 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先查 Redismiss 时查 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
}
// 回填 RedisSET + SADDTTL 使用保守默认值
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)
}

View File

@@ -1,6 +1,7 @@
// ============================================================
// HTTP API 客户端
// 职责:封装 REST API 请求auth、conversations 等)
// 内置 401 拦截 + 自动刷新 + 重试机制
// ============================================================
const API_BASE = "/api";
@@ -11,9 +12,69 @@ interface ApiResponse<T> {
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<boolean> | 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<boolean> {
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<boolean> {
if (!refreshPromise) {
refreshPromise = doRefresh().finally(() => {
refreshPromise = null;
});
}
return refreshPromise;
}
// ---- 核心请求函数 ----
async function request<T>(
path: string,
options: RequestInit = {}
options: RequestInit = {},
_retry = false
): Promise<ApiResponse<T>> {
const url = `${API_BASE}${path}`;
const headers: Record<string, string> = {
@@ -21,6 +82,14 @@ async function request<T>(
...(options.headers as Record<string, string>),
};
// 对非公开路径自动附加 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<T>(
const body = await res.json();
// 401 拦截:尝试刷新 token 后重试(仅重试一次)
if (res.status === 401 && !_retry && !isPublicPath(path) && authCallbacks) {
const refreshed = await refreshWithLock();
if (refreshed) {
return request<T>(path, options, true);
}
}
if (!res.ok) {
return {
error: { code: body.error || "UNKNOWN", message: body.message || "请求失败" },
@@ -39,7 +116,7 @@ async function request<T>(
}
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<ApiResponse<AuthResponse>> {
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<ApiResponse<AuthResponse>> {
@@ -96,11 +200,11 @@ export async function refreshToken(
export async function logout(
accessToken: string,
refreshToken: string
refreshTokenStr: string
): Promise<ApiResponse<{ message: string }>> {
return request<{ message: string }>("/auth/logout", {
method: "POST",
headers: authHeaders(accessToken),
body: JSON.stringify({ refresh_token: refreshToken }),
body: JSON.stringify({ refresh_token: refreshTokenStr }),
});
}

View File

@@ -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 () => {