- Claims 新增 TokenType 字段("access" / "refresh") - GeneratePair 为 access/refresh token 分别设置 token_type - ValidateAccess 校验后检查 token_type == "access" - ValidateRefresh 校验后检查 token_type == "refresh" - 增加 token 类型交叉校验测试
166 lines
5.1 KiB
Go
166 lines
5.1 KiB
Go
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)))
|
||
}
|