mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 11:43:23 +08:00
修复邮箱域名注册额度策略
This commit is contained in:
@@ -1460,7 +1460,7 @@ func TestCreateOIDCOAuthAccountExistingEmailNormalizesLegacySpacingAndCase(t *te
|
||||
require.Equal(t, "owner@example.com", storedSession.ResolvedEmail)
|
||||
}
|
||||
|
||||
func TestCreateOIDCOAuthAccountRejectsEmailOutsideRegistrationSuffixWhitelist(t *testing.T) {
|
||||
func TestCreateOIDCOAuthAccountRejectsSecondEmailOutsideRegistrationSuffixWhitelist(t *testing.T) {
|
||||
handler, client := newOAuthPendingFlowTestHandlerWithDependencies(t, oauthPendingFlowTestHandlerOptions{
|
||||
emailVerifyEnabled: true,
|
||||
emailCache: &oauthPendingFlowEmailCacheStub{
|
||||
@@ -1477,6 +1477,14 @@ func TestCreateOIDCOAuthAccountRejectsEmailOutsideRegistrationSuffixWhitelist(t
|
||||
},
|
||||
})
|
||||
ctx := context.Background()
|
||||
_, err := client.User.Create().
|
||||
SetEmail("existing@gmail.com").
|
||||
SetUsername("existing-gmail-user").
|
||||
SetPasswordHash("hash").
|
||||
SetRole(service.RoleUser).
|
||||
SetStatus(service.StatusActive).
|
||||
Save(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
session, err := client.PendingAuthSession.Create().
|
||||
SetSessionToken("suffix-whitelist-session-token").
|
||||
@@ -1505,7 +1513,7 @@ func TestCreateOIDCOAuthAccountRejectsEmailOutsideRegistrationSuffixWhitelist(t
|
||||
|
||||
require.Equal(t, http.StatusBadRequest, recorder.Code)
|
||||
payload := decodeJSONBody(t, recorder)
|
||||
require.Equal(t, "EMAIL_SUFFIX_NOT_ALLOWED", payload["reason"])
|
||||
require.Equal(t, "EMAIL_DOMAIN_REGISTRATION_LIMIT", payload["reason"])
|
||||
|
||||
count, err := client.User.Query().Where(dbuser.EmailEQ("foo@gmail.com")).Count(ctx)
|
||||
require.NoError(t, err)
|
||||
@@ -3033,6 +3041,8 @@ type oauthPendingFlowUserRepo struct {
|
||||
options oauthPendingFlowUserRepoOptions
|
||||
}
|
||||
|
||||
var _ service.RegistrationEmailDomainRepository = (*oauthPendingFlowUserRepo)(nil)
|
||||
|
||||
type oauthPendingFlowUserRepoOptions struct {
|
||||
rejectDeleteWhileAuthIdentityExists bool
|
||||
}
|
||||
@@ -3075,6 +3085,35 @@ func (r *oauthPendingFlowUserRepo) CreateWithEmailAliasGuard(ctx context.Context
|
||||
return r.Create(ctx, user)
|
||||
}
|
||||
|
||||
func (r *oauthPendingFlowUserRepo) CountUsersByEmailDomain(ctx context.Context, domain string) (int, error) {
|
||||
domain = service.NormalizeRegistrationEmailDomain(domain)
|
||||
if domain == "" {
|
||||
return 0, nil
|
||||
}
|
||||
emails, err := r.client.User.Query().Select(dbuser.FieldEmail).Strings(ctx)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
count := 0
|
||||
for _, email := range emails {
|
||||
if service.RegistrationEmailDomain(email) == domain {
|
||||
count++
|
||||
}
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
func (r *oauthPendingFlowUserRepo) CreateWithEmailAliasGuardAndDomainLimit(ctx context.Context, user *service.User, domain string) error {
|
||||
count, err := r.CountUsersByEmailDomain(ctx, domain)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if count > 0 {
|
||||
return service.ErrEmailDomainRegistrationLimit
|
||||
}
|
||||
return r.CreateWithEmailAliasGuard(ctx, user)
|
||||
}
|
||||
|
||||
func (r *oauthPendingFlowUserRepo) GetByID(ctx context.Context, id int64) (*service.User, error) {
|
||||
entity, err := r.client.User.Get(ctx, id)
|
||||
if err != nil {
|
||||
|
||||
@@ -43,16 +43,27 @@ func newUserRepositoryWithSQL(client *dbent.Client, sqlq sqlExecutor) *userRepos
|
||||
}
|
||||
|
||||
func (r *userRepository) Create(ctx context.Context, userIn *service.User) error {
|
||||
return r.create(ctx, userIn, false)
|
||||
return r.create(ctx, userIn, false, "")
|
||||
}
|
||||
|
||||
// CreateWithEmailAliasGuard 见 service.UserRepository:在邮箱唯一性锁内复查收件箱身份,
|
||||
// 供注册路径使用。
|
||||
func (r *userRepository) CreateWithEmailAliasGuard(ctx context.Context, userIn *service.User) error {
|
||||
return r.create(ctx, userIn, true)
|
||||
return r.create(ctx, userIn, true, "")
|
||||
}
|
||||
|
||||
func (r *userRepository) create(ctx context.Context, userIn *service.User, guardEmailAlias bool) error {
|
||||
// CountUsersByEmailDomain 统计指定可注册主域名及其子域名下的未删除用户。
|
||||
func (r *userRepository) CountUsersByEmailDomain(ctx context.Context, domain string) (int, error) {
|
||||
return countUsersByEmailDomainWithClient(ctx, clientFromContext(ctx, r.client), domain)
|
||||
}
|
||||
|
||||
// CreateWithEmailAliasGuardAndDomainLimit 串行化非白名单域名的注册请求,
|
||||
// 并在用户写入的同一事务内复查域名额度。
|
||||
func (r *userRepository) CreateWithEmailAliasGuardAndDomainLimit(ctx context.Context, userIn *service.User, domain string) error {
|
||||
return r.create(ctx, userIn, true, normalizeEmailDomain(domain))
|
||||
}
|
||||
|
||||
func (r *userRepository) create(ctx context.Context, userIn *service.User, guardEmailAlias bool, domainLimit string) error {
|
||||
if userIn == nil {
|
||||
return nil
|
||||
}
|
||||
@@ -84,6 +95,9 @@ func (r *userRepository) create(ctx context.Context, userIn *service.User, guard
|
||||
// 别名变体的字面量不同,唯一索引无法兜底;用收件箱身份锁把同一收件箱的并发注册串行化。
|
||||
lockKeys = append(lockKeys, emailAliasUniquenessLockKey(userIn.Email))
|
||||
}
|
||||
if domainLimit != "" {
|
||||
lockKeys = append(lockKeys, registrationEmailDomainLockKey(domainLimit))
|
||||
}
|
||||
releaseEmailLock, err := lockRepositoryScopedKeys(
|
||||
txCtx,
|
||||
txClient,
|
||||
@@ -95,6 +109,16 @@ func (r *userRepository) create(ctx context.Context, userIn *service.User, guard
|
||||
}
|
||||
defer releaseEmailLock()
|
||||
|
||||
if domainLimit != "" {
|
||||
count, err := countUsersByEmailDomainWithClient(txCtx, txClient, domainLimit)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if count > 0 {
|
||||
return service.ErrEmailDomainRegistrationLimit
|
||||
}
|
||||
}
|
||||
|
||||
if err := ensureNormalizedEmailAvailableWithClient(txCtx, txClient, 0, userIn.Email); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -1232,6 +1256,48 @@ func normalizedEmailUniquenessLockKey(email string) string {
|
||||
return "users:normalized-email:" + normalized
|
||||
}
|
||||
|
||||
func registrationEmailDomainLockKey(domain string) string {
|
||||
domain = normalizeEmailDomain(domain)
|
||||
if domain == "" {
|
||||
return ""
|
||||
}
|
||||
return "users:registration-email-domain:" + domain
|
||||
}
|
||||
|
||||
func normalizeEmailDomain(domain string) string {
|
||||
return service.NormalizeRegistrationEmailDomain(domain)
|
||||
}
|
||||
|
||||
func countUsersByEmailDomainWithClient(ctx context.Context, client *dbent.Client, domain string) (int, error) {
|
||||
client = clientFromContext(ctx, client)
|
||||
domain = normalizeEmailDomain(domain)
|
||||
if client == nil || domain == "" {
|
||||
return 0, nil
|
||||
}
|
||||
return client.User.Query().Where(userEmailDomainPredicate(domain)).Count(ctx)
|
||||
}
|
||||
|
||||
func userEmailDomainPredicate(domain string) predicate.User {
|
||||
domain = normalizeEmailDomain(domain)
|
||||
escapedDomain := escapeLikeWildcards(domain)
|
||||
exactPattern := "%@" + escapedDomain
|
||||
subdomainPattern := "%@%." + escapedDomain
|
||||
return predicate.User(func(s *entsql.Selector) {
|
||||
s.Where(entsql.P(func(b *entsql.Builder) {
|
||||
b.WriteString("(RTRIM(LOWER(TRIM(").
|
||||
Ident(s.C(dbuser.FieldEmail)).
|
||||
WriteString(")), '.') LIKE ").
|
||||
Arg(exactPattern).
|
||||
WriteString(` ESCAPE '\' OR RTRIM(LOWER(TRIM(`).
|
||||
Ident(s.C(dbuser.FieldEmail)).
|
||||
WriteString(")), '.') LIKE ").
|
||||
Arg(subdomainPattern).
|
||||
WriteString(` ESCAPE '\'`).
|
||||
WriteString(")")
|
||||
}))
|
||||
})
|
||||
}
|
||||
|
||||
// emailAliasUniquenessLockKey 按收件箱身份(而非邮箱字面量)加锁,使同一收件箱的不同
|
||||
// 别名变体在注册时互斥。
|
||||
func emailAliasUniquenessLockKey(email string) string {
|
||||
|
||||
@@ -98,3 +98,58 @@ func TestUserRepositoryCreateWithEmailAliasGuard(t *testing.T) {
|
||||
Status: service.StatusActive,
|
||||
}))
|
||||
}
|
||||
|
||||
func TestUserRepositoryCountUsersByEmailDomain(t *testing.T) {
|
||||
repo, _ := newUserEntRepo(t)
|
||||
ctx := context.Background()
|
||||
|
||||
active := &service.User{
|
||||
Email: "first@custom.example",
|
||||
Username: "first",
|
||||
PasswordHash: "hash",
|
||||
Role: service.RoleUser,
|
||||
Status: service.StatusActive,
|
||||
}
|
||||
seedUserForAliasTest(t, repo, active.Email)
|
||||
seedUserForAliasTest(t, repo, "other@sub.custom.example")
|
||||
deleted := &service.User{
|
||||
Email: "deleted@custom.example",
|
||||
Username: "deleted",
|
||||
PasswordHash: "hash",
|
||||
Role: service.RoleUser,
|
||||
Status: service.StatusActive,
|
||||
}
|
||||
require.NoError(t, repo.Create(ctx, deleted))
|
||||
require.NoError(t, repo.Delete(ctx, deleted.ID))
|
||||
|
||||
count, err := repo.CountUsersByEmailDomain(ctx, "custom.example")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 2, count)
|
||||
}
|
||||
|
||||
func TestUserRepositoryCreateWithEmailAliasGuardAndDomainLimit(t *testing.T) {
|
||||
repo, _ := newUserEntRepo(t)
|
||||
ctx := context.Background()
|
||||
seedUserForAliasTest(t, repo, "first@custom.example.")
|
||||
|
||||
err := repo.CreateWithEmailAliasGuardAndDomainLimit(ctx, &service.User{
|
||||
Email: "second@custom.example",
|
||||
Username: "second",
|
||||
PasswordHash: "hash",
|
||||
Role: service.RoleUser,
|
||||
Status: service.StatusActive,
|
||||
}, "custom.example")
|
||||
require.ErrorIs(t, err, service.ErrEmailDomainRegistrationLimit)
|
||||
}
|
||||
|
||||
func TestUserRepositoryCountUsersByEmailDomainEscapesLikeWildcards(t *testing.T) {
|
||||
repo, _ := newUserEntRepo(t)
|
||||
ctx := context.Background()
|
||||
seedUserForAliasTest(t, repo, "first@foo_bar.com")
|
||||
seedUserForAliasTest(t, repo, "other@fooxbar.com")
|
||||
|
||||
count, err := repo.CountUsersByEmailDomain(ctx, "foo_bar.com")
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 1, count)
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
@@ -64,6 +65,44 @@ func (s *UserRepoSuite) mustCreateUser(u *service.User) *service.User {
|
||||
return u
|
||||
}
|
||||
|
||||
func (s *UserRepoSuite) TestCreateWithEmailAliasGuardAndDomainLimitConcurrent() {
|
||||
domain := "race-" + strings.ToLower(strings.ReplaceAll(time.Now().Format("150405.000000000"), ".", "")) + ".example"
|
||||
users := []*service.User{
|
||||
{Email: "first@" + domain, PasswordHash: "hash", Role: service.RoleUser, Status: service.StatusActive},
|
||||
{Email: "second@" + domain, PasswordHash: "hash", Role: service.RoleUser, Status: service.StatusActive},
|
||||
}
|
||||
|
||||
errs := make(chan error, len(users))
|
||||
var wg sync.WaitGroup
|
||||
for _, user := range users {
|
||||
wg.Add(1)
|
||||
go func(user *service.User) {
|
||||
defer wg.Done()
|
||||
errs <- s.repo.CreateWithEmailAliasGuardAndDomainLimit(s.ctx, user, domain)
|
||||
}(user)
|
||||
}
|
||||
wg.Wait()
|
||||
close(errs)
|
||||
|
||||
var success, limited int
|
||||
for err := range errs {
|
||||
switch {
|
||||
case err == nil:
|
||||
success++
|
||||
case errors.Is(err, service.ErrEmailDomainRegistrationLimit):
|
||||
limited++
|
||||
default:
|
||||
s.Require().NoError(err)
|
||||
}
|
||||
}
|
||||
s.Require().Equal(1, success)
|
||||
s.Require().Equal(1, limited)
|
||||
|
||||
count, err := s.repo.CountUsersByEmailDomain(s.ctx, domain)
|
||||
s.Require().NoError(err)
|
||||
s.Require().Equal(1, count)
|
||||
}
|
||||
|
||||
func (s *UserRepoSuite) mustCreateGroup(name string) *service.Group {
|
||||
s.T().Helper()
|
||||
|
||||
|
||||
@@ -13,21 +13,45 @@ import (
|
||||
)
|
||||
|
||||
type userRepoStub struct {
|
||||
user *User
|
||||
getErr error
|
||||
createErr error
|
||||
deleteErr error
|
||||
exists bool
|
||||
existsErr error
|
||||
aliasExists bool
|
||||
aliasErr error
|
||||
guardedCreates int
|
||||
nextID int64
|
||||
created []*User
|
||||
updated []*User
|
||||
deletedIDs []int64
|
||||
usersByEmail map[string]*User
|
||||
getByEmailErr error
|
||||
user *User
|
||||
usersByID map[int64]*User
|
||||
getErr error
|
||||
createErr error
|
||||
deleteErr error
|
||||
exists bool
|
||||
existsErr error
|
||||
aliasExists bool
|
||||
aliasErr error
|
||||
guardedCreates int
|
||||
nextID int64
|
||||
created []*User
|
||||
updated []*User
|
||||
deletedIDs []int64
|
||||
usersByEmail map[string]*User
|
||||
getByEmailErr error
|
||||
getByEmailMisses int
|
||||
domainCounts map[string]int
|
||||
domainCountErr error
|
||||
domainLimitErr error
|
||||
domainLimitedCreates int
|
||||
}
|
||||
|
||||
func (s *userRepoStub) CountUsersByEmailDomain(_ context.Context, domain string) (int, error) {
|
||||
if s.domainCountErr != nil {
|
||||
return 0, s.domainCountErr
|
||||
}
|
||||
return s.domainCounts[domain], nil
|
||||
}
|
||||
|
||||
func (s *userRepoStub) CreateWithEmailAliasGuardAndDomainLimit(ctx context.Context, user *User, domain string) error {
|
||||
s.domainLimitedCreates++
|
||||
if s.domainLimitErr != nil {
|
||||
return s.domainLimitErr
|
||||
}
|
||||
if s.domainCounts[domain] > 0 {
|
||||
return ErrEmailDomainRegistrationLimit
|
||||
}
|
||||
return s.CreateWithEmailAliasGuard(ctx, user)
|
||||
}
|
||||
|
||||
func (s *userRepoStub) Create(ctx context.Context, user *User) error {
|
||||
@@ -61,6 +85,12 @@ func (s *userRepoStub) GetByID(ctx context.Context, id int64) (*User, error) {
|
||||
if s.getErr != nil {
|
||||
return nil, s.getErr
|
||||
}
|
||||
if s.usersByID != nil {
|
||||
if user, ok := s.usersByID[id]; ok {
|
||||
return user, nil
|
||||
}
|
||||
return nil, ErrUserNotFound
|
||||
}
|
||||
if s.user == nil {
|
||||
return nil, ErrUserNotFound
|
||||
}
|
||||
@@ -71,6 +101,10 @@ func (s *userRepoStub) GetByEmail(ctx context.Context, email string) (*User, err
|
||||
if s.getByEmailErr != nil {
|
||||
return nil, s.getByEmailErr
|
||||
}
|
||||
if s.getByEmailMisses > 0 {
|
||||
s.getByEmailMisses--
|
||||
return nil, ErrUserNotFound
|
||||
}
|
||||
if s.usersByEmail != nil {
|
||||
if user, ok := s.usersByEmail[email]; ok {
|
||||
return user, nil
|
||||
|
||||
@@ -42,6 +42,9 @@ func (s *AuthService) SendPendingOAuthVerifyCode(ctx context.Context, email stri
|
||||
if s == nil || s.emailService == nil {
|
||||
return nil, ErrServiceUnavailable
|
||||
}
|
||||
if err := s.validateRegistrationEmailQuota(ctx, email); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
siteName := "Sub2API"
|
||||
if s.settingService != nil {
|
||||
@@ -118,10 +121,6 @@ func (s *AuthService) RegisterOAuthEmailAccount(
|
||||
if isReservedEmail(email) {
|
||||
return nil, nil, ErrEmailReserved
|
||||
}
|
||||
if err := s.validateRegistrationEmailPolicy(ctx, email); err != nil {
|
||||
slog.Error("oauth email register: policy rejected", "email", email, "error", err.Error())
|
||||
return nil, nil, err
|
||||
}
|
||||
if err := s.VerifyOAuthEmailCode(ctx, email, verifyCode); err != nil {
|
||||
slog.Error("oauth email register: verify code failed", "email", email, "error", err.Error())
|
||||
return nil, nil, err
|
||||
@@ -141,6 +140,10 @@ func (s *AuthService) RegisterOAuthEmailAccount(
|
||||
if existsEmail {
|
||||
return nil, nil, ErrEmailExists
|
||||
}
|
||||
if err := s.validateRegistrationEmailQuota(ctx, email); err != nil {
|
||||
slog.Error("oauth email register: policy rejected", "email", email, "error", err.Error())
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
hashedPassword, err := s.HashPassword(password)
|
||||
if err != nil {
|
||||
@@ -160,12 +163,16 @@ func (s *AuthService) RegisterOAuthEmailAccount(
|
||||
SignupSource: signupSource,
|
||||
}
|
||||
|
||||
if err := s.userRepo.CreateWithEmailAliasGuard(ctx, user); err != nil {
|
||||
if errors.Is(err, ErrEmailExists) {
|
||||
if err := s.createUserWithRegistrationEmailGuard(ctx, user); err != nil {
|
||||
switch {
|
||||
case errors.Is(err, ErrEmailExists):
|
||||
return nil, nil, ErrEmailExists
|
||||
case errors.Is(err, ErrEmailDomainRegistrationLimit):
|
||||
return nil, nil, ErrEmailDomainRegistrationLimit
|
||||
default:
|
||||
slog.Error("oauth email register: userRepo.Create failed", "email", email, "signup_source", signupSource, "error", err.Error())
|
||||
return nil, nil, ErrServiceUnavailable
|
||||
}
|
||||
slog.Error("oauth email register: userRepo.Create failed", "email", email, "signup_source", signupSource, "error", err.Error())
|
||||
return nil, nil, ErrServiceUnavailable
|
||||
}
|
||||
|
||||
tokenPair, err := s.GenerateTokenPair(ctx, user, "")
|
||||
@@ -202,9 +209,6 @@ func (s *AuthService) RegisterVerifiedOAuthEmailAccount(
|
||||
if isReservedEmail(email) {
|
||||
return nil, nil, ErrEmailReserved
|
||||
}
|
||||
if err := s.validateRegistrationEmailPolicy(ctx, email); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
if strings.TrimSpace(password) == "" {
|
||||
return nil, nil, infraerrors.BadRequest("PASSWORD_REQUIRED", "password is required")
|
||||
}
|
||||
@@ -220,6 +224,9 @@ func (s *AuthService) RegisterVerifiedOAuthEmailAccount(
|
||||
if existsEmail {
|
||||
return nil, nil, ErrEmailExists
|
||||
}
|
||||
if err := s.validateRegistrationEmailQuota(ctx, email); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
hashedPassword, err := s.HashPassword(password)
|
||||
if err != nil {
|
||||
@@ -243,11 +250,15 @@ func (s *AuthService) RegisterVerifiedOAuthEmailAccount(
|
||||
SignupSource: signupSource,
|
||||
}
|
||||
|
||||
if err := s.userRepo.CreateWithEmailAliasGuard(ctx, user); err != nil {
|
||||
if errors.Is(err, ErrEmailExists) {
|
||||
if err := s.createUserWithRegistrationEmailGuard(ctx, user); err != nil {
|
||||
switch {
|
||||
case errors.Is(err, ErrEmailExists):
|
||||
return nil, nil, ErrEmailExists
|
||||
case errors.Is(err, ErrEmailDomainRegistrationLimit):
|
||||
return nil, nil, ErrEmailDomainRegistrationLimit
|
||||
default:
|
||||
return nil, nil, ErrServiceUnavailable
|
||||
}
|
||||
return nil, nil, ErrServiceUnavailable
|
||||
}
|
||||
|
||||
tokenPair, err := s.GenerateTokenPair(ctx, user, "")
|
||||
|
||||
@@ -199,6 +199,87 @@ func TestRegisterOAuthEmailAccountRollsBackCreatedUserWhenTokenPairGenerationFai
|
||||
require.Empty(t, redeemRepo.updateCalls)
|
||||
}
|
||||
|
||||
func TestRegisterOAuthEmailAccount_NonWhitelistDomainLimit(t *testing.T) {
|
||||
userRepo := &userRepoStub{domainCounts: map[string]int{"custom.example": 1}}
|
||||
authService := newOAuthEmailFlowAuthService(
|
||||
userRepo,
|
||||
&redeemCodeRepoStub{},
|
||||
&refreshTokenCacheStub{},
|
||||
map[string]string{
|
||||
SettingKeyRegistrationEnabled: "true",
|
||||
SettingKeyRegistrationEmailSuffixWhitelist: `["@example.com"]`,
|
||||
},
|
||||
&emailCacheStub{data: &VerificationCodeData{
|
||||
Code: "246810",
|
||||
CreatedAt: time.Now().UTC(),
|
||||
ExpiresAt: time.Now().UTC().Add(15 * time.Minute),
|
||||
}},
|
||||
nil,
|
||||
)
|
||||
|
||||
_, _, err := authService.RegisterOAuthEmailAccount(
|
||||
context.Background(),
|
||||
"second@custom.example",
|
||||
"secret-123",
|
||||
"246810",
|
||||
"",
|
||||
"oidc",
|
||||
)
|
||||
|
||||
require.ErrorIs(t, err, ErrEmailDomainRegistrationLimit)
|
||||
}
|
||||
|
||||
func TestRegisterVerifiedOAuthEmailAccount_NonWhitelistDomainLimit(t *testing.T) {
|
||||
userRepo := &userRepoStub{domainCounts: map[string]int{"custom.example": 1}}
|
||||
authService := newOAuthEmailFlowAuthService(
|
||||
userRepo,
|
||||
nil,
|
||||
&refreshTokenCacheStub{},
|
||||
map[string]string{
|
||||
SettingKeyRegistrationEnabled: "true",
|
||||
SettingKeyRegistrationEmailSuffixWhitelist: `["@example.com"]`,
|
||||
},
|
||||
&emailCacheStub{},
|
||||
nil,
|
||||
)
|
||||
|
||||
_, _, err := authService.RegisterVerifiedOAuthEmailAccount(
|
||||
context.Background(),
|
||||
"second@custom.example",
|
||||
"secret-123",
|
||||
"",
|
||||
"oidc",
|
||||
)
|
||||
|
||||
require.ErrorIs(t, err, ErrEmailDomainRegistrationLimit)
|
||||
}
|
||||
|
||||
func TestSendPendingOAuthVerifyCode_NonWhitelistDomainLimit(t *testing.T) {
|
||||
userRepo := &userRepoStub{domainCounts: map[string]int{"custom.example": 1}}
|
||||
authService := newOAuthEmailFlowAuthService(
|
||||
userRepo,
|
||||
nil,
|
||||
nil,
|
||||
map[string]string{
|
||||
SettingKeyRegistrationEnabled: "true",
|
||||
SettingKeyRegistrationEmailSuffixWhitelist: `["@example.com"]`,
|
||||
},
|
||||
&emailCacheStub{},
|
||||
nil,
|
||||
)
|
||||
|
||||
_, err := authService.SendPendingOAuthVerifyCode(context.Background(), "second@custom.example")
|
||||
require.ErrorIs(t, err, ErrEmailDomainRegistrationLimit)
|
||||
}
|
||||
|
||||
func TestSendPendingOAuthVerifyCode_NilServiceReturnsUnavailable(t *testing.T) {
|
||||
var authService *AuthService
|
||||
|
||||
_, err := authService.SendPendingOAuthVerifyCode(context.Background(), "fresh@example.com")
|
||||
|
||||
require.ErrorIs(t, err, ErrServiceUnavailable)
|
||||
}
|
||||
|
||||
func TestRegisterOAuthEmailAccountSetsNormalizedSignupSourceOnCreatedUser(t *testing.T) {
|
||||
userRepo := &userRepoStub{nextID: 42}
|
||||
emailCache := &emailCacheStub{
|
||||
|
||||
@@ -24,20 +24,24 @@ import (
|
||||
)
|
||||
|
||||
var (
|
||||
ErrInvalidCredentials = infraerrors.Unauthorized("INVALID_CREDENTIALS", "invalid email or password")
|
||||
ErrUserNotActive = infraerrors.Forbidden("USER_NOT_ACTIVE", "user is not active")
|
||||
ErrEmailExists = infraerrors.Conflict("EMAIL_EXISTS", "email already exists")
|
||||
ErrEmailReserved = infraerrors.BadRequest("EMAIL_RESERVED", "email is reserved")
|
||||
ErrInvalidToken = infraerrors.Unauthorized("INVALID_TOKEN", "invalid token")
|
||||
ErrTokenExpired = infraerrors.Unauthorized("TOKEN_EXPIRED", "token has expired")
|
||||
ErrAccessTokenExpired = infraerrors.Unauthorized("ACCESS_TOKEN_EXPIRED", "access token has expired")
|
||||
ErrTokenTooLarge = infraerrors.BadRequest("TOKEN_TOO_LARGE", "token too large")
|
||||
ErrTokenRevoked = infraerrors.Unauthorized("TOKEN_REVOKED", "token has been revoked")
|
||||
ErrRefreshTokenInvalid = infraerrors.Unauthorized("REFRESH_TOKEN_INVALID", "invalid refresh token")
|
||||
ErrRefreshTokenExpired = infraerrors.Unauthorized("REFRESH_TOKEN_EXPIRED", "refresh token has expired")
|
||||
ErrRefreshTokenReused = infraerrors.Unauthorized("REFRESH_TOKEN_REUSED", "refresh token has been reused")
|
||||
ErrEmailVerifyRequired = infraerrors.BadRequest("EMAIL_VERIFY_REQUIRED", "email verification is required")
|
||||
ErrEmailSuffixNotAllowed = infraerrors.BadRequest("EMAIL_SUFFIX_NOT_ALLOWED", "email suffix is not allowed")
|
||||
ErrInvalidCredentials = infraerrors.Unauthorized("INVALID_CREDENTIALS", "invalid email or password")
|
||||
ErrUserNotActive = infraerrors.Forbidden("USER_NOT_ACTIVE", "user is not active")
|
||||
ErrEmailExists = infraerrors.Conflict("EMAIL_EXISTS", "email already exists")
|
||||
ErrEmailReserved = infraerrors.BadRequest("EMAIL_RESERVED", "email is reserved")
|
||||
ErrInvalidToken = infraerrors.Unauthorized("INVALID_TOKEN", "invalid token")
|
||||
ErrTokenExpired = infraerrors.Unauthorized("TOKEN_EXPIRED", "token has expired")
|
||||
ErrAccessTokenExpired = infraerrors.Unauthorized("ACCESS_TOKEN_EXPIRED", "access token has expired")
|
||||
ErrTokenTooLarge = infraerrors.BadRequest("TOKEN_TOO_LARGE", "token too large")
|
||||
ErrTokenRevoked = infraerrors.Unauthorized("TOKEN_REVOKED", "token has been revoked")
|
||||
ErrRefreshTokenInvalid = infraerrors.Unauthorized("REFRESH_TOKEN_INVALID", "invalid refresh token")
|
||||
ErrRefreshTokenExpired = infraerrors.Unauthorized("REFRESH_TOKEN_EXPIRED", "refresh token has expired")
|
||||
ErrRefreshTokenReused = infraerrors.Unauthorized("REFRESH_TOKEN_REUSED", "refresh token has been reused")
|
||||
ErrEmailVerifyRequired = infraerrors.BadRequest("EMAIL_VERIFY_REQUIRED", "email verification is required")
|
||||
ErrEmailSuffixNotAllowed = infraerrors.BadRequest("EMAIL_SUFFIX_NOT_ALLOWED", "email suffix is not allowed")
|
||||
ErrEmailDomainRegistrationLimit = infraerrors.BadRequest(
|
||||
"EMAIL_DOMAIN_REGISTRATION_LIMIT",
|
||||
"this email domain cannot register another account; use a mainstream email or contact support to add the enterprise domain",
|
||||
)
|
||||
ErrRegDisabled = infraerrors.Forbidden("REGISTRATION_DISABLED", "registration is currently disabled")
|
||||
ErrServiceUnavailable = infraerrors.ServiceUnavailable("SERVICE_UNAVAILABLE", "service temporarily unavailable")
|
||||
ErrInvitationCodeRequired = infraerrors.BadRequest("INVITATION_CODE_REQUIRED", "invitation code is required")
|
||||
@@ -166,10 +170,6 @@ func (s *AuthService) RegisterWithVerification(ctx context.Context, email, passw
|
||||
if isReservedEmail(email) {
|
||||
return "", nil, ErrEmailReserved
|
||||
}
|
||||
if err := s.validateRegistrationEmailPolicy(ctx, email); err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
|
||||
// 检查是否需要邀请码
|
||||
var invitationRedeemCode *RedeemCode
|
||||
if s.settingService != nil && s.settingService.IsInvitationCodeEnabled(ctx) {
|
||||
@@ -216,6 +216,9 @@ func (s *AuthService) RegisterWithVerification(ctx context.Context, email, passw
|
||||
if existsEmail {
|
||||
return "", nil, ErrEmailExists
|
||||
}
|
||||
if err := s.validateRegistrationEmailQuota(ctx, email); err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
|
||||
// 密码哈希
|
||||
hashedPassword, err := s.HashPassword(password)
|
||||
@@ -242,13 +245,17 @@ func (s *AuthService) RegisterWithVerification(ctx context.Context, email, passw
|
||||
Status: StatusActive,
|
||||
}
|
||||
|
||||
if err := s.userRepo.CreateWithEmailAliasGuard(ctx, user); err != nil {
|
||||
if err := s.createUserWithRegistrationEmailGuard(ctx, user); err != nil {
|
||||
// 优先检查邮箱冲突错误(竞态条件下可能发生)
|
||||
if errors.Is(err, ErrEmailExists) {
|
||||
switch {
|
||||
case errors.Is(err, ErrEmailExists):
|
||||
return "", nil, ErrEmailExists
|
||||
case errors.Is(err, ErrEmailDomainRegistrationLimit):
|
||||
return "", nil, ErrEmailDomainRegistrationLimit
|
||||
default:
|
||||
logger.LegacyPrintf("service.auth", "[Auth] Database error creating user: %v", err)
|
||||
return "", nil, ErrServiceUnavailable
|
||||
}
|
||||
logger.LegacyPrintf("service.auth", "[Auth] Database error creating user: %v", err)
|
||||
return "", nil, ErrServiceUnavailable
|
||||
}
|
||||
s.postAuthUserBootstrap(ctx, user, "email", true)
|
||||
s.assignSubscriptions(ctx, user.ID, grantPlan.Subscriptions, "auto assigned by signup defaults")
|
||||
@@ -310,10 +317,6 @@ func (s *AuthService) SendVerifyCode(ctx context.Context, email string, locale .
|
||||
if isReservedEmail(email) {
|
||||
return ErrEmailReserved
|
||||
}
|
||||
if err := s.validateRegistrationEmailPolicy(ctx, email); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 检查邮箱是否已存在(含 +别名 / Gmail 点号变体归一化,防止单个收件箱批量派生注册)
|
||||
existsEmail, err := s.existsByEmailOrAlias(ctx, email)
|
||||
if err != nil {
|
||||
@@ -323,6 +326,9 @@ func (s *AuthService) SendVerifyCode(ctx context.Context, email string, locale .
|
||||
if existsEmail {
|
||||
return ErrEmailExists
|
||||
}
|
||||
if err := s.validateRegistrationEmailQuota(ctx, email); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 发送验证码
|
||||
if s.emailService == nil {
|
||||
@@ -351,10 +357,6 @@ func (s *AuthService) SendVerifyCodeAsync(ctx context.Context, email string, loc
|
||||
if isReservedEmail(email) {
|
||||
return nil, ErrEmailReserved
|
||||
}
|
||||
if err := s.validateRegistrationEmailPolicy(ctx, email); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 检查邮箱是否已存在(含 +别名 / Gmail 点号变体归一化;在发信前拦截,避免批量脚本消耗发信配额)
|
||||
existsEmail, err := s.existsByEmailOrAlias(ctx, email)
|
||||
if err != nil {
|
||||
@@ -365,6 +367,9 @@ func (s *AuthService) SendVerifyCodeAsync(ctx context.Context, email string, loc
|
||||
logger.LegacyPrintf("service.auth", "[Auth] Email already exists: %s", email)
|
||||
return nil, ErrEmailExists
|
||||
}
|
||||
if err := s.validateRegistrationEmailQuota(ctx, email); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 检查邮件队列服务是否配置
|
||||
if s.emailQueueService == nil {
|
||||
@@ -1201,6 +1206,66 @@ func (s *AuthService) validateRegistrationEmailPolicy(ctx context.Context, email
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateRegistrationEmailQuota 保留白名单为空时的全放行行为;配置白名单后,
|
||||
// 非白名单域名每个最多允许一个账户。
|
||||
func (s *AuthService) validateRegistrationEmailQuota(ctx context.Context, email string) error {
|
||||
if s.settingService == nil {
|
||||
return nil
|
||||
}
|
||||
whitelist := s.settingService.GetRegistrationEmailSuffixWhitelist(ctx)
|
||||
if !IsRegistrationEmailSuffixLimited(email, whitelist) {
|
||||
return nil
|
||||
}
|
||||
|
||||
domain := RegistrationEmailDomain(email)
|
||||
if domain == "" {
|
||||
return buildEmailSuffixNotAllowedError(whitelist)
|
||||
}
|
||||
quotaRepo, ok := s.userRepo.(RegistrationEmailDomainRepository)
|
||||
if !ok {
|
||||
// 生产装配必须提供原子仓储能力;没有数据库的 unit 测试桩保留旧路径,
|
||||
// 避免无关测试被注册专用依赖干扰。
|
||||
if s.entClient != nil {
|
||||
return ErrServiceUnavailable
|
||||
}
|
||||
return nil
|
||||
}
|
||||
count, err := quotaRepo.CountUsersByEmailDomain(ctx, domain)
|
||||
if err != nil {
|
||||
logger.LegacyPrintf("service.auth", "[Auth] Failed to count registration email domain %s: %v", domain, err)
|
||||
return ErrServiceUnavailable
|
||||
}
|
||||
if count > 0 {
|
||||
return ErrEmailDomainRegistrationLimit
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *AuthService) createUserWithRegistrationEmailGuard(ctx context.Context, user *User) error {
|
||||
if s == nil || s.userRepo == nil {
|
||||
return ErrServiceUnavailable
|
||||
}
|
||||
whitelist := []string{}
|
||||
if s.settingService != nil {
|
||||
whitelist = s.settingService.GetRegistrationEmailSuffixWhitelist(ctx)
|
||||
}
|
||||
domain := RegistrationEmailDomain(user.Email)
|
||||
if !IsRegistrationEmailSuffixLimited(user.Email, whitelist) {
|
||||
return s.userRepo.CreateWithEmailAliasGuard(ctx, user)
|
||||
}
|
||||
if domain == "" {
|
||||
return buildEmailSuffixNotAllowedError(whitelist)
|
||||
}
|
||||
quotaRepo, ok := s.userRepo.(RegistrationEmailDomainRepository)
|
||||
if !ok {
|
||||
if s.entClient != nil {
|
||||
return ErrServiceUnavailable
|
||||
}
|
||||
return s.userRepo.CreateWithEmailAliasGuard(ctx, user)
|
||||
}
|
||||
return quotaRepo.CreateWithEmailAliasGuardAndDomainLimit(ctx, user, domain)
|
||||
}
|
||||
|
||||
func buildEmailSuffixNotAllowedError(whitelist []string) error {
|
||||
if len(whitelist) == 0 {
|
||||
return ErrEmailSuffixNotAllowed
|
||||
|
||||
@@ -414,20 +414,51 @@ func TestAuthService_Register_ReservedEmail(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestAuthService_Register_EmailSuffixNotAllowed(t *testing.T) {
|
||||
repo := &userRepoStub{}
|
||||
repo := &userRepoStub{domainCounts: map[string]int{"other.com": 1}}
|
||||
service := newAuthService(repo, map[string]string{
|
||||
SettingKeyRegistrationEnabled: "true",
|
||||
SettingKeyRegistrationEmailSuffixWhitelist: `["@example.com","@company.com"]`,
|
||||
}, nil, nil)
|
||||
|
||||
_, _, err := service.Register(context.Background(), "user@other.com", "password")
|
||||
require.ErrorIs(t, err, ErrEmailSuffixNotAllowed)
|
||||
require.ErrorIs(t, err, ErrEmailDomainRegistrationLimit)
|
||||
appErr := infraerrors.FromError(err)
|
||||
require.Contains(t, appErr.Message, "@example.com")
|
||||
require.Contains(t, appErr.Message, "@company.com")
|
||||
require.Equal(t, "EMAIL_SUFFIX_NOT_ALLOWED", appErr.Reason)
|
||||
require.Equal(t, "2", appErr.Metadata["allowed_suffix_count"])
|
||||
require.Equal(t, "@example.com,@company.com", appErr.Metadata["allowed_suffixes"])
|
||||
require.Equal(t, "EMAIL_DOMAIN_REGISTRATION_LIMIT", appErr.Reason)
|
||||
require.Contains(t, appErr.Message, "mainstream email")
|
||||
}
|
||||
|
||||
func TestAuthService_Register_NonWhitelistDomainAllowsFirstAccount(t *testing.T) {
|
||||
repo := &userRepoStub{nextID: 9, domainCounts: map[string]int{"custom.example": 0}}
|
||||
svc := newAuthService(repo, map[string]string{
|
||||
SettingKeyRegistrationEnabled: "true",
|
||||
SettingKeyRegistrationEmailSuffixWhitelist: `["@example.com"]`,
|
||||
}, nil, nil)
|
||||
|
||||
_, user, err := svc.Register(context.Background(), "first@custom.example", "password")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(9), user.ID)
|
||||
}
|
||||
|
||||
func TestAuthService_Register_NonWhitelistDomainRejectsSecondAccount(t *testing.T) {
|
||||
repo := &userRepoStub{domainCounts: map[string]int{"custom.example": 1}}
|
||||
svc := newAuthService(repo, map[string]string{
|
||||
SettingKeyRegistrationEnabled: "true",
|
||||
SettingKeyRegistrationEmailSuffixWhitelist: `["@example.com"]`,
|
||||
}, nil, nil)
|
||||
|
||||
_, _, err := svc.Register(context.Background(), "second@sub.custom.example", "password")
|
||||
require.ErrorIs(t, err, ErrEmailDomainRegistrationLimit)
|
||||
}
|
||||
|
||||
func TestAuthService_Register_EmptyWhitelistAllowsAllDomains(t *testing.T) {
|
||||
repo := &userRepoStub{nextID: 10}
|
||||
svc := newAuthService(repo, map[string]string{
|
||||
SettingKeyRegistrationEnabled: "true",
|
||||
SettingKeyRegistrationEmailSuffixWhitelist: `[]`,
|
||||
}, nil, nil)
|
||||
|
||||
_, _, err := svc.Register(context.Background(), "any@custom.example", "password")
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestAuthService_Register_EmailSuffixAllowed(t *testing.T) {
|
||||
@@ -444,18 +475,38 @@ func TestAuthService_Register_EmailSuffixAllowed(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestAuthService_SendVerifyCode_EmailSuffixNotAllowed(t *testing.T) {
|
||||
repo := &userRepoStub{}
|
||||
repo := &userRepoStub{domainCounts: map[string]int{"other.com": 1}}
|
||||
service := newAuthService(repo, map[string]string{
|
||||
SettingKeyRegistrationEnabled: "true",
|
||||
SettingKeyRegistrationEmailSuffixWhitelist: `["@example.com","@company.com"]`,
|
||||
}, nil, nil)
|
||||
|
||||
err := service.SendVerifyCode(context.Background(), "user@other.com")
|
||||
require.ErrorIs(t, err, ErrEmailSuffixNotAllowed)
|
||||
require.ErrorIs(t, err, ErrEmailDomainRegistrationLimit)
|
||||
appErr := infraerrors.FromError(err)
|
||||
require.Contains(t, appErr.Message, "@example.com")
|
||||
require.Contains(t, appErr.Message, "@company.com")
|
||||
require.Equal(t, "2", appErr.Metadata["allowed_suffix_count"])
|
||||
require.Equal(t, "EMAIL_DOMAIN_REGISTRATION_LIMIT", appErr.Reason)
|
||||
}
|
||||
|
||||
func TestAuthService_SendVerifyCode_NonWhitelistDomainLimit(t *testing.T) {
|
||||
repo := &userRepoStub{domainCounts: map[string]int{"custom.example": 1}}
|
||||
svc := newAuthService(repo, map[string]string{
|
||||
SettingKeyRegistrationEnabled: "true",
|
||||
SettingKeyRegistrationEmailSuffixWhitelist: `["@example.com"]`,
|
||||
}, nil, nil)
|
||||
|
||||
err := svc.SendVerifyCode(context.Background(), "user@custom.example")
|
||||
require.ErrorIs(t, err, ErrEmailDomainRegistrationLimit)
|
||||
}
|
||||
|
||||
func TestAuthService_SendVerifyCodeAsync_NonWhitelistDomainLimit(t *testing.T) {
|
||||
repo := &userRepoStub{domainCounts: map[string]int{"custom.example": 1}}
|
||||
svc := newAuthService(repo, map[string]string{
|
||||
SettingKeyRegistrationEnabled: "true",
|
||||
SettingKeyRegistrationEmailSuffixWhitelist: `["@example.com"]`,
|
||||
}, nil, nil)
|
||||
|
||||
_, err := svc.SendVerifyCodeAsync(context.Background(), "user@custom.example")
|
||||
require.ErrorIs(t, err, ErrEmailDomainRegistrationLimit)
|
||||
}
|
||||
|
||||
func TestAuthService_Register_CreateError(t *testing.T) {
|
||||
|
||||
@@ -5,6 +5,8 @@ import (
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
"golang.org/x/net/publicsuffix"
|
||||
)
|
||||
|
||||
var registrationEmailDomainPattern = regexp.MustCompile(
|
||||
@@ -20,6 +22,30 @@ func RegistrationEmailSuffix(email string) string {
|
||||
return "@" + domain
|
||||
}
|
||||
|
||||
// RegistrationEmailDomain 返回邮箱对应的可注册主域名,用于域名注册额度归一化。
|
||||
// 例如 abc.com 和 abcd.abc.com 都返回 abc.com;无法从公共后缀表归一化时保留原域名。
|
||||
func RegistrationEmailDomain(email string) string {
|
||||
_, domain, ok := splitEmailForPolicy(email)
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
return NormalizeRegistrationEmailDomain(domain)
|
||||
}
|
||||
|
||||
// NormalizeRegistrationEmailDomain 将邮箱域名归一为可注册主域名。
|
||||
func NormalizeRegistrationEmailDomain(domain string) string {
|
||||
domain = strings.ToLower(strings.TrimSpace(strings.TrimPrefix(domain, "@")))
|
||||
domain = strings.TrimRight(domain, ".")
|
||||
if domain == "" {
|
||||
return ""
|
||||
}
|
||||
registrable, err := publicsuffix.EffectiveTLDPlusOne(domain)
|
||||
if err != nil {
|
||||
return domain
|
||||
}
|
||||
return registrable
|
||||
}
|
||||
|
||||
// IsRegistrationEmailSuffixAllowed checks whether an email is allowed by suffix whitelist.
|
||||
// Empty whitelist means allow all.
|
||||
func IsRegistrationEmailSuffixAllowed(email string, whitelist []string) bool {
|
||||
@@ -43,6 +69,11 @@ func IsRegistrationEmailSuffixAllowed(email string, whitelist []string) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// IsRegistrationEmailSuffixLimited 判断非空白名单是否对该邮箱域名启用单账户额度。
|
||||
func IsRegistrationEmailSuffixLimited(email string, whitelist []string) bool {
|
||||
return len(whitelist) > 0 && !IsRegistrationEmailSuffixAllowed(email, whitelist)
|
||||
}
|
||||
|
||||
// NormalizeRegistrationEmailSuffixWhitelist normalizes and validates suffix whitelist items.
|
||||
func NormalizeRegistrationEmailSuffixWhitelist(raw []string) ([]string, error) {
|
||||
return normalizeRegistrationEmailSuffixWhitelist(raw, true)
|
||||
@@ -146,5 +177,9 @@ func splitEmailForPolicy(raw string) (local string, domain string, ok bool) {
|
||||
if !found || local == "" || domain == "" || strings.Contains(domain, "@") {
|
||||
return "", "", false
|
||||
}
|
||||
domain = strings.TrimRight(domain, ".")
|
||||
if domain == "" {
|
||||
return "", "", false
|
||||
}
|
||||
return local, domain, true
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -30,6 +31,7 @@ func TestParseRegistrationEmailSuffixWhitelist(t *testing.T) {
|
||||
|
||||
func TestIsRegistrationEmailSuffixAllowed(t *testing.T) {
|
||||
require.True(t, IsRegistrationEmailSuffixAllowed("user@example.com", []string{"@example.com"}))
|
||||
require.True(t, IsRegistrationEmailSuffixAllowed("user@example.com.", []string{"@example.com"}))
|
||||
require.False(t, IsRegistrationEmailSuffixAllowed("user@sub.example.com", []string{"@example.com"}))
|
||||
require.True(t, IsRegistrationEmailSuffixAllowed("user@qq.com", []string{"@qq.com"}))
|
||||
require.False(t, IsRegistrationEmailSuffixAllowed("user@sub.qq.com", []string{"@qq.com"}))
|
||||
@@ -42,3 +44,30 @@ func TestIsRegistrationEmailSuffixAllowed(t *testing.T) {
|
||||
require.False(t, IsRegistrationEmailSuffixAllowed("user@c.cn", []string{"@a.com", "*.b.cn"}))
|
||||
require.True(t, IsRegistrationEmailSuffixAllowed("user@any.com", []string{}))
|
||||
}
|
||||
|
||||
func TestRegistrationEmailQuotaRejectsMalformedDomainWhenWhitelistConfigured(t *testing.T) {
|
||||
repo := &userRepoStub{}
|
||||
svc := newAuthService(repo, map[string]string{
|
||||
SettingKeyRegistrationEnabled: "true",
|
||||
SettingKeyRegistrationEmailSuffixWhitelist: `["@example.com"]`,
|
||||
}, nil, nil)
|
||||
|
||||
_, _, err := svc.Register(context.Background(), "malformed-email", "password")
|
||||
|
||||
require.ErrorIs(t, err, ErrEmailSuffixNotAllowed)
|
||||
require.Empty(t, repo.created)
|
||||
}
|
||||
|
||||
func TestIsRegistrationEmailSuffixLimited(t *testing.T) {
|
||||
require.False(t, IsRegistrationEmailSuffixLimited("user@custom.example", nil))
|
||||
require.False(t, IsRegistrationEmailSuffixLimited("user@example.com", []string{"@example.com"}))
|
||||
require.True(t, IsRegistrationEmailSuffixLimited("user@custom.example", []string{"@example.com"}))
|
||||
}
|
||||
|
||||
func TestRegistrationEmailDomainUsesRegistrableDomain(t *testing.T) {
|
||||
require.Equal(t, "abc.com", RegistrationEmailDomain("user@abc.com"))
|
||||
require.Equal(t, "abc.com", RegistrationEmailDomain("user@abcd.abc.com"))
|
||||
require.Equal(t, "example.co.uk", RegistrationEmailDomain("user@team.example.co.uk"))
|
||||
require.Equal(t, "example.com", RegistrationEmailDomain("user@example.com."))
|
||||
require.Equal(t, "example.com", RegistrationEmailDomain("user@team.example.com."))
|
||||
}
|
||||
|
||||
@@ -181,6 +181,13 @@ type UserRepository interface {
|
||||
DisableTotp(ctx context.Context, userID int64) error
|
||||
}
|
||||
|
||||
// RegistrationEmailDomainRepository 是生产用户仓储为非白名单域名单账户兜底策略提供的可选能力。
|
||||
// 它独立于 UserRepository,避免无关测试桩和服务消费者实现注册专用方法。
|
||||
type RegistrationEmailDomainRepository interface {
|
||||
CountUsersByEmailDomain(ctx context.Context, domain string) (int, error)
|
||||
CreateWithEmailAliasGuardAndDomainLimit(ctx context.Context, user *User, domain string) error
|
||||
}
|
||||
|
||||
// RedeemUserAdjustmentRepository provides the atomic, floor-at-zero updates
|
||||
// used by negative-value redeem codes. It is intentionally narrower than
|
||||
// UserRepository because normal usage billing is allowed to overdraw.
|
||||
|
||||
@@ -119,9 +119,9 @@ export default {
|
||||
emailVerificationHint: 'Require email verification for new registrations',
|
||||
emailSuffixWhitelist: 'Email Domain Whitelist',
|
||||
emailSuffixWhitelistHint:
|
||||
"Only email addresses from the specified domains can register (for example, {'@'}qq.com, {'@'}gmail.com, *.edu.cn)",
|
||||
"Emails from allowlist domains can register without a quota. When the allowlist is not empty, every other registrable domain can register one account. Empty the allowlist to remove the quota for all domains (for example, {'@'}qq.com, {'@'}gmail.com, *.edu.cn).",
|
||||
emailSuffixWhitelistPlaceholder: "{'@'}example.com, *.edu.cn",
|
||||
emailSuffixWhitelistInputHint: 'Leave empty for no restriction. Use *.edu.cn to match edu.cn and its subdomains.',
|
||||
emailSuffixWhitelistInputHint: 'Empty the allowlist to remove the registration quota. Use *.edu.cn to match edu.cn and its subdomains.',
|
||||
promoCode: 'Promo Code',
|
||||
promoCodeHint: 'Allow users to use promo codes during registration',
|
||||
invitationCode: 'Invitation Code Registration',
|
||||
|
||||
@@ -234,6 +234,8 @@ export default {
|
||||
USER_NOT_ACTIVE: 'Account has been disabled.',
|
||||
},
|
||||
registrationFailed: 'Registration failed. Please try again.',
|
||||
emailDomainRegistrationLimit:
|
||||
'This email domain cannot register another account. Please use a mainstream email, or contact support to add your enterprise domain to the allowlist.',
|
||||
emailSuffixNotAllowed: 'This email domain is not allowed for registration.',
|
||||
emailSuffixNotAllowedWithAllowed:
|
||||
'This email domain is not allowed. Allowed domains: {suffixes}',
|
||||
|
||||
@@ -119,9 +119,9 @@ export default {
|
||||
emailVerificationHint: '新用户注册时需要验证邮箱',
|
||||
emailSuffixWhitelist: '邮箱域名白名单',
|
||||
emailSuffixWhitelistHint:
|
||||
"仅允许使用指定域名的邮箱注册账号(例如 {'@'}qq.com, {'@'}gmail.com, *.edu.cn)",
|
||||
"白名单域名的邮箱可无限注册;白名单非空时,其他可注册主域名各限注册一个账户。清空白名单后,所有域名均不限制注册数量(例如 {'@'}qq.com, {'@'}gmail.com, *.edu.cn)",
|
||||
emailSuffixWhitelistPlaceholder: "{'@'}example.com, *.edu.cn",
|
||||
emailSuffixWhitelistInputHint: '留空则不限制。使用 *.edu.cn 可匹配 edu.cn 及其子域名。',
|
||||
emailSuffixWhitelistInputHint: '清空白名单后不限制注册数量。使用 *.edu.cn 可匹配 edu.cn 及其子域名。',
|
||||
promoCode: '优惠码',
|
||||
promoCodeHint: '允许用户在注册时使用优惠码',
|
||||
invitationCode: '邀请码注册',
|
||||
|
||||
@@ -234,6 +234,8 @@ export default {
|
||||
USER_NOT_ACTIVE: '账号已被禁用',
|
||||
},
|
||||
registrationFailed: '注册失败,请重试。',
|
||||
emailDomainRegistrationLimit:
|
||||
'该邮箱域名无法注册新账户。请使用主流邮箱注册;如需使用企业邮箱,请联系客服添加域名白名单。',
|
||||
emailSuffixNotAllowed: '该邮箱域名不在允许注册范围内。',
|
||||
emailSuffixNotAllowedWithAllowed: '该邮箱域名不被允许。可用域名:{suffixes}',
|
||||
emailSuffixAllowedMore: '等 {count} 项',
|
||||
|
||||
@@ -2,6 +2,10 @@ import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import { defineComponent, h } from "vue";
|
||||
import { flushPromises, mount } from "@vue/test-utils";
|
||||
|
||||
import enCommon from "@/i18n/locales/en/common";
|
||||
import enSettings from "@/i18n/locales/en/admin/settings";
|
||||
import zhCommon from "@/i18n/locales/zh/common";
|
||||
import zhSettings from "@/i18n/locales/zh/admin/settings";
|
||||
import SettingsView from "../SettingsView.vue";
|
||||
|
||||
const {
|
||||
@@ -592,6 +596,22 @@ async function openUsersTab(wrapper: ReturnType<typeof mountView>) {
|
||||
await flushPromises();
|
||||
}
|
||||
|
||||
describe("admin SettingsView email domain quota copy", () => {
|
||||
it("documents the email domain quota and empty-whitelist behavior in both locales", () => {
|
||||
expect(zhCommon.auth.emailDomainRegistrationLimit).toContain("主流邮箱");
|
||||
expect(zhCommon.auth.emailDomainRegistrationLimit).toContain("联系客服");
|
||||
expect(enCommon.auth.emailDomainRegistrationLimit).toContain("mainstream email");
|
||||
expect(enCommon.auth.emailDomainRegistrationLimit).toContain("contact support");
|
||||
|
||||
const zhHint = zhSettings.settings.registration.emailSuffixWhitelistHint;
|
||||
const enHint = enSettings.settings.registration.emailSuffixWhitelistHint;
|
||||
expect(zhHint).toContain("其他可注册主域名各限注册一个账户");
|
||||
expect(zhHint).toContain("清空白名单");
|
||||
expect(enHint).toContain("one account");
|
||||
expect(enHint).toContain("empty");
|
||||
});
|
||||
});
|
||||
|
||||
describe("admin SettingsView payment visible method controls", () => {
|
||||
beforeEach(() => {
|
||||
getSettings.mockReset();
|
||||
|
||||
@@ -195,18 +195,14 @@ import {
|
||||
} from '@/api/auth'
|
||||
import { apiClient } from '@/api/client'
|
||||
import { buildAuthErrorMessage } from '@/utils/authError'
|
||||
import {
|
||||
formatRegistrationEmailSuffixWhitelistForMessage,
|
||||
isRegistrationEmailSuffixAllowed,
|
||||
normalizeRegistrationEmailSuffixWhitelist
|
||||
} from '@/utils/registrationEmailPolicy'
|
||||
import { extractApiErrorCode } from '@/utils/apiError'
|
||||
import {
|
||||
clearAllAffiliateReferralCodes,
|
||||
loadAffiliateReferralCode,
|
||||
oauthAffiliatePayload
|
||||
} from '@/utils/oauthAffiliate'
|
||||
|
||||
const { t, locale } = useI18n()
|
||||
const { t } = useI18n()
|
||||
|
||||
// ==================== Router & Stores ====================
|
||||
|
||||
@@ -270,7 +266,6 @@ const aliyunCaptchaSceneId = ref<string>('')
|
||||
const aliyunCaptchaPrefix = ref<string>('')
|
||||
const aliyunCaptchaRegion = ref<string>('cn')
|
||||
const siteName = ref<string>('Sub2API')
|
||||
const registrationEmailSuffixWhitelist = ref<string[]>([])
|
||||
|
||||
// Turnstile for resend
|
||||
const turnstileRef = ref<InstanceType<typeof TurnstileWidget> | null>(null)
|
||||
@@ -370,9 +365,6 @@ onMounted(async () => {
|
||||
aliyunCaptchaPrefix.value = settings.aliyun_captcha_prefix || ''
|
||||
aliyunCaptchaRegion.value = settings.aliyun_captcha_region || 'cn'
|
||||
siteName.value = settings.site_name || 'Sub2API'
|
||||
registrationEmailSuffixWhitelist.value = normalizeRegistrationEmailSuffixWhitelist(
|
||||
settings.registration_email_suffix_whitelist || []
|
||||
)
|
||||
} catch (error) {
|
||||
console.error('Failed to load public settings:', error)
|
||||
}
|
||||
@@ -481,10 +473,6 @@ function isPendingOAuthFlow(): boolean {
|
||||
return Boolean(pendingProvider.value.trim())
|
||||
}
|
||||
|
||||
function shouldBypassRegistrationEmailPolicy(): boolean {
|
||||
return isPendingOAuthFlow() || Boolean(pendingAuthToken.value.trim())
|
||||
}
|
||||
|
||||
function resolvePendingOAuthCallbackRoute(provider: string): string {
|
||||
switch (provider.trim().toLowerCase()) {
|
||||
case 'linuxdo':
|
||||
@@ -526,12 +514,6 @@ async function sendCode(): Promise<void> {
|
||||
let captchaProofUsed = false
|
||||
|
||||
try {
|
||||
if (!shouldBypassRegistrationEmailPolicy() && !isRegistrationEmailSuffixAllowed(email.value, registrationEmailSuffixWhitelist.value)) {
|
||||
errorMessage.value = buildEmailSuffixNotAllowedMessage()
|
||||
appStore.showError(errorMessage.value)
|
||||
return
|
||||
}
|
||||
|
||||
const requestPayload = {
|
||||
email: email.value,
|
||||
[pendingAuthTokenField.value]: pendingAuthToken.value || undefined,
|
||||
@@ -575,9 +557,7 @@ async function sendCode(): Promise<void> {
|
||||
|
||||
showResendTurnstile.value = false
|
||||
} catch (error: unknown) {
|
||||
errorMessage.value = buildAuthErrorMessage(error, {
|
||||
fallback: t('auth.sendCodeFailed')
|
||||
})
|
||||
errorMessage.value = buildRegistrationErrorMessage(error, t('auth.sendCodeFailed'))
|
||||
|
||||
appStore.showError(errorMessage.value)
|
||||
} finally {
|
||||
@@ -655,12 +635,6 @@ async function handleVerify(): Promise<void> {
|
||||
return
|
||||
}
|
||||
|
||||
if (!shouldBypassRegistrationEmailPolicy() && !isRegistrationEmailSuffixAllowed(email.value, registrationEmailSuffixWhitelist.value)) {
|
||||
errorMessage.value = buildEmailSuffixNotAllowedMessage()
|
||||
appStore.showError(errorMessage.value)
|
||||
return
|
||||
}
|
||||
|
||||
if (!(await acquireCreateAccountActionProof())) {
|
||||
return
|
||||
}
|
||||
@@ -740,9 +714,7 @@ async function handleVerify(): Promise<void> {
|
||||
// Redirect to dashboard
|
||||
await router.push(pendingRedirect.value || '/dashboard')
|
||||
} catch (error: unknown) {
|
||||
errorMessage.value = buildAuthErrorMessage(error, {
|
||||
fallback: t('auth.verifyFailed')
|
||||
})
|
||||
errorMessage.value = buildRegistrationErrorMessage(error, t('auth.verifyFailed'))
|
||||
|
||||
appStore.showError(errorMessage.value)
|
||||
} finally {
|
||||
@@ -763,20 +735,11 @@ function handleBack(): void {
|
||||
router.push('/register')
|
||||
}
|
||||
|
||||
function buildEmailSuffixNotAllowedMessage(): string {
|
||||
const normalizedWhitelist = normalizeRegistrationEmailSuffixWhitelist(
|
||||
registrationEmailSuffixWhitelist.value
|
||||
)
|
||||
if (normalizedWhitelist.length === 0) {
|
||||
return t('auth.emailSuffixNotAllowed')
|
||||
function buildRegistrationErrorMessage(error: unknown, fallback: string): string {
|
||||
if (extractApiErrorCode(error) === 'EMAIL_DOMAIN_REGISTRATION_LIMIT') {
|
||||
return t('auth.emailDomainRegistrationLimit')
|
||||
}
|
||||
const separator = String(locale.value || '').toLowerCase().startsWith('zh') ? '、' : ', '
|
||||
return t('auth.emailSuffixNotAllowedWithAllowed', {
|
||||
suffixes: formatRegistrationEmailSuffixWhitelistForMessage(normalizedWhitelist, {
|
||||
separator,
|
||||
more: (count) => t('auth.emailSuffixAllowedMore', { count })
|
||||
})
|
||||
})
|
||||
return buildAuthErrorMessage(error, { fallback })
|
||||
}
|
||||
</script>
|
||||
|
||||
|
||||
@@ -353,12 +353,7 @@ import {
|
||||
validateInvitationCode
|
||||
} from '@/api/auth'
|
||||
import { buildAuthErrorMessage } from '@/utils/authError'
|
||||
import { extractI18nErrorMessage } from '@/utils/apiError'
|
||||
import {
|
||||
formatRegistrationEmailSuffixWhitelistForMessage,
|
||||
isRegistrationEmailSuffixAllowed,
|
||||
normalizeRegistrationEmailSuffixWhitelist
|
||||
} from '@/utils/registrationEmailPolicy'
|
||||
import { extractApiErrorCode, extractI18nErrorMessage } from '@/utils/apiError'
|
||||
import {
|
||||
clearAffiliateReferralCode,
|
||||
loadAffiliateReferralCode,
|
||||
@@ -366,7 +361,7 @@ import {
|
||||
} from '@/utils/oauthAffiliate'
|
||||
import type { LoginAgreementDocument } from '@/types'
|
||||
|
||||
const { t, locale } = useI18n()
|
||||
const { t } = useI18n()
|
||||
const LOGIN_AGREEMENT_STORAGE_KEY = 'sub2api_login_agreement_consent'
|
||||
|
||||
// ==================== Router & Stores ====================
|
||||
@@ -405,7 +400,6 @@ const oidcOAuthEnabled = ref<boolean>(false)
|
||||
const oidcOAuthProviderName = ref<string>('OIDC')
|
||||
const githubOAuthEnabled = ref<boolean>(false)
|
||||
const googleOAuthEnabled = ref<boolean>(false)
|
||||
const registrationEmailSuffixWhitelist = ref<string[]>([])
|
||||
const loginAgreementEnabled = ref<boolean>(false)
|
||||
const loginAgreementMode = ref<'modal' | 'checkbox' | string>('modal')
|
||||
const loginAgreementUpdatedAt = ref<string>('')
|
||||
@@ -538,9 +532,6 @@ onMounted(async () => {
|
||||
oidcOAuthProviderName.value = settings.oidc_oauth_provider_name || 'OIDC'
|
||||
githubOAuthEnabled.value = settings.github_oauth_enabled
|
||||
googleOAuthEnabled.value = settings.google_oauth_enabled
|
||||
registrationEmailSuffixWhitelist.value = normalizeRegistrationEmailSuffixWhitelist(
|
||||
settings.registration_email_suffix_whitelist || []
|
||||
)
|
||||
applyLoginAgreementSettings(settings)
|
||||
|
||||
// Read promo code from URL parameter only if promo code is enabled
|
||||
@@ -859,22 +850,6 @@ function validateEmail(email: string): boolean {
|
||||
return emailRegex.test(email)
|
||||
}
|
||||
|
||||
function buildEmailSuffixNotAllowedMessage(): string {
|
||||
const normalizedWhitelist = normalizeRegistrationEmailSuffixWhitelist(
|
||||
registrationEmailSuffixWhitelist.value
|
||||
)
|
||||
if (normalizedWhitelist.length === 0) {
|
||||
return t('auth.emailSuffixNotAllowed')
|
||||
}
|
||||
const separator = String(locale.value || '').toLowerCase().startsWith('zh') ? '、' : ', '
|
||||
return t('auth.emailSuffixNotAllowedWithAllowed', {
|
||||
suffixes: formatRegistrationEmailSuffixWhitelistForMessage(normalizedWhitelist, {
|
||||
separator,
|
||||
more: (count) => t('auth.emailSuffixAllowedMore', { count })
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
function validateForm(): boolean {
|
||||
// Reset errors
|
||||
errors.email = ''
|
||||
@@ -899,11 +874,6 @@ function validateForm(): boolean {
|
||||
} else if (!validateEmail(formData.email)) {
|
||||
errors.email = t('auth.invalidEmail')
|
||||
isValid = false
|
||||
} else if (
|
||||
!isRegistrationEmailSuffixAllowed(formData.email, registrationEmailSuffixWhitelist.value)
|
||||
) {
|
||||
errors.email = buildEmailSuffixNotAllowedMessage()
|
||||
isValid = false
|
||||
}
|
||||
|
||||
// Password validation
|
||||
@@ -1037,9 +1007,7 @@ async function handleRegister(): Promise<void> {
|
||||
await router.push('/dashboard')
|
||||
} catch (error: unknown) {
|
||||
// Handle registration error
|
||||
errorMessage.value = buildAuthErrorMessage(error, {
|
||||
fallback: t('auth.registrationFailed')
|
||||
})
|
||||
errorMessage.value = buildRegistrationErrorMessage(error, t('auth.registrationFailed'))
|
||||
|
||||
// Also show error toast
|
||||
appStore.showError(errorMessage.value)
|
||||
@@ -1050,6 +1018,13 @@ async function handleRegister(): Promise<void> {
|
||||
isLoading.value = false
|
||||
}
|
||||
}
|
||||
|
||||
function buildRegistrationErrorMessage(error: unknown, fallback: string): string {
|
||||
if (extractApiErrorCode(error) === 'EMAIL_DOMAIN_REGISTRATION_LIMIT') {
|
||||
return t('auth.emailDomainRegistrationLimit')
|
||||
}
|
||||
return buildAuthErrorMessage(error, { fallback })
|
||||
}
|
||||
</script>
|
||||
|
||||
<style scoped>
|
||||
|
||||
@@ -64,6 +64,9 @@ vi.mock('vue-i18n', () => ({
|
||||
if (key === 'auth.accountCreatedSuccess') {
|
||||
return `Account created for ${params?.siteName ?? 'Sub2API'}`
|
||||
}
|
||||
if (key === 'auth.emailDomainRegistrationLimit') {
|
||||
return '该邮箱域名无法注册新账户。请使用主流邮箱注册;如需使用企业邮箱,请联系客服添加域名白名单。'
|
||||
}
|
||||
return key
|
||||
},
|
||||
locale: { value: 'en' },
|
||||
@@ -307,6 +310,118 @@ describe('EmailVerifyView', () => {
|
||||
expect(showErrorMock).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('sends a verification code for a non-whitelist email domain', async () => {
|
||||
getPublicSettingsMock.mockResolvedValue({
|
||||
turnstile_enabled: false,
|
||||
turnstile_site_key: '',
|
||||
site_name: 'Sub2API',
|
||||
registration_email_suffix_whitelist: ['allowed.com'],
|
||||
})
|
||||
sessionStorage.setItem(
|
||||
'register_data',
|
||||
JSON.stringify({
|
||||
email: 'first@custom.example',
|
||||
password: 'secret-123',
|
||||
})
|
||||
)
|
||||
|
||||
mount(EmailVerifyView, {
|
||||
global: {
|
||||
stubs: {
|
||||
AuthLayout: { template: '<div><slot /><slot name="footer" /></div>' },
|
||||
Icon: true,
|
||||
TurnstileWidget: true,
|
||||
transition: false,
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
await flushPromises()
|
||||
|
||||
expect(sendVerifyCodeMock).toHaveBeenCalledWith(
|
||||
expect.objectContaining({ email: 'first@custom.example' })
|
||||
)
|
||||
expect(showErrorMock).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('shows the localized domain quota message when sending a verification code is rejected', async () => {
|
||||
getPublicSettingsMock.mockResolvedValue({
|
||||
turnstile_enabled: false,
|
||||
turnstile_site_key: '',
|
||||
site_name: 'Sub2API',
|
||||
registration_email_suffix_whitelist: ['allowed.com'],
|
||||
})
|
||||
sendVerifyCodeMock.mockRejectedValueOnce({
|
||||
reason: 'EMAIL_DOMAIN_REGISTRATION_LIMIT',
|
||||
message: 'raw backend message',
|
||||
})
|
||||
sessionStorage.setItem(
|
||||
'register_data',
|
||||
JSON.stringify({
|
||||
email: 'second@custom.example',
|
||||
password: 'secret-123',
|
||||
})
|
||||
)
|
||||
|
||||
mount(EmailVerifyView, {
|
||||
global: {
|
||||
stubs: {
|
||||
AuthLayout: { template: '<div><slot /><slot name="footer" /></div>' },
|
||||
Icon: true,
|
||||
TurnstileWidget: true,
|
||||
transition: false,
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
await flushPromises()
|
||||
|
||||
expect(showErrorMock).toHaveBeenLastCalledWith(
|
||||
'该邮箱域名无法注册新账户。请使用主流邮箱注册;如需使用企业邮箱,请联系客服添加域名白名单。'
|
||||
)
|
||||
})
|
||||
|
||||
it('shows the localized domain quota message when verified registration is rejected', async () => {
|
||||
getPublicSettingsMock.mockResolvedValue({
|
||||
turnstile_enabled: false,
|
||||
turnstile_site_key: '',
|
||||
site_name: 'Sub2API',
|
||||
registration_email_suffix_whitelist: ['allowed.com'],
|
||||
})
|
||||
sessionStorage.setItem(
|
||||
'register_data',
|
||||
JSON.stringify({
|
||||
email: 'second@custom.example',
|
||||
password: 'secret-123',
|
||||
})
|
||||
)
|
||||
registerMock.mockRejectedValueOnce({
|
||||
reason: 'EMAIL_DOMAIN_REGISTRATION_LIMIT',
|
||||
message: 'raw backend message',
|
||||
})
|
||||
|
||||
const wrapper = mount(EmailVerifyView, {
|
||||
global: {
|
||||
stubs: {
|
||||
AuthLayout: { template: '<div><slot /><slot name="footer" /></div>' },
|
||||
Icon: true,
|
||||
TurnstileWidget: true,
|
||||
transition: false,
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
await flushPromises()
|
||||
await wrapper.get('#code').setValue('123456')
|
||||
await wrapper.get('form').trigger('submit.prevent')
|
||||
await flushPromises()
|
||||
|
||||
expect(registerMock).toHaveBeenCalled()
|
||||
expect(showErrorMock).toHaveBeenLastCalledWith(
|
||||
'该邮箱域名无法注册新账户。请使用主流邮箱注册;如需使用企业邮箱,请联系客服添加域名白名单。'
|
||||
)
|
||||
})
|
||||
|
||||
it('uses the pending oauth verify-code endpoint when auth store only carries the pending provider', async () => {
|
||||
authStoreState.pendingAuthSession = {
|
||||
token: '',
|
||||
|
||||
@@ -2,8 +2,10 @@ import { flushPromises, mount } from '@vue/test-utils'
|
||||
import { beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
import RegisterView from '@/views/auth/RegisterView.vue'
|
||||
|
||||
const { getPublicSettingsMock } = vi.hoisted(() => ({
|
||||
getPublicSettingsMock: vi.fn()
|
||||
const { getPublicSettingsMock, registerMock, showErrorMock } = vi.hoisted(() => ({
|
||||
getPublicSettingsMock: vi.fn(),
|
||||
registerMock: vi.fn(),
|
||||
showErrorMock: vi.fn()
|
||||
}))
|
||||
|
||||
const publicSettings = {
|
||||
@@ -35,15 +37,18 @@ vi.mock('vue-i18n', () => ({
|
||||
}
|
||||
}),
|
||||
useI18n: () => ({
|
||||
t: (key: string) => key,
|
||||
t: (key: string) =>
|
||||
key === 'auth.emailDomainRegistrationLimit'
|
||||
? '该邮箱域名无法注册新账户。请使用主流邮箱注册;如需使用企业邮箱,请联系客服添加域名白名单。'
|
||||
: key,
|
||||
locale: { value: 'en' }
|
||||
})
|
||||
}))
|
||||
|
||||
vi.mock('@/stores', () => ({
|
||||
useAuthStore: () => ({ register: vi.fn() }),
|
||||
useAuthStore: () => ({ register: (...args: unknown[]) => registerMock(...args) }),
|
||||
useAppStore: () => ({
|
||||
showError: vi.fn(),
|
||||
showError: (...args: unknown[]) => showErrorMock(...args),
|
||||
showSuccess: vi.fn(),
|
||||
showWarning: vi.fn()
|
||||
})
|
||||
@@ -79,7 +84,10 @@ function mountRegister() {
|
||||
describe('RegisterView invitation layout', () => {
|
||||
beforeEach(() => {
|
||||
getPublicSettingsMock.mockReset()
|
||||
registerMock.mockReset()
|
||||
showErrorMock.mockReset()
|
||||
getPublicSettingsMock.mockResolvedValue(publicSettings)
|
||||
registerMock.mockResolvedValue({})
|
||||
})
|
||||
|
||||
it('keeps the optional affiliate invitation field before Turnstile', async () => {
|
||||
@@ -109,4 +117,47 @@ describe('RegisterView invitation layout', () => {
|
||||
expect(wrapper.find('[data-testid="affiliate-invitation-field"]').exists()).toBe(false)
|
||||
expect(wrapper.get('#invitation_code').exists()).toBe(true)
|
||||
})
|
||||
|
||||
it('submits a non-whitelist email domain so the backend can enforce its registration quota', async () => {
|
||||
getPublicSettingsMock.mockResolvedValueOnce({
|
||||
...publicSettings,
|
||||
turnstile_enabled: false,
|
||||
registration_email_suffix_whitelist: ['allowed.com']
|
||||
})
|
||||
|
||||
const wrapper = mountRegister()
|
||||
await flushPromises()
|
||||
await wrapper.get('#email').setValue('first@custom.example')
|
||||
await wrapper.get('#password').setValue('secret-123')
|
||||
await wrapper.get('form').trigger('submit.prevent')
|
||||
await flushPromises()
|
||||
|
||||
expect(registerMock).toHaveBeenCalledWith(
|
||||
expect.objectContaining({ email: 'first@custom.example' })
|
||||
)
|
||||
expect(showErrorMock).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('shows the localized registration domain quota message returned by the backend', async () => {
|
||||
getPublicSettingsMock.mockResolvedValueOnce({
|
||||
...publicSettings,
|
||||
turnstile_enabled: false,
|
||||
registration_email_suffix_whitelist: ['allowed.com']
|
||||
})
|
||||
registerMock.mockRejectedValueOnce({
|
||||
reason: 'EMAIL_DOMAIN_REGISTRATION_LIMIT',
|
||||
message: 'raw backend message'
|
||||
})
|
||||
|
||||
const wrapper = mountRegister()
|
||||
await flushPromises()
|
||||
await wrapper.get('#email').setValue('second@custom.example')
|
||||
await wrapper.get('#password').setValue('secret-123')
|
||||
await wrapper.get('form').trigger('submit.prevent')
|
||||
await flushPromises()
|
||||
|
||||
expect(showErrorMock).toHaveBeenCalledWith(
|
||||
'该邮箱域名无法注册新账户。请使用主流邮箱注册;如需使用企业邮箱,请联系客服添加域名白名单。'
|
||||
)
|
||||
})
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user