feat: 完善鉴权模块 #142

Merged
huanghaosheng merged 6 commits from fix/auth into develop 2026-06-20 16:34:48 +08:00
2 changed files with 69 additions and 8 deletions
Showing only changes of commit 87c54b80c0 - Show all commits

View File

@@ -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。

View File

@@ -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