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