Files
gochat/internal/auth/jwt_test.go_BAK
T
2026-06-04 15:44:48 +08:00

201 lines
5.3 KiB
Plaintext

package auth
import (
"testing"
"time"
"github.com/golang-jwt/jwt/v5"
"github.com/stretchr/testify/assert"
"github.com/gochat/gochat/internal/config"
"github.com/gochat/gochat/internal/model"
)
func newTestJWTConfig() *config.JWTConfig {
return &config.JWTConfig{
Secret: "test-secret-key-for-unit-tests",
ExpiryHours: 1,
RefreshExpiryHours: 168, // 7 days
}
}
func newTestUser() *model.User {
return &model.User{
Base: model.Base{ID: 42},
Name: "Test Agent",
Email: "test@example.com",
Provider: "email",
Role: string(model.AccountUserRoleAgent),
}
}
// --- JWTService GenerateTokenPair Tests ---
func TestJWTService_GenerateTokenPair(t *testing.T) {
cfg := newTestJWTConfig()
svc := NewJWTService(cfg)
user := newTestUser()
pair, err := svc.GenerateTokenPair(user, 1, "agent")
assert.NoError(t, err)
assert.NotEmpty(t, pair.AccessToken)
assert.NotEmpty(t, pair.RefreshToken)
assert.True(t, pair.ExpiresAt.After(time.Now()))
}
func TestJWTService_GenerateTokenPair_AdminRole(t *testing.T) {
cfg := newTestJWTConfig()
svc := NewJWTService(cfg)
user := newTestUser()
pair, err := svc.GenerateTokenPair(user, 1, "administrator")
assert.NoError(t, err)
assert.NotEmpty(t, pair.AccessToken)
}
func TestJWTService_GenerateTokenPair_CustomRole(t *testing.T) {
cfg := newTestJWTConfig()
svc := NewJWTService(cfg)
user := newTestUser()
user.CustomRoleID = func() *uint { v := uint(10); return &v }()
pair, err := svc.GenerateTokenPair(user, 5, "custom_role")
assert.NoError(t, err)
assert.NotEmpty(t, pair.AccessToken)
// Verify claims include custom_role_id
claims, err := svc.ValidateAccessToken(pair.AccessToken)
assert.NoError(t, err)
assert.Equal(t, uint(10), claims.CustomRoleID)
}
// --- JWTService ValidateAccessToken Tests ---
func TestJWTService_ValidateAccessToken_Valid(t *testing.T) {
cfg := newTestJWTConfig()
svc := NewJWTService(cfg)
user := newTestUser()
pair, err := svc.GenerateTokenPair(user, 1, "agent")
assert.NoError(t, err)
claims, err := svc.ValidateAccessToken(pair.AccessToken)
assert.NoError(t, err)
assert.Equal(t, user.ID, claims.UserID)
assert.Equal(t, uint(1), claims.AccountID)
assert.Equal(t, "agent", claims.Role)
assert.Equal(t, "email", claims.Provider)
}
func TestJWTService_ValidateAccessToken_Expired(t *testing.T) {
cfg := &config.JWTConfig{
Secret: "test-secret-key-for-unit-tests",
ExpiryHours: 0, // effectively expired immediately (1h min, but let's use negative trick)
RefreshExpiryHours: 168,
}
// Create a manually-expired token
claims := &Claims{
UserID: 1,
AccountID: 1,
Role: "agent",
Provider: "email",
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(time.Now().Add(-1 * time.Hour)),
IssuedAt: jwt.NewNumericDate(time.Now().Add(-2 * time.Hour)),
Subject: "1",
},
}
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
tokenString, err := token.SignedString([]byte(cfg.Secret))
assert.NoError(t, err)
svc := NewJWTService(cfg)
validatedClaims, err := svc.ValidateAccessToken(tokenString)
assert.Error(t, err)
assert.Nil(t, validatedClaims)
}
func TestJWTService_ValidateAccessToken_InvalidSignature(t *testing.T) {
cfg := newTestJWTConfig()
svc := NewJWTService(cfg)
user := newTestUser()
pair, err := svc.GenerateTokenPair(user, 1, "agent")
assert.NoError(t, err)
// Validate with a different secret
differentCfg := &config.JWTConfig{
Secret: "different-secret-key",
ExpiryHours: 1,
RefreshExpiryHours: 168,
}
differentSvc := NewJWTService(differentCfg)
claims, err := differentSvc.ValidateAccessToken(pair.AccessToken)
assert.Error(t, err)
assert.Nil(t, claims)
}
func TestJWTService_ValidateAccessToken_Malformed(t *testing.T) {
cfg := newTestJWTConfig()
svc := NewJWTService(cfg)
claims, err := svc.ValidateAccessToken("not-a-valid-token")
assert.Error(t, err)
assert.Nil(t, claims)
}
// --- JWTService ValidateRefreshToken Tests ---
func TestJWTService_ValidateRefreshToken_Valid(t *testing.T) {
cfg := newTestJWTConfig()
svc := NewJWTService(cfg)
user := newTestUser()
pair, err := svc.GenerateTokenPair(user, 1, "agent")
assert.NoError(t, err)
claims, err := svc.ValidateRefreshToken(pair.RefreshToken)
assert.NoError(t, err)
assert.Equal(t, user.ID, claims.UserID)
assert.Equal(t, "email", claims.Provider)
}
// --- Claims Structure Tests ---
func TestClaims_Fields(t *testing.T) {
claims := &Claims{
UserID: 1,
AccountID: 2,
Role: "administrator",
Provider: "google",
CustomRoleID: 5,
}
assert.Equal(t, uint(1), claims.UserID)
assert.Equal(t, uint(2), claims.AccountID)
assert.Equal(t, "administrator", claims.Role)
assert.Equal(t, "google", claims.Provider)
assert.Equal(t, uint(5), claims.CustomRoleID)
}
// --- TokenPair Structure Tests ---
func TestTokenPair_Fields(t *testing.T) {
now := time.Now()
pair := &TokenPair{
AccessToken: "access-token-value",
RefreshToken: "refresh-token-value",
ExpiresAt: now,
}
assert.Equal(t, "access-token-value", pair.AccessToken)
assert.Equal(t, "refresh-token-value", pair.RefreshToken)
assert.Equal(t, now, pair.ExpiresAt)
}
// --- NewJWTService Constructor Tests ---
func TestNewJWTService(t *testing.T) {
cfg := newTestJWTConfig()
svc := NewJWTService(cfg)
assert.NotNil(t, svc)
assert.Equal(t, cfg, svc.cfg)
}