From 87c54b80c0c222a74b6310c1f20bf0de65f0bf28 Mon Sep 17 00:00:00 2001 From: hhs <386998068@qq.com> Date: Sat, 20 Jun 2026 14:50:19 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20JWT=20Claims=20=E5=A2=9E=E5=8A=A0=20Tok?= =?UTF-8?q?enType=20=E5=AD=97=E6=AE=B5=EF=BC=8CValidateAccess/ValidateRefr?= =?UTF-8?q?esh=20=E5=8C=BA=E5=88=86=E6=A0=A1=E9=AA=8C=E7=B1=BB=E5=9E=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Claims 新增 TokenType 字段("access" / "refresh") - GeneratePair 为 access/refresh token 分别设置 token_type - ValidateAccess 校验后检查 token_type == "access" - ValidateRefresh 校验后检查 token_type == "refresh" - 增加 token 类型交叉校验测试 --- backend/internal/auth/jwt.go | 39 ++++++++++++++++++++++++------- backend/internal/auth/jwt_test.go | 38 ++++++++++++++++++++++++++++++ 2 files changed, 69 insertions(+), 8 deletions(-) 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