feat: JWT Claims 增加 TokenType 字段,ValidateAccess/ValidateRefresh 区分校验类型
- Claims 新增 TokenType 字段("access" / "refresh") - GeneratePair 为 access/refresh token 分别设置 token_type - ValidateAccess 校验后检查 token_type == "access" - ValidateRefresh 校验后检查 token_type == "refresh" - 增加 token 类型交叉校验测试
This commit is contained in:
@@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -47,6 +54,7 @@ func (tm *TokenManager) GeneratePair(userID, username string) (access, refresh s
|
|||||||
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),
|
||||||
@@ -64,6 +72,7 @@ func (tm *TokenManager) GeneratePair(userID, username string) (access, refresh s
|
|||||||
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
|
||||||
|
|||||||
Reference in New Issue
Block a user