Files
CamTalk/backend/internal/auth/jwt_test.go
hhs 87c54b80c0 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 类型交叉校验测试
2026-06-20 14:50:19 +08:00

166 lines
5.1 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package auth
import (
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestGeneratePair_ReturnsNonEmptyTokens(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)
assert.NotEmpty(t, access)
assert.NotEmpty(t, refresh)
assert.NotEqual(t, access, refresh)
}
func TestValidateAccess_ValidToken(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)
claims, err := tm.ValidateAccess(access)
require.NoError(t, err)
assert.Equal(t, "user-123", claims.UserID)
assert.Equal(t, "alice", claims.Username)
assert.Equal(t, "camtalk", claims.Issuer)
}
func TestValidateRefresh_ValidToken(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)
claims, err := tm.ValidateRefresh(refresh)
require.NoError(t, err)
assert.Equal(t, "user-123", claims.UserID)
assert.Equal(t, "alice", claims.Username)
assert.NotEmpty(t, claims.ID) // refresh token 应含唯一 ID
}
func TestValidateAccess_ExpiredToken(t *testing.T) {
// 使用极短的 TTL
tm := NewTokenManager("test-secret-key", -1*time.Second, -1*time.Second)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
_, err = tm.ValidateAccess(access)
assert.ErrorIs(t, err, ErrInvalidToken)
}
func TestValidateAccess_WrongSecret(t *testing.T) {
tm1 := NewTokenManager("secret-1", 15*time.Minute, 7*24*time.Hour)
tm2 := NewTokenManager("secret-2", 15*time.Minute, 7*24*time.Hour)
access, _, err := tm1.GeneratePair("user-123", "alice")
require.NoError(t, err)
_, err = tm2.ValidateAccess(access)
assert.ErrorIs(t, err, ErrInvalidToken)
}
func TestValidateAccess_InvalidFormat(t *testing.T) {
tm := NewTokenManager("test-secret-key", 15*time.Minute, 7*24*time.Hour)
_, err := tm.ValidateAccess("not-a-valid-token")
assert.ErrorIs(t, err, ErrInvalidToken)
}
func TestValidateAccess_EmptyString(t *testing.T) {
tm := NewTokenManager("test-secret-key", 15*time.Minute, 7*24*time.Hour)
_, err := tm.ValidateAccess("")
assert.ErrorIs(t, err, ErrInvalidToken)
}
func TestHashToken_Deterministic(t *testing.T) {
hash1 := HashToken("some-token-value")
hash2 := HashToken("some-token-value")
assert.Equal(t, hash1, hash2)
assert.Len(t, hash1, 64) // SHA256 hex = 64 chars
}
func TestHashToken_DifferentInputsDifferentHashes(t *testing.T) {
hash1 := HashToken("token-a")
hash2 := HashToken("token-b")
assert.NotEqual(t, hash1, hash2)
}
func TestValidateRefresh_ExpiredToken(t *testing.T) {
tm := NewTokenManager("test-secret-key", -1*time.Minute, -1*time.Minute)
_, refresh, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
_, err = tm.ValidateRefresh(refresh)
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
tm := NewTokenManager("test-secret-key", accessTTL, refreshTTL)
before := time.Now()
access, refresh, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
after := time.Now()
// 校验 access token 有效期范围
accessClaims, err := tm.ValidateAccess(access)
require.NoError(t, err)
assert.True(t, accessClaims.ExpiresAt.Time.After(before.Add(accessTTL).Add(-1*time.Second)))
assert.True(t, accessClaims.ExpiresAt.Time.Before(after.Add(accessTTL).Add(1*time.Second)))
// 校验 refresh token 有效期范围
refreshClaims, err := tm.ValidateRefresh(refresh)
require.NoError(t, err)
assert.True(t, refreshClaims.ExpiresAt.Time.After(before.Add(refreshTTL).Add(-1*time.Second)))
assert.True(t, refreshClaims.ExpiresAt.Time.Before(after.Add(refreshTTL).Add(1*time.Second)))
}