fix: 完善鉴权模块,修复 CI/CD 中 dubious ownership 错误 #143
@@ -10,6 +10,7 @@ jobs:
|
|||||||
steps:
|
steps:
|
||||||
- name: Deploy
|
- name: Deploy
|
||||||
run: |
|
run: |
|
||||||
|
git config --global --add safe.directory /root/camtalk
|
||||||
if [ -d /root/camtalk/.git ]; then
|
if [ -d /root/camtalk/.git ]; then
|
||||||
cd /root/camtalk
|
cd /root/camtalk
|
||||||
git fetch origin ${GITHUB_REF_NAME} --depth=1
|
git fetch origin ${GITHUB_REF_NAME} --depth=1
|
||||||
|
|||||||
@@ -86,9 +86,10 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// L2: Redis(热数据分布式会话层)
|
// L2: Redis(热数据分布式会话层)
|
||||||
|
var rdb *redis.Client
|
||||||
var redisMgr *session.RedisManager
|
var redisMgr *session.RedisManager
|
||||||
if cfg.Storage.Redis.Enabled {
|
if cfg.Storage.Redis.Enabled {
|
||||||
rdb := redis.NewClient(&redis.Options{
|
rdb = redis.NewClient(&redis.Options{
|
||||||
Addr: cfg.Redis.Addr,
|
Addr: cfg.Redis.Addr,
|
||||||
Password: cfg.Redis.Password,
|
Password: cfg.Redis.Password,
|
||||||
DB: cfg.Redis.DB,
|
DB: cfg.Redis.DB,
|
||||||
@@ -102,9 +103,12 @@ func main() {
|
|||||||
time.Duration(cfg.Session.TTL)*time.Minute,
|
time.Duration(cfg.Session.TTL)*time.Minute,
|
||||||
cfg.Session.MaxHistory,
|
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",
|
logger.Log.Infow("L2 Redis storage initialized",
|
||||||
"addr", cfg.Redis.Addr,
|
"addr", cfg.Redis.Addr,
|
||||||
"db", cfg.Redis.DB)
|
"db", cfg.Redis.DB,
|
||||||
|
"cached_user_repo", true)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 初始化 Session Manager(三级存储)
|
// 初始化 Session Manager(三级存储)
|
||||||
|
|||||||
@@ -15,10 +15,17 @@ var (
|
|||||||
ErrInvalidToken = errors.New("invalid or expired token")
|
ErrInvalidToken = errors.New("invalid or expired token")
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// 令牌类型常量。
|
||||||
|
const (
|
||||||
|
TokenTypeAccess = "access"
|
||||||
|
TokenTypeRefresh = "refresh"
|
||||||
|
)
|
||||||
|
|
||||||
// Claims JWT 声明。
|
// Claims JWT 声明。
|
||||||
type Claims struct {
|
type Claims struct {
|
||||||
UserID string `json:"user_id"`
|
UserID string `json:"user_id"`
|
||||||
Username string `json:"username"`
|
Username string `json:"username"`
|
||||||
|
TokenType string `json:"token_type"`
|
||||||
jwt.RegisteredClaims
|
jwt.RegisteredClaims
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -45,8 +52,9 @@ func (tm *TokenManager) GeneratePair(userID, username string) (access, refresh s
|
|||||||
|
|
||||||
// access token
|
// access token
|
||||||
accessClaims := &Claims{
|
accessClaims := &Claims{
|
||||||
UserID: userID,
|
UserID: userID,
|
||||||
Username: username,
|
Username: username,
|
||||||
|
TokenType: TokenTypeAccess,
|
||||||
RegisteredClaims: jwt.RegisteredClaims{
|
RegisteredClaims: jwt.RegisteredClaims{
|
||||||
ExpiresAt: jwt.NewNumericDate(now.Add(tm.accessTTL)),
|
ExpiresAt: jwt.NewNumericDate(now.Add(tm.accessTTL)),
|
||||||
IssuedAt: jwt.NewNumericDate(now),
|
IssuedAt: jwt.NewNumericDate(now),
|
||||||
@@ -62,8 +70,9 @@ func (tm *TokenManager) GeneratePair(userID, username string) (access, refresh s
|
|||||||
// refresh token(含唯一 token_id 用于 DB 关联)
|
// refresh token(含唯一 token_id 用于 DB 关联)
|
||||||
tokenID := uuid.New().String()
|
tokenID := uuid.New().String()
|
||||||
refreshClaims := &Claims{
|
refreshClaims := &Claims{
|
||||||
UserID: userID,
|
UserID: userID,
|
||||||
Username: username,
|
Username: username,
|
||||||
|
TokenType: TokenTypeRefresh,
|
||||||
RegisteredClaims: jwt.RegisteredClaims{
|
RegisteredClaims: jwt.RegisteredClaims{
|
||||||
ID: tokenID,
|
ID: tokenID,
|
||||||
ExpiresAt: jwt.NewNumericDate(now.Add(tm.refreshTTL)),
|
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。
|
// ValidateAccess 校验 access token 并返回 Claims。
|
||||||
func (tm *TokenManager) ValidateAccess(tokenStr string) (*Claims, error) {
|
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。
|
// ValidateRefresh 校验 refresh token 并返回 Claims。
|
||||||
func (tm *TokenManager) ValidateRefresh(tokenStr string) (*Claims, error) {
|
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。
|
// validate 解析并校验 JWT。
|
||||||
|
|||||||
@@ -103,6 +103,44 @@ func TestValidateRefresh_ExpiredToken(t *testing.T) {
|
|||||||
assert.ErrorIs(t, err, ErrInvalidToken)
|
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) {
|
func TestGeneratePair_TokenClaimsContainCorrectExpiry(t *testing.T) {
|
||||||
accessTTL := 15 * time.Minute
|
accessTTL := 15 * time.Minute
|
||||||
refreshTTL := 7 * 24 * time.Hour
|
refreshTTL := 7 * 24 * time.Hour
|
||||||
|
|||||||
@@ -166,6 +166,9 @@ func (s *authService) Refresh(ctx context.Context, req RefreshRequest) (*AuthRes
|
|||||||
userID, err := s.userRepo.FindRefreshToken(ctx, tokenHash)
|
userID, err := s.userRepo.FindRefreshToken(ctx, tokenHash)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, store.ErrRefreshTokenNotFound) {
|
if errors.Is(err, store.ErrRefreshTokenNotFound) {
|
||||||
|
// JWT 校验已通过但 DB 中不存在 → token 已被 rotation 删除,属于复用行为
|
||||||
|
// 吊销该用户全部 refresh token,强制所有设备重新登录
|
||||||
|
_ = s.userRepo.DeleteUserRefreshTokens(ctx, claims.UserID)
|
||||||
return nil, ErrRefreshTokenUsed
|
return nil, ErrRefreshTokenUsed
|
||||||
}
|
}
|
||||||
return nil, err
|
return nil, err
|
||||||
|
|||||||
@@ -188,3 +188,48 @@ func TestLogout_Success(t *testing.T) {
|
|||||||
})
|
})
|
||||||
assert.ErrorIs(t, err, auth.ErrRefreshTokenUsed)
|
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 失败已间接验证
|
||||||
|
}
|
||||||
|
|||||||
168
backend/internal/store/cached_user.go
Normal file
168
backend/internal/store/cached_user.go
Normal 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
|
||||||
|
}
|
||||||
|
|
||||||
|
// 写 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)
|
||||||
|
}
|
||||||
@@ -1,6 +1,7 @@
|
|||||||
// ============================================================
|
// ============================================================
|
||||||
// HTTP API 客户端
|
// HTTP API 客户端
|
||||||
// 职责:封装 REST API 请求(auth、conversations 等)
|
// 职责:封装 REST API 请求(auth、conversations 等)
|
||||||
|
// 内置 401 拦截 + 自动刷新 + 重试机制
|
||||||
// ============================================================
|
// ============================================================
|
||||||
|
|
||||||
const API_BASE = "/api";
|
const API_BASE = "/api";
|
||||||
@@ -11,9 +12,69 @@ interface ApiResponse<T> {
|
|||||||
status: number;
|
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>(
|
async function request<T>(
|
||||||
path: string,
|
path: string,
|
||||||
options: RequestInit = {}
|
options: RequestInit = {},
|
||||||
|
_retry = false
|
||||||
): Promise<ApiResponse<T>> {
|
): Promise<ApiResponse<T>> {
|
||||||
const url = `${API_BASE}${path}`;
|
const url = `${API_BASE}${path}`;
|
||||||
const headers: Record<string, string> = {
|
const headers: Record<string, string> = {
|
||||||
@@ -21,6 +82,14 @@ async function request<T>(
|
|||||||
...(options.headers as Record<string, string>),
|
...(options.headers as Record<string, string>),
|
||||||
};
|
};
|
||||||
|
|
||||||
|
// 对非公开路径自动附加 access token
|
||||||
|
if (!isPublicPath(path) && !headers["Authorization"]) {
|
||||||
|
const token = authCallbacks?.getAccessToken();
|
||||||
|
if (token) {
|
||||||
|
headers["Authorization"] = `Bearer ${token}`;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
try {
|
try {
|
||||||
const res = await fetch(url, { ...options, headers });
|
const res = await fetch(url, { ...options, headers });
|
||||||
const status = res.status;
|
const status = res.status;
|
||||||
@@ -31,6 +100,14 @@ async function request<T>(
|
|||||||
|
|
||||||
const body = await res.json();
|
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) {
|
if (!res.ok) {
|
||||||
return {
|
return {
|
||||||
error: { code: body.error || "UNKNOWN", message: body.message || "请求失败" },
|
error: { code: body.error || "UNKNOWN", message: body.message || "请求失败" },
|
||||||
@@ -39,7 +116,7 @@ async function request<T>(
|
|||||||
}
|
}
|
||||||
|
|
||||||
return { data: body as T, status };
|
return { data: body as T, status };
|
||||||
} catch (err) {
|
} catch {
|
||||||
return {
|
return {
|
||||||
error: { code: "NETWORK_ERROR", message: "网络连接失败,请检查网络" },
|
error: { code: "NETWORK_ERROR", message: "网络连接失败,请检查网络" },
|
||||||
status: 0,
|
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(
|
export async function refreshToken(
|
||||||
refresh_token: string
|
refresh_token: string
|
||||||
): Promise<ApiResponse<AuthResponse>> {
|
): Promise<ApiResponse<AuthResponse>> {
|
||||||
@@ -96,11 +200,11 @@ export async function refreshToken(
|
|||||||
|
|
||||||
export async function logout(
|
export async function logout(
|
||||||
accessToken: string,
|
accessToken: string,
|
||||||
refreshToken: string
|
refreshTokenStr: string
|
||||||
): Promise<ApiResponse<{ message: string }>> {
|
): Promise<ApiResponse<{ message: string }>> {
|
||||||
return request<{ message: string }>("/auth/logout", {
|
return request<{ message: string }>("/auth/logout", {
|
||||||
method: "POST",
|
method: "POST",
|
||||||
headers: authHeaders(accessToken),
|
headers: authHeaders(accessToken),
|
||||||
body: JSON.stringify({ refresh_token: refreshToken }),
|
body: JSON.stringify({ refresh_token: refreshTokenStr }),
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ import {
|
|||||||
} from "react";
|
} from "react";
|
||||||
import * as api from "./api";
|
import * as api from "./api";
|
||||||
import type { AuthUser } from "./api";
|
import type { AuthUser } from "./api";
|
||||||
|
import { setAuthCallbacks } from "./api";
|
||||||
import {
|
import {
|
||||||
clearAuth,
|
clearAuth,
|
||||||
loadAccessToken,
|
loadAccessToken,
|
||||||
@@ -108,6 +109,23 @@ export function AuthProvider({ children }: { children: ReactNode }) {
|
|||||||
[clearRefreshTimer, persistAuth]
|
[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 并尝试刷新
|
// 初始化:检查已有 token 并尝试刷新
|
||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
const init = async () => {
|
const init = async () => {
|
||||||
|
|||||||
Reference in New Issue
Block a user