fix: 完善鉴权模块,修复 CI/CD 中 dubious ownership 错误 #143
@@ -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(三级存储)
|
||||
|
||||
@@ -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。
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 失败已间接验证
|
||||
}
|
||||
|
||||
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 客户端
|
||||
// 职责:封装 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 }),
|
||||
});
|
||||
}
|
||||
|
||||
@@ -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 () => {
|
||||
|
||||
Reference in New Issue
Block a user