Skip to content

Commit 445bfdf

Browse files
authored
Merge pull request Wei-Shaw#706 from PMExtra/feat/default-subscriptions-on-user-create
feat(settings): add default subscriptions for new users
2 parents fc5b9c8 + 0fba190 commit 445bfdf

22 files changed

Lines changed: 751 additions & 50 deletions

backend/cmd/jwtgen/main.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,7 @@ func main() {
3333
}()
3434

3535
userRepo := repository.NewUserRepository(client, sqlDB)
36-
authService := service.NewAuthService(userRepo, nil, nil, cfg, nil, nil, nil, nil, nil)
36+
authService := service.NewAuthService(userRepo, nil, nil, cfg, nil, nil, nil, nil, nil, nil)
3737

3838
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
3939
defer cancel()

backend/cmd/server/wire_gen.go

Lines changed: 5 additions & 5 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

backend/go.mod

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -180,6 +180,8 @@ require (
180180
golang.org/x/text v0.34.0 // indirect
181181
golang.org/x/tools v0.41.0 // indirect
182182
google.golang.org/genproto/googleapis/rpc v0.0.0-20250929231259-57b25ae835d4 // indirect
183+
google.golang.org/grpc v1.75.1 // indirect
184+
google.golang.org/protobuf v1.36.10 // indirect
183185
gopkg.in/ini.v1 v1.67.0 // indirect
184186
modernc.org/libc v1.67.6 // indirect
185187
modernc.org/mathutil v1.7.1 // indirect

backend/internal/handler/admin/setting_handler.go

Lines changed: 60 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -51,6 +51,13 @@ func (h *SettingHandler) GetSettings(c *gin.Context) {
5151

5252
// Check if ops monitoring is enabled (respects config.ops.enabled)
5353
opsEnabled := h.opsService != nil && h.opsService.IsMonitoringEnabled(c.Request.Context())
54+
defaultSubscriptions := make([]dto.DefaultSubscriptionSetting, 0, len(settings.DefaultSubscriptions))
55+
for _, sub := range settings.DefaultSubscriptions {
56+
defaultSubscriptions = append(defaultSubscriptions, dto.DefaultSubscriptionSetting{
57+
GroupID: sub.GroupID,
58+
ValidityDays: sub.ValidityDays,
59+
})
60+
}
5461

5562
response.Success(c, dto.SystemSettings{
5663
RegistrationEnabled: settings.RegistrationEnabled,
@@ -87,6 +94,7 @@ func (h *SettingHandler) GetSettings(c *gin.Context) {
8794
SoraClientEnabled: settings.SoraClientEnabled,
8895
DefaultConcurrency: settings.DefaultConcurrency,
8996
DefaultBalance: settings.DefaultBalance,
97+
DefaultSubscriptions: defaultSubscriptions,
9098
EnableModelFallback: settings.EnableModelFallback,
9199
FallbackModelAnthropic: settings.FallbackModelAnthropic,
92100
FallbackModelOpenAI: settings.FallbackModelOpenAI,
@@ -146,8 +154,9 @@ type UpdateSettingsRequest struct {
146154
SoraClientEnabled bool `json:"sora_client_enabled"`
147155

148156
// 默认配置
149-
DefaultConcurrency int `json:"default_concurrency"`
150-
DefaultBalance float64 `json:"default_balance"`
157+
DefaultConcurrency int `json:"default_concurrency"`
158+
DefaultBalance float64 `json:"default_balance"`
159+
DefaultSubscriptions []dto.DefaultSubscriptionSetting `json:"default_subscriptions"`
151160

152161
// Model fallback configuration
153162
EnableModelFallback bool `json:"enable_model_fallback"`
@@ -194,6 +203,7 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
194203
if req.SMTPPort <= 0 {
195204
req.SMTPPort = 587
196205
}
206+
req.DefaultSubscriptions = normalizeDefaultSubscriptions(req.DefaultSubscriptions)
197207

198208
// Turnstile 参数验证
199209
if req.TurnstileEnabled {
@@ -300,6 +310,13 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
300310
}
301311
req.OpsMetricsIntervalSeconds = &v
302312
}
313+
defaultSubscriptions := make([]service.DefaultSubscriptionSetting, 0, len(req.DefaultSubscriptions))
314+
for _, sub := range req.DefaultSubscriptions {
315+
defaultSubscriptions = append(defaultSubscriptions, service.DefaultSubscriptionSetting{
316+
GroupID: sub.GroupID,
317+
ValidityDays: sub.ValidityDays,
318+
})
319+
}
303320

304321
// 验证最低版本号格式(空字符串=禁用,或合法 semver)
305322
if req.MinClaudeCodeVersion != "" {
@@ -343,6 +360,7 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
343360
SoraClientEnabled: req.SoraClientEnabled,
344361
DefaultConcurrency: req.DefaultConcurrency,
345362
DefaultBalance: req.DefaultBalance,
363+
DefaultSubscriptions: defaultSubscriptions,
346364
EnableModelFallback: req.EnableModelFallback,
347365
FallbackModelAnthropic: req.FallbackModelAnthropic,
348366
FallbackModelOpenAI: req.FallbackModelOpenAI,
@@ -390,6 +408,13 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
390408
response.ErrorFrom(c, err)
391409
return
392410
}
411+
updatedDefaultSubscriptions := make([]dto.DefaultSubscriptionSetting, 0, len(updatedSettings.DefaultSubscriptions))
412+
for _, sub := range updatedSettings.DefaultSubscriptions {
413+
updatedDefaultSubscriptions = append(updatedDefaultSubscriptions, dto.DefaultSubscriptionSetting{
414+
GroupID: sub.GroupID,
415+
ValidityDays: sub.ValidityDays,
416+
})
417+
}
393418

394419
response.Success(c, dto.SystemSettings{
395420
RegistrationEnabled: updatedSettings.RegistrationEnabled,
@@ -426,6 +451,7 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
426451
SoraClientEnabled: updatedSettings.SoraClientEnabled,
427452
DefaultConcurrency: updatedSettings.DefaultConcurrency,
428453
DefaultBalance: updatedSettings.DefaultBalance,
454+
DefaultSubscriptions: updatedDefaultSubscriptions,
429455
EnableModelFallback: updatedSettings.EnableModelFallback,
430456
FallbackModelAnthropic: updatedSettings.FallbackModelAnthropic,
431457
FallbackModelOpenAI: updatedSettings.FallbackModelOpenAI,
@@ -547,6 +573,9 @@ func diffSettings(before *service.SystemSettings, after *service.SystemSettings,
547573
if before.DefaultBalance != after.DefaultBalance {
548574
changed = append(changed, "default_balance")
549575
}
576+
if !equalDefaultSubscriptions(before.DefaultSubscriptions, after.DefaultSubscriptions) {
577+
changed = append(changed, "default_subscriptions")
578+
}
550579
if before.EnableModelFallback != after.EnableModelFallback {
551580
changed = append(changed, "enable_model_fallback")
552581
}
@@ -586,6 +615,35 @@ func diffSettings(before *service.SystemSettings, after *service.SystemSettings,
586615
return changed
587616
}
588617

618+
func normalizeDefaultSubscriptions(input []dto.DefaultSubscriptionSetting) []dto.DefaultSubscriptionSetting {
619+
if len(input) == 0 {
620+
return nil
621+
}
622+
normalized := make([]dto.DefaultSubscriptionSetting, 0, len(input))
623+
for _, item := range input {
624+
if item.GroupID <= 0 || item.ValidityDays <= 0 {
625+
continue
626+
}
627+
if item.ValidityDays > service.MaxValidityDays {
628+
item.ValidityDays = service.MaxValidityDays
629+
}
630+
normalized = append(normalized, item)
631+
}
632+
return normalized
633+
}
634+
635+
func equalDefaultSubscriptions(a, b []service.DefaultSubscriptionSetting) bool {
636+
if len(a) != len(b) {
637+
return false
638+
}
639+
for i := range a {
640+
if a[i].GroupID != b[i].GroupID || a[i].ValidityDays != b[i].ValidityDays {
641+
return false
642+
}
643+
}
644+
return true
645+
}
646+
589647
// TestSMTPRequest 测试SMTP连接请求
590648
type TestSMTPRequest struct {
591649
SMTPHost string `json:"smtp_host" binding:"required"`

backend/internal/handler/dto/settings.go

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -39,8 +39,9 @@ type SystemSettings struct {
3939
PurchaseSubscriptionURL string `json:"purchase_subscription_url"`
4040
SoraClientEnabled bool `json:"sora_client_enabled"`
4141

42-
DefaultConcurrency int `json:"default_concurrency"`
43-
DefaultBalance float64 `json:"default_balance"`
42+
DefaultConcurrency int `json:"default_concurrency"`
43+
DefaultBalance float64 `json:"default_balance"`
44+
DefaultSubscriptions []DefaultSubscriptionSetting `json:"default_subscriptions"`
4445

4546
// Model fallback configuration
4647
EnableModelFallback bool `json:"enable_model_fallback"`
@@ -62,6 +63,11 @@ type SystemSettings struct {
6263
MinClaudeCodeVersion string `json:"min_claude_code_version"`
6364
}
6465

66+
type DefaultSubscriptionSetting struct {
67+
GroupID int64 `json:"group_id"`
68+
ValidityDays int `json:"validity_days"`
69+
}
70+
6571
type PublicSettings struct {
6672
RegistrationEnabled bool `json:"registration_enabled"`
6773
EmailVerifyEnabled bool `json:"email_verify_enabled"`

backend/internal/server/api_contract_test.go

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -499,6 +499,7 @@ func TestAPIContracts(t *testing.T) {
499499
"doc_url": "https://docs.example.com",
500500
"default_concurrency": 5,
501501
"default_balance": 1.25,
502+
"default_subscriptions": [],
502503
"enable_model_fallback": false,
503504
"fallback_model_anthropic": "claude-3-5-sonnet-20241022",
504505
"fallback_model_antigravity": "gemini-2.5-pro",
@@ -620,7 +621,7 @@ func newContractDeps(t *testing.T) *contractDeps {
620621
settingRepo := newStubSettingRepo()
621622
settingService := service.NewSettingService(settingRepo, cfg)
622623

623-
adminService := service.NewAdminService(userRepo, groupRepo, &accountRepo, nil, proxyRepo, apiKeyRepo, redeemRepo, nil, nil, nil, nil, nil, nil)
624+
adminService := service.NewAdminService(userRepo, groupRepo, &accountRepo, nil, proxyRepo, apiKeyRepo, redeemRepo, nil, nil, nil, nil, nil, nil, nil, nil)
624625
authHandler := handler.NewAuthHandler(cfg, nil, userService, settingService, nil, redeemService, nil)
625626
apiKeyHandler := handler.NewAPIKeyHandler(apiKeyService)
626627
usageHandler := handler.NewUsageHandler(usageService, apiKeyService)

backend/internal/server/middleware/admin_auth_test.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,7 @@ func TestAdminAuthJWTValidatesTokenVersion(t *testing.T) {
1919
gin.SetMode(gin.TestMode)
2020

2121
cfg := &config.Config{JWT: config.JWTConfig{Secret: "test-secret", ExpireHour: 1}}
22-
authService := service.NewAuthService(nil, nil, nil, cfg, nil, nil, nil, nil, nil)
22+
authService := service.NewAuthService(nil, nil, nil, cfg, nil, nil, nil, nil, nil, nil)
2323

2424
admin := &service.User{
2525
ID: 1,

backend/internal/server/middleware/jwt_auth_test.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -40,7 +40,7 @@ func newJWTTestEnv(users map[int64]*service.User) (*gin.Engine, *service.AuthSer
4040
cfg.JWT.AccessTokenExpireMinutes = 60
4141

4242
userRepo := &stubJWTUserRepo{users: users}
43-
authSvc := service.NewAuthService(userRepo, nil, nil, cfg, nil, nil, nil, nil, nil)
43+
authSvc := service.NewAuthService(userRepo, nil, nil, cfg, nil, nil, nil, nil, nil, nil)
4444
userSvc := service.NewUserService(userRepo, nil, nil)
4545
mw := NewJWTAuthMiddleware(authSvc, userSvc)
4646

backend/internal/service/admin_service.go

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -420,6 +420,8 @@ type adminServiceImpl struct {
420420
proxyLatencyCache ProxyLatencyCache
421421
authCacheInvalidator APIKeyAuthCacheInvalidator
422422
entClient *dbent.Client // 用于开启数据库事务
423+
settingService *SettingService
424+
defaultSubAssigner DefaultSubscriptionAssigner
423425
}
424426

425427
type userGroupRateBatchReader interface {
@@ -445,6 +447,8 @@ func NewAdminService(
445447
proxyLatencyCache ProxyLatencyCache,
446448
authCacheInvalidator APIKeyAuthCacheInvalidator,
447449
entClient *dbent.Client,
450+
settingService *SettingService,
451+
defaultSubAssigner DefaultSubscriptionAssigner,
448452
) AdminService {
449453
return &adminServiceImpl{
450454
userRepo: userRepo,
@@ -460,6 +464,8 @@ func NewAdminService(
460464
proxyLatencyCache: proxyLatencyCache,
461465
authCacheInvalidator: authCacheInvalidator,
462466
entClient: entClient,
467+
settingService: settingService,
468+
defaultSubAssigner: defaultSubAssigner,
463469
}
464470
}
465471

@@ -544,9 +550,27 @@ func (s *adminServiceImpl) CreateUser(ctx context.Context, input *CreateUserInpu
544550
if err := s.userRepo.Create(ctx, user); err != nil {
545551
return nil, err
546552
}
553+
s.assignDefaultSubscriptions(ctx, user.ID)
547554
return user, nil
548555
}
549556

557+
func (s *adminServiceImpl) assignDefaultSubscriptions(ctx context.Context, userID int64) {
558+
if s.settingService == nil || s.defaultSubAssigner == nil || userID <= 0 {
559+
return
560+
}
561+
items := s.settingService.GetDefaultSubscriptions(ctx)
562+
for _, item := range items {
563+
if _, _, err := s.defaultSubAssigner.AssignOrExtendSubscription(ctx, &AssignSubscriptionInput{
564+
UserID: userID,
565+
GroupID: item.GroupID,
566+
ValidityDays: item.ValidityDays,
567+
Notes: "auto assigned by default user subscriptions setting",
568+
}); err != nil {
569+
logger.LegacyPrintf("service.admin", "failed to assign default subscription: user_id=%d group_id=%d err=%v", userID, item.GroupID, err)
570+
}
571+
}
572+
}
573+
550574
func (s *adminServiceImpl) UpdateUser(ctx context.Context, id int64, input *UpdateUserInput) (*User, error) {
551575
user, err := s.userRepo.GetByID(ctx, id)
552576
if err != nil {

backend/internal/service/admin_service_create_user_test.go

Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@ import (
77
"errors"
88
"testing"
99

10+
"github.com/Wei-Shaw/sub2api/internal/config"
1011
"github.com/stretchr/testify/require"
1112
)
1213

@@ -65,3 +66,32 @@ func TestAdminService_CreateUser_CreateError(t *testing.T) {
6566
require.ErrorIs(t, err, createErr)
6667
require.Empty(t, repo.created)
6768
}
69+
70+
func TestAdminService_CreateUser_AssignsDefaultSubscriptions(t *testing.T) {
71+
repo := &userRepoStub{nextID: 21}
72+
assigner := &defaultSubscriptionAssignerStub{}
73+
cfg := &config.Config{
74+
Default: config.DefaultConfig{
75+
UserBalance: 0,
76+
UserConcurrency: 1,
77+
},
78+
}
79+
settingService := NewSettingService(&settingRepoStub{values: map[string]string{
80+
SettingKeyDefaultSubscriptions: `[{"group_id":5,"validity_days":30}]`,
81+
}}, cfg)
82+
svc := &adminServiceImpl{
83+
userRepo: repo,
84+
settingService: settingService,
85+
defaultSubAssigner: assigner,
86+
}
87+
88+
_, err := svc.CreateUser(context.Background(), &CreateUserInput{
89+
Email: "new-user@test.com",
90+
Password: "password",
91+
})
92+
require.NoError(t, err)
93+
require.Len(t, assigner.calls, 1)
94+
require.Equal(t, int64(21), assigner.calls[0].UserID)
95+
require.Equal(t, int64(5), assigner.calls[0].GroupID)
96+
require.Equal(t, 30, assigner.calls[0].ValidityDays)
97+
}

0 commit comments

Comments
 (0)