From fc66a30ffc2ccbc4caa7478598095b93cbf6d2e4 Mon Sep 17 00:00:00 2001 From: superman2003 <2112076433zcr@gmail.com> Date: Fri, 10 Jul 2026 10:45:27 +0800 Subject: [PATCH] fix: harden billing concurrency and payment recovery --- Makefile | 12 +- backend/internal/repository/user_repo.go | 44 +++ .../repository/user_repo_integration_test.go | 47 ++++ .../user_repo_redeem_adjustment_test.go | 56 ++++ .../repository/user_subscription_repo.go | 72 ++++- ...user_subscription_repo_integration_test.go | 49 +++- backend/internal/server/api_contract_test.go | 9 +- .../server/middleware/api_key_auth.go | 15 +- .../server/middleware/api_key_auth_google.go | 14 +- .../middleware/api_key_auth_google_test.go | 13 +- .../server/middleware/api_key_auth_test.go | 97 ++++++- .../internal/service/payment_fulfillment.go | 266 ++++++++++++++---- .../service/payment_fulfillment_test.go | 233 +++++++++++++++ backend/internal/service/redeem_service.go | 31 +- .../subscription_assign_idempotency_test.go | 47 +++- .../subscription_expiry_service_test.go | 10 +- .../service/subscription_reset_quota_test.go | 43 ++- .../internal/service/subscription_service.go | 94 ++++--- backend/internal/service/user_service.go | 8 + .../user_subscription_daily_quota_test.go | 2 +- .../service/user_subscription_port.go | 7 +- .../router/__tests__/feature-access.spec.ts | 177 ++++++++++++ frontend/src/router/index.ts | 29 +- frontend/src/stores/__tests__/app.spec.ts | 127 +++++++++ frontend/src/stores/app.ts | 49 +++- 25 files changed, 1353 insertions(+), 198 deletions(-) create mode 100644 backend/internal/repository/user_repo_redeem_adjustment_test.go create mode 100644 frontend/src/router/__tests__/feature-access.spec.ts diff --git a/Makefile b/Makefile index d00d0c4f5e..c878526f96 100644 --- a/Makefile +++ b/Makefile @@ -1,4 +1,4 @@ -.PHONY: build build-backend build-frontend build-datamanagementd test test-backend test-frontend test-frontend-critical test-datamanagementd secret-scan +.PHONY: build build-backend build-frontend test test-backend test-frontend test-frontend-critical FRONTEND_CRITICAL_VITEST := \ src/views/auth/__tests__/LinuxDoCallbackView.spec.ts \ @@ -19,10 +19,6 @@ build-backend: build-frontend: @pnpm --dir frontend run build -# 编译 datamanagementd(宿主机数据管理进程) -build-datamanagementd: - @cd datamanagement && go build -o datamanagementd ./cmd/datamanagementd - # 运行测试(后端 + 前端) test: test-backend test-frontend @@ -36,9 +32,3 @@ test-frontend: test-frontend-critical: @pnpm --dir frontend exec vitest run $(FRONTEND_CRITICAL_VITEST) - -test-datamanagementd: - @cd datamanagement && go test ./... - -secret-scan: - @python3 tools/secret_scan.py diff --git a/backend/internal/repository/user_repo.go b/backend/internal/repository/user_repo.go index 3ac8dcfbf8..d4c0d0dc28 100644 --- a/backend/internal/repository/user_repo.go +++ b/backend/internal/repository/user_repo.go @@ -32,6 +32,8 @@ type userRepository struct { sql sqlExecutor } +var _ service.RedeemUserAdjustmentRepository = (*userRepository)(nil) + func NewUserRepository(client *dbent.Client, sqlDB *sql.DB) service.UserRepository { return newUserRepositoryWithSQL(client, sqlDB) } @@ -751,6 +753,27 @@ func (r *userRepository) UpdateBalance(ctx context.Context, id int64, amount flo return nil } +func (r *userRepository) ApplyRedeemBalanceAdjustment(ctx context.Context, id int64, delta float64) error { + const updateSQL = ` + UPDATE users + SET balance = GREATEST(balance + $1, 0), updated_at = NOW() + WHERE id = $2 AND deleted_at IS NULL + ` + client := clientFromContext(ctx, r.client) + result, err := client.ExecContext(ctx, updateSQL, delta, id) + if err != nil { + return err + } + affected, err := result.RowsAffected() + if err != nil { + return err + } + if affected == 0 { + return service.ErrUserNotFound + } + return nil +} + // DeductBalance 扣除用户余额 // 透支策略:允许余额变为负数,确保当前请求能够完成 // 中间件会阻止余额 <= 0 的用户发起后续请求 @@ -792,6 +815,27 @@ func (r *userRepository) UpdateConcurrency(ctx context.Context, id int64, amount return nil } +func (r *userRepository) ApplyRedeemConcurrencyAdjustment(ctx context.Context, id int64, delta int) error { + const updateSQL = ` + UPDATE users + SET concurrency = GREATEST(concurrency + $1, 0), updated_at = NOW() + WHERE id = $2 AND deleted_at IS NULL + ` + client := clientFromContext(ctx, r.client) + result, err := client.ExecContext(ctx, updateSQL, delta, id) + if err != nil { + return err + } + affected, err := result.RowsAffected() + if err != nil { + return err + } + if affected == 0 { + return service.ErrUserNotFound + } + return nil +} + func (r *userRepository) BatchSetConcurrency(ctx context.Context, userIDs []int64, value int) (int, error) { if len(userIDs) == 0 { return 0, nil diff --git a/backend/internal/repository/user_repo_integration_test.go b/backend/internal/repository/user_repo_integration_test.go index 13a605a2f5..42d10af632 100644 --- a/backend/internal/repository/user_repo_integration_test.go +++ b/backend/internal/repository/user_repo_integration_test.go @@ -4,6 +4,7 @@ package repository import ( "context" + "sync" "testing" "time" @@ -353,6 +354,29 @@ func (s *UserRepoSuite) TestUpdateBalance_Negative() { s.Require().InDelta(7.0, got.Balance, 1e-6) } +func (s *UserRepoSuite) TestApplyRedeemBalanceAdjustment_ConcurrentNeverNegative() { + user := s.mustCreateUser(&service.User{Email: "redeem-bal-concurrent@test.com", Balance: 10}) + + var wg sync.WaitGroup + errs := make(chan error, 2) + for i := 0; i < 2; i++ { + wg.Add(1) + go func() { + defer wg.Done() + errs <- s.repo.ApplyRedeemBalanceAdjustment(context.Background(), user.ID, -7) + }() + } + wg.Wait() + close(errs) + for err := range errs { + s.Require().NoError(err) + } + + got, err := s.repo.GetByID(s.ctx, user.ID) + s.Require().NoError(err) + s.Require().InDelta(0, got.Balance, 1e-6) +} + func (s *UserRepoSuite) TestDeductBalance() { user := s.mustCreateUser(&service.User{Email: "deduct@test.com", Balance: 10}) @@ -425,6 +449,29 @@ func (s *UserRepoSuite) TestUpdateConcurrency_Negative() { s.Require().Equal(3, got.Concurrency) } +func (s *UserRepoSuite) TestApplyRedeemConcurrencyAdjustment_ConcurrentNeverNegative() { + user := s.mustCreateUser(&service.User{Email: "redeem-concurrency-concurrent@test.com", Concurrency: 10}) + + var wg sync.WaitGroup + errs := make(chan error, 2) + for i := 0; i < 2; i++ { + wg.Add(1) + go func() { + defer wg.Done() + errs <- s.repo.ApplyRedeemConcurrencyAdjustment(context.Background(), user.ID, -7) + }() + } + wg.Wait() + close(errs) + for err := range errs { + s.Require().NoError(err) + } + + got, err := s.repo.GetByID(s.ctx, user.ID) + s.Require().NoError(err) + s.Require().Equal(0, got.Concurrency) +} + // --- ExistsByEmail --- func (s *UserRepoSuite) TestExistsByEmail() { diff --git a/backend/internal/repository/user_repo_redeem_adjustment_test.go b/backend/internal/repository/user_repo_redeem_adjustment_test.go new file mode 100644 index 0000000000..2c0d21e4bd --- /dev/null +++ b/backend/internal/repository/user_repo_redeem_adjustment_test.go @@ -0,0 +1,56 @@ +package repository + +import ( + "context" + "testing" + + "github.com/DATA-DOG/go-sqlmock" + dbent "github.com/Wei-Shaw/sub2api/ent" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" + + "entgo.io/ent/dialect" + entsql "entgo.io/ent/dialect/sql" +) + +func newRedeemAdjustmentRepoMock(t *testing.T) (*userRepository, sqlmock.Sqlmock) { + t.Helper() + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + driver := entsql.OpenDB(dialect.Postgres, db) + client := dbent.NewClient(dbent.Driver(driver)) + t.Cleanup(func() { _ = client.Close() }) + return newUserRepositoryWithSQL(client, db), mock +} + +func TestApplyRedeemBalanceAdjustment_UsesAtomicFloor(t *testing.T) { + repo, mock := newRedeemAdjustmentRepoMock(t) + mock.ExpectExec(`UPDATE users SET balance = GREATEST\(balance \+ \$1, 0\), updated_at = NOW\(\) WHERE id = \$2 AND deleted_at IS NULL`). + WithArgs(-7.0, int64(42)). + WillReturnResult(sqlmock.NewResult(0, 1)) + + require.NoError(t, repo.ApplyRedeemBalanceAdjustment(context.Background(), 42, -7)) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestApplyRedeemConcurrencyAdjustment_UsesAtomicFloor(t *testing.T) { + repo, mock := newRedeemAdjustmentRepoMock(t) + mock.ExpectExec(`UPDATE users SET concurrency = GREATEST\(concurrency \+ \$1, 0\), updated_at = NOW\(\) WHERE id = \$2 AND deleted_at IS NULL`). + WithArgs(-7, int64(42)). + WillReturnResult(sqlmock.NewResult(0, 1)) + + require.NoError(t, repo.ApplyRedeemConcurrencyAdjustment(context.Background(), 42, -7)) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestApplyRedeemAdjustment_MissingUser(t *testing.T) { + repo, mock := newRedeemAdjustmentRepoMock(t) + mock.ExpectExec(`UPDATE users SET balance = GREATEST\(balance \+ \$1, 0\), updated_at = NOW\(\) WHERE id = \$2 AND deleted_at IS NULL`). + WithArgs(-1.0, int64(404)). + WillReturnResult(sqlmock.NewResult(0, 0)) + + err := repo.ApplyRedeemBalanceAdjustment(context.Background(), 404, -1) + require.ErrorIs(t, err, service.ErrUserNotFound) + require.NoError(t, mock.ExpectationsWereMet()) +} diff --git a/backend/internal/repository/user_subscription_repo.go b/backend/internal/repository/user_subscription_repo.go index 6326c9711f..37f06a038b 100644 --- a/backend/internal/repository/user_subscription_repo.go +++ b/backend/internal/repository/user_subscription_repo.go @@ -366,31 +366,85 @@ func (r *userSubscriptionRepository) ActivateWindows(ctx context.Context, id int return translatePersistenceError(err, service.ErrSubscriptionNotFound, nil) } -func (r *userSubscriptionRepository) ResetDailyUsage(ctx context.Context, id int64, newWindowStart time.Time) error { +func (r *userSubscriptionRepository) ResetUsageWindows(ctx context.Context, id int64, resetDaily, resetWeekly, resetMonthly bool, newWindowStart time.Time) error { client := clientFromContext(ctx, r.client) - _, err := client.UserSubscription.UpdateOneID(id). + update := client.UserSubscription.UpdateOneID(id) + if resetDaily { + update.SetDailyUsageUsd(0).SetDailyWindowStart(newWindowStart) + } + if resetWeekly { + update.SetWeeklyUsageUsd(0).SetWeeklyWindowStart(newWindowStart) + } + if resetMonthly { + update.SetMonthlyUsageUsd(0).SetMonthlyWindowStart(newWindowStart) + } + _, err := update.Save(ctx) + return translatePersistenceError(err, service.ErrSubscriptionNotFound, nil) +} + +func (r *userSubscriptionRepository) ResetDailyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error { + client := clientFromContext(ctx, r.client) + query := client.UserSubscription.Update().Where(usersubscription.IDEQ(id)) + if expectedWindowStart == nil { + query = query.Where(usersubscription.DailyWindowStartIsNil()) + } else { + query = query.Where(usersubscription.DailyWindowStartEQ(*expectedWindowStart)) + } + n, err := query. SetDailyUsageUsd(0). SetDailyWindowStart(newWindowStart). Save(ctx) - return translatePersistenceError(err, service.ErrSubscriptionNotFound, nil) + return r.translateConditionalWindowReset(ctx, client, id, n, err) } -func (r *userSubscriptionRepository) ResetWeeklyUsage(ctx context.Context, id int64, newWindowStart time.Time) error { +func (r *userSubscriptionRepository) ResetWeeklyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error { client := clientFromContext(ctx, r.client) - _, err := client.UserSubscription.UpdateOneID(id). + query := client.UserSubscription.Update().Where(usersubscription.IDEQ(id)) + if expectedWindowStart == nil { + query = query.Where(usersubscription.WeeklyWindowStartIsNil()) + } else { + query = query.Where(usersubscription.WeeklyWindowStartEQ(*expectedWindowStart)) + } + n, err := query. SetWeeklyUsageUsd(0). SetWeeklyWindowStart(newWindowStart). Save(ctx) - return translatePersistenceError(err, service.ErrSubscriptionNotFound, nil) + return r.translateConditionalWindowReset(ctx, client, id, n, err) } -func (r *userSubscriptionRepository) ResetMonthlyUsage(ctx context.Context, id int64, newWindowStart time.Time) error { +func (r *userSubscriptionRepository) ResetMonthlyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error { client := clientFromContext(ctx, r.client) - _, err := client.UserSubscription.UpdateOneID(id). + query := client.UserSubscription.Update().Where(usersubscription.IDEQ(id)) + if expectedWindowStart == nil { + query = query.Where(usersubscription.MonthlyWindowStartIsNil()) + } else { + query = query.Where(usersubscription.MonthlyWindowStartEQ(*expectedWindowStart)) + } + n, err := query. SetMonthlyUsageUsd(0). SetMonthlyWindowStart(newWindowStart). Save(ctx) - return translatePersistenceError(err, service.ErrSubscriptionNotFound, nil) + return r.translateConditionalWindowReset(ctx, client, id, n, err) +} + +func (r *userSubscriptionRepository) translateConditionalWindowReset(ctx context.Context, client *dbent.Client, id int64, affected int, err error) error { + if err != nil { + return translatePersistenceError(err, service.ErrSubscriptionNotFound, nil) + } + if affected > 0 { + return nil + } + + // A stale reset is an expected no-op: another request already advanced the + // window. Preserve not-found semantics for callers that target a missing row. + exists, err := client.UserSubscription.Query().Where(usersubscription.IDEQ(id)).Exist(ctx) + if err != nil { + return translatePersistenceError(err, service.ErrSubscriptionNotFound, nil) + } + if !exists { + return service.ErrSubscriptionNotFound + } + return nil } // IncrementUsage 原子性地累加订阅用量。 diff --git a/backend/internal/repository/user_subscription_repo_integration_test.go b/backend/internal/repository/user_subscription_repo_integration_test.go index caa88cc640..96eead494e 100644 --- a/backend/internal/repository/user_subscription_repo_integration_test.go +++ b/backend/internal/repository/user_subscription_repo_integration_test.go @@ -472,7 +472,7 @@ func (s *UserSubscriptionRepoSuite) TestResetDailyUsage() { }) resetAt := time.Date(2025, 1, 2, 0, 0, 0, 0, time.UTC) - err := s.repo.ResetDailyUsage(s.ctx, sub.ID, resetAt) + err := s.repo.ResetDailyUsage(s.ctx, sub.ID, sub.DailyWindowStart, resetAt) s.Require().NoError(err, "ResetDailyUsage") got, err := s.repo.GetByID(s.ctx, sub.ID) @@ -483,6 +483,47 @@ func (s *UserSubscriptionRepoSuite) TestResetDailyUsage() { s.Require().WithinDuration(resetAt, *got.DailyWindowStart, time.Microsecond) } +func (s *UserSubscriptionRepoSuite) TestResetDailyUsage_StaleResetDoesNotClearNewWindowUsage() { + user := s.mustCreateUser("resetd-cas@test.com", service.RoleUser) + group := s.mustCreateGroup("g-resetd-cas") + oldWindowStart := time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC) + sub := s.mustCreateSubscription(user.ID, group.ID, func(c *dbent.UserSubscriptionCreate) { + c.SetDailyWindowStart(oldWindowStart) + c.SetDailyUsageUsd(10) + }) + + newWindowStart := oldWindowStart.Add(24 * time.Hour) + s.Require().NoError(s.repo.ResetDailyUsage(s.ctx, sub.ID, &oldWindowStart, newWindowStart)) + s.Require().NoError(s.repo.IncrementUsage(s.ctx, sub.ID, 3)) + // Simulate a second request carrying the stale old-window snapshot. + s.Require().NoError(s.repo.ResetDailyUsage(s.ctx, sub.ID, &oldWindowStart, newWindowStart)) + + got, err := s.repo.GetByID(s.ctx, sub.ID) + s.Require().NoError(err) + s.Require().InDelta(3, got.DailyUsageUSD, 1e-6) + s.Require().WithinDuration(newWindowStart, *got.DailyWindowStart, time.Microsecond) +} + +func (s *UserSubscriptionRepoSuite) TestResetUsageWindows_ClearsUsageAfterAutomaticWindowAdvance() { + user := s.mustCreateUser("admin-reset-current@test.com", service.RoleUser) + group := s.mustCreateGroup("g-admin-reset-current") + oldWindowStart := time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC) + sub := s.mustCreateSubscription(user.ID, group.ID, func(c *dbent.UserSubscriptionCreate) { + c.SetDailyWindowStart(oldWindowStart) + c.SetDailyUsageUsd(10) + }) + + newWindowStart := oldWindowStart.Add(24 * time.Hour) + s.Require().NoError(s.repo.ResetDailyUsage(s.ctx, sub.ID, &oldWindowStart, newWindowStart)) + s.Require().NoError(s.repo.IncrementUsage(s.ctx, sub.ID, 3)) + s.Require().NoError(s.repo.ResetUsageWindows(s.ctx, sub.ID, true, false, false, newWindowStart)) + + got, err := s.repo.GetByID(s.ctx, sub.ID) + s.Require().NoError(err) + s.Require().InDelta(0, got.DailyUsageUSD, 1e-6) + s.Require().WithinDuration(newWindowStart, *got.DailyWindowStart, time.Microsecond) +} + func (s *UserSubscriptionRepoSuite) TestResetWeeklyUsage() { user := s.mustCreateUser("resetw@test.com", service.RoleUser) group := s.mustCreateGroup("g-resetw") @@ -492,7 +533,7 @@ func (s *UserSubscriptionRepoSuite) TestResetWeeklyUsage() { }) resetAt := time.Date(2025, 1, 6, 0, 0, 0, 0, time.UTC) - err := s.repo.ResetWeeklyUsage(s.ctx, sub.ID, resetAt) + err := s.repo.ResetWeeklyUsage(s.ctx, sub.ID, sub.WeeklyWindowStart, resetAt) s.Require().NoError(err, "ResetWeeklyUsage") got, err := s.repo.GetByID(s.ctx, sub.ID) @@ -511,7 +552,7 @@ func (s *UserSubscriptionRepoSuite) TestResetMonthlyUsage() { }) resetAt := time.Date(2025, 2, 1, 0, 0, 0, 0, time.UTC) - err := s.repo.ResetMonthlyUsage(s.ctx, sub.ID, resetAt) + err := s.repo.ResetMonthlyUsage(s.ctx, sub.ID, sub.MonthlyWindowStart, resetAt) s.Require().NoError(err, "ResetMonthlyUsage") got, err := s.repo.GetByID(s.ctx, sub.ID) @@ -723,7 +764,7 @@ func (s *UserSubscriptionRepoSuite) TestActiveExpiredBoundaries_UsageAndReset_Ba s.Require().NotNil(after.MonthlyWindowStart, "expected MonthlyWindowStart activated") resetAt := time.Now().Truncate(time.Microsecond) // truncate to microsecond for DB precision - s.Require().NoError(s.repo.ResetDailyUsage(s.ctx, active.ID, resetAt), "ResetDailyUsage") + s.Require().NoError(s.repo.ResetDailyUsage(s.ctx, active.ID, after.DailyWindowStart, resetAt), "ResetDailyUsage") afterReset, err := s.repo.GetByID(s.ctx, active.ID) s.Require().NoError(err, "GetByID after reset") s.Require().InDelta(0.0, afterReset.DailyUsageUSD, 1e-6) diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index f48ecb060a..d260afe738 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -2123,13 +2123,16 @@ func (stubUserSubscriptionRepo) UpdateNotes(ctx context.Context, subscriptionID func (stubUserSubscriptionRepo) ActivateWindows(ctx context.Context, id int64, start time.Time) error { return errors.New("not implemented") } -func (stubUserSubscriptionRepo) ResetDailyUsage(ctx context.Context, id int64, newWindowStart time.Time) error { +func (stubUserSubscriptionRepo) ResetUsageWindows(ctx context.Context, id int64, resetDaily, resetWeekly, resetMonthly bool, newWindowStart time.Time) error { return errors.New("not implemented") } -func (stubUserSubscriptionRepo) ResetWeeklyUsage(ctx context.Context, id int64, newWindowStart time.Time) error { +func (stubUserSubscriptionRepo) ResetDailyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error { return errors.New("not implemented") } -func (stubUserSubscriptionRepo) ResetMonthlyUsage(ctx context.Context, id int64, newWindowStart time.Time) error { +func (stubUserSubscriptionRepo) ResetWeeklyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error { + return errors.New("not implemented") +} +func (stubUserSubscriptionRepo) ResetMonthlyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error { return errors.New("not implemented") } func (stubUserSubscriptionRepo) IncrementUsage(ctx context.Context, id int64, costUSD float64) error { diff --git a/backend/internal/server/middleware/api_key_auth.go b/backend/internal/server/middleware/api_key_auth.go index 1610390440..04a09862b5 100644 --- a/backend/internal/server/middleware/api_key_auth.go +++ b/backend/internal/server/middleware/api_key_auth.go @@ -193,6 +193,15 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti // 订阅模式:验证订阅限额 if subscription != nil { needsMaintenance, validateErr := subscriptionService.ValidateAndCheckLimits(subscription, apiKey.Group) + if needsMaintenance { + refreshed, maintenanceErr := subscriptionService.EnsureWindowMaintenance(c.Request.Context(), subscription) + if maintenanceErr != nil { + AbortWithError(c, 500, "SUBSCRIPTION_MAINTENANCE_FAILED", "Failed to maintain subscription usage windows") + return + } + subscription = refreshed + _, validateErr = subscriptionService.ValidateAndCheckLimits(subscription, apiKey.Group) + } if validateErr != nil { code := "SUBSCRIPTION_INVALID" status := 403 @@ -205,12 +214,6 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti AbortWithError(c, status, code, validateErr.Error()) return } - - // 窗口维护异步化(不阻塞请求) - if needsMaintenance { - maintenanceCopy := *subscription - subscriptionService.DoWindowMaintenance(&maintenanceCopy) - } } else { // 非订阅模式 或 订阅模式但 subscriptionService 未注入:回退到余额检查 if apiKeyBalanceBelowAuthThreshold(apiKey.User.Balance, cfg) { diff --git a/backend/internal/server/middleware/api_key_auth_google.go b/backend/internal/server/middleware/api_key_auth_google.go index c75d5b99f2..b910dd9fa4 100644 --- a/backend/internal/server/middleware/api_key_auth_google.go +++ b/backend/internal/server/middleware/api_key_auth_google.go @@ -141,6 +141,15 @@ func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subs } needsMaintenance, err := subscriptionService.ValidateAndCheckLimits(subscription, apiKey.Group) + if needsMaintenance { + refreshed, maintenanceErr := subscriptionService.EnsureWindowMaintenance(c.Request.Context(), subscription) + if maintenanceErr != nil { + abortWithGoogleError(c, 500, "Failed to maintain subscription usage windows") + return + } + subscription = refreshed + _, err = subscriptionService.ValidateAndCheckLimits(subscription, apiKey.Group) + } if err != nil { status := 403 if errors.Is(err, service.ErrDailyLimitExceeded) || @@ -153,11 +162,6 @@ func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subs } c.Set(string(ContextKeySubscription), subscription) - - if needsMaintenance { - maintenanceCopy := *subscription - subscriptionService.DoWindowMaintenance(&maintenanceCopy) - } } else { if apiKeyBalanceBelowAuthThreshold(apiKey.User.Balance, cfg) { abortWithGoogleError(c, 403, "Insufficient account balance") diff --git a/backend/internal/server/middleware/api_key_auth_google_test.go b/backend/internal/server/middleware/api_key_auth_google_test.go index 45ddb0bf93..746d238ca5 100644 --- a/backend/internal/server/middleware/api_key_auth_google_test.go +++ b/backend/internal/server/middleware/api_key_auth_google_test.go @@ -24,6 +24,7 @@ type fakeAPIKeyRepo struct { } type fakeGoogleSubscriptionRepo struct { + getByID func(ctx context.Context, id int64) (*service.UserSubscription, error) getActive func(ctx context.Context, userID, groupID int64) (*service.UserSubscription, error) updateStatus func(ctx context.Context, subscriptionID int64, status string) error activateWindow func(ctx context.Context, id int64, start time.Time) error @@ -115,6 +116,9 @@ func (f fakeGoogleSubscriptionRepo) Create(ctx context.Context, sub *service.Use return errors.New("not implemented") } func (f fakeGoogleSubscriptionRepo) GetByID(ctx context.Context, id int64) (*service.UserSubscription, error) { + if f.getByID != nil { + return f.getByID(ctx, id) + } return nil, errors.New("not implemented") } func (f fakeGoogleSubscriptionRepo) GetByIDIncludeDeleted(ctx context.Context, id int64) (*service.UserSubscription, error) { @@ -174,19 +178,22 @@ func (f fakeGoogleSubscriptionRepo) ActivateWindows(ctx context.Context, id int6 } return errors.New("not implemented") } -func (f fakeGoogleSubscriptionRepo) ResetDailyUsage(ctx context.Context, id int64, start time.Time) error { +func (f fakeGoogleSubscriptionRepo) ResetUsageWindows(context.Context, int64, bool, bool, bool, time.Time) error { + return errors.New("not implemented") +} +func (f fakeGoogleSubscriptionRepo) ResetDailyUsage(ctx context.Context, id int64, _ *time.Time, start time.Time) error { if f.resetDaily != nil { return f.resetDaily(ctx, id, start) } return errors.New("not implemented") } -func (f fakeGoogleSubscriptionRepo) ResetWeeklyUsage(ctx context.Context, id int64, start time.Time) error { +func (f fakeGoogleSubscriptionRepo) ResetWeeklyUsage(ctx context.Context, id int64, _ *time.Time, start time.Time) error { if f.resetWeekly != nil { return f.resetWeekly(ctx, id, start) } return errors.New("not implemented") } -func (f fakeGoogleSubscriptionRepo) ResetMonthlyUsage(ctx context.Context, id int64, start time.Time) error { +func (f fakeGoogleSubscriptionRepo) ResetMonthlyUsage(ctx context.Context, id int64, _ *time.Time, start time.Time) error { if f.resetMonthly != nil { return f.resetMonthly(ctx, id, start) } diff --git a/backend/internal/server/middleware/api_key_auth_test.go b/backend/internal/server/middleware/api_key_auth_test.go index d5fbc46098..abb84ab852 100644 --- a/backend/internal/server/middleware/api_key_auth_test.go +++ b/backend/internal/server/middleware/api_key_auth_test.go @@ -58,7 +58,7 @@ func TestSimpleModeBypassesQuotaCheck(t *testing.T) { }, } - t.Run("standard_mode_needs_maintenance_does_not_block_request", func(t *testing.T) { + t.Run("standard_mode_completes_maintenance_before_request", func(t *testing.T) { cfg := &config.Config{RunMode: config.RunModeStandard} cfg.SubscriptionMaintenance.WorkerCount = 1 cfg.SubscriptionMaintenance.QueueSize = 1 @@ -67,16 +67,22 @@ func TestSimpleModeBypassesQuotaCheck(t *testing.T) { past := time.Now().Add(-48 * time.Hour) sub := &service.UserSubscription{ - ID: 55, - UserID: user.ID, - GroupID: group.ID, - Status: service.SubscriptionStatusActive, - ExpiresAt: time.Now().Add(24 * time.Hour), - DailyWindowStart: &past, - DailyUsageUSD: 0, + ID: 55, + UserID: user.ID, + GroupID: group.ID, + Status: service.SubscriptionStatusActive, + ExpiresAt: time.Now().Add(24 * time.Hour), + DailyWindowStart: &past, + WeeklyWindowStart: &past, + MonthlyWindowStart: &past, + DailyUsageUSD: 0, } maintenanceCalled := make(chan struct{}, 1) subscriptionRepo := &stubUserSubscriptionRepo{ + getByID: func(ctx context.Context, id int64) (*service.UserSubscription, error) { + clone := *sub + return &clone, nil + }, getActive: func(ctx context.Context, userID, groupID int64) (*service.UserSubscription, error) { clone := *sub return &clone, nil @@ -84,11 +90,19 @@ func TestSimpleModeBypassesQuotaCheck(t *testing.T) { updateStatus: func(ctx context.Context, subscriptionID int64, status string) error { return nil }, activateWindow: func(ctx context.Context, id int64, start time.Time) error { return nil }, resetDaily: func(ctx context.Context, id int64, start time.Time) error { + sub.DailyWindowStart = &start + sub.DailyUsageUSD = 0 maintenanceCalled <- struct{}{} return nil }, - resetWeekly: func(ctx context.Context, id int64, start time.Time) error { return nil }, - resetMonthly: func(ctx context.Context, id int64, start time.Time) error { return nil }, + resetWeekly: func(ctx context.Context, id int64, start time.Time) error { + sub.WeeklyWindowStart = &start + return nil + }, + resetMonthly: func(ctx context.Context, id int64, start time.Time) error { + sub.MonthlyWindowStart = &start + return nil + }, } subscriptionService := service.NewSubscriptionService(nil, subscriptionRepo, nil, nil, cfg) t.Cleanup(subscriptionService.Stop) @@ -105,10 +119,57 @@ func TestSimpleModeBypassesQuotaCheck(t *testing.T) { case <-maintenanceCalled: // ok case <-time.After(time.Second): - t.Fatalf("expected maintenance to be scheduled") + t.Fatalf("expected maintenance to complete before response") } }) + t.Run("standard_mode_revalidates_cas_loser_from_database", func(t *testing.T) { + cfg := &config.Config{RunMode: config.RunModeStandard} + apiKeyService := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, cfg) + + past := time.Now().Add(-48 * time.Hour) + current := time.Now() + stale := &service.UserSubscription{ + ID: 56, + UserID: user.ID, + GroupID: group.ID, + Status: service.SubscriptionStatusActive, + ExpiresAt: current.Add(24 * time.Hour), + DailyWindowStart: &past, + WeeklyWindowStart: &past, + MonthlyWindowStart: &past, + DailyUsageUSD: 10, + } + fresh := *stale + fresh.DailyWindowStart = ¤t + fresh.WeeklyWindowStart = ¤t + fresh.MonthlyWindowStart = ¤t + fresh.DailyUsageUSD = 2 + + subscriptionRepo := &stubUserSubscriptionRepo{ + getActive: func(context.Context, int64, int64) (*service.UserSubscription, error) { + clone := *stale + return &clone, nil + }, + getByID: func(context.Context, int64) (*service.UserSubscription, error) { + clone := fresh + return &clone, nil + }, + resetDaily: func(context.Context, int64, time.Time) error { return nil }, + resetWeekly: func(context.Context, int64, time.Time) error { return nil }, + resetMonthly: func(context.Context, int64, time.Time) error { return nil }, + } + subscriptionService := service.NewSubscriptionService(nil, subscriptionRepo, nil, nil, cfg) + router := newAuthTestRouter(apiKeyService, subscriptionService, cfg) + + w := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/t", nil) + req.Header.Set("x-api-key", apiKey.Key) + router.ServeHTTP(w, req) + + require.Equal(t, http.StatusTooManyRequests, w.Code) + }) + t.Run("simple_mode_bypasses_quota_check", func(t *testing.T) { cfg := &config.Config{RunMode: config.RunModeSimple} apiKeyService := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, cfg) @@ -1210,6 +1271,7 @@ func (r *stubApiKeyRepo) GetRateLimitData(ctx context.Context, id int64) (*servi } type stubUserSubscriptionRepo struct { + getByID func(ctx context.Context, id int64) (*service.UserSubscription, error) getActive func(ctx context.Context, userID, groupID int64) (*service.UserSubscription, error) updateStatus func(ctx context.Context, subscriptionID int64, status string) error activateWindow func(ctx context.Context, id int64, start time.Time) error @@ -1258,6 +1320,9 @@ func (r *stubUserSubscriptionRepo) Create(ctx context.Context, sub *service.User } func (r *stubUserSubscriptionRepo) GetByID(ctx context.Context, id int64) (*service.UserSubscription, error) { + if r.getByID != nil { + return r.getByID(ctx, id) + } return nil, errors.New("not implemented") } @@ -1334,21 +1399,25 @@ func (r *stubUserSubscriptionRepo) ActivateWindows(ctx context.Context, id int64 return errors.New("not implemented") } -func (r *stubUserSubscriptionRepo) ResetDailyUsage(ctx context.Context, id int64, newWindowStart time.Time) error { +func (r *stubUserSubscriptionRepo) ResetUsageWindows(context.Context, int64, bool, bool, bool, time.Time) error { + return errors.New("not implemented") +} + +func (r *stubUserSubscriptionRepo) ResetDailyUsage(ctx context.Context, id int64, _ *time.Time, newWindowStart time.Time) error { if r.resetDaily != nil { return r.resetDaily(ctx, id, newWindowStart) } return errors.New("not implemented") } -func (r *stubUserSubscriptionRepo) ResetWeeklyUsage(ctx context.Context, id int64, newWindowStart time.Time) error { +func (r *stubUserSubscriptionRepo) ResetWeeklyUsage(ctx context.Context, id int64, _ *time.Time, newWindowStart time.Time) error { if r.resetWeekly != nil { return r.resetWeekly(ctx, id, newWindowStart) } return errors.New("not implemented") } -func (r *stubUserSubscriptionRepo) ResetMonthlyUsage(ctx context.Context, id int64, newWindowStart time.Time) error { +func (r *stubUserSubscriptionRepo) ResetMonthlyUsage(ctx context.Context, id int64, _ *time.Time, newWindowStart time.Time) error { if r.resetMonthly != nil { return r.resetMonthly(ctx, id, newWindowStart) } diff --git a/backend/internal/service/payment_fulfillment.go b/backend/internal/service/payment_fulfillment.go index 51ecefe250..4d442f3d1e 100644 --- a/backend/internal/service/payment_fulfillment.go +++ b/backend/internal/service/payment_fulfillment.go @@ -28,6 +28,12 @@ import ( // misconfigured to point at us, or when our orders table has been wiped). var ErrOrderNotFound = errors.New("payment order not found") +const paymentFulfillmentLeaseDuration = 5 * time.Minute + +type paymentFulfillmentLease struct { + version time.Time +} + // --- Payment Notification & Fulfillment --- func (s *PaymentService) HandlePaymentNotification(ctx context.Context, n *payment.PaymentNotification, pk string) error { @@ -188,10 +194,8 @@ func (s *PaymentService) alreadyProcessed(ctx context.Context, o *dbent.PaymentO switch cur.Status { case OrderStatusCompleted, OrderStatusRefunded: return nil - case OrderStatusFailed: + case OrderStatusFailed, OrderStatusPaid, OrderStatusRecharging: return s.executeFulfillment(ctx, o.ID) - case OrderStatusPaid, OrderStatusRecharging: - return fmt.Errorf("order %d is being processed", o.ID) case OrderStatusExpired: slog.Warn("webhook payment success for expired order beyond grace period", "orderID", o.ID, @@ -231,23 +235,74 @@ func (s *PaymentService) ExecuteBalanceFulfillment(ctx context.Context, oid int6 if psIsRefundStatus(o.Status) { return infraerrors.BadRequest("INVALID_STATUS", "refund-related order cannot fulfill") } - if o.Status != OrderStatusPaid && o.Status != OrderStatusFailed { + if o.Status != OrderStatusPaid && o.Status != OrderStatusFailed && o.Status != OrderStatusRecharging { return infraerrors.BadRequest("INVALID_STATUS", "order cannot fulfill in status "+o.Status) } - c, err := s.entClient.PaymentOrder.Update().Where(paymentorder.IDEQ(oid), paymentorder.StatusIn(OrderStatusPaid, OrderStatusFailed)).SetStatus(OrderStatusRecharging).Save(ctx) + lease, err := s.acquirePaymentFulfillmentLease(ctx, o) if err != nil { - return fmt.Errorf("lock: %w", err) + return err } - if c == 0 { + if lease == nil { return nil } - if err := s.doBalance(ctx, o); err != nil { - s.markFailed(ctx, oid, err) + if err := s.doBalance(ctx, o, lease); err != nil { + s.markFailed(ctx, oid, lease, err) return err } return nil } +func (s *PaymentService) acquirePaymentFulfillmentLease(ctx context.Context, o *dbent.PaymentOrder) (*paymentFulfillmentLease, error) { + if o == nil { + return nil, infraerrors.BadRequest("INVALID_STATUS", "nil payment order") + } + + now := time.Now().UTC().Truncate(time.Microsecond) + staleBefore := now.Add(-paymentFulfillmentLeaseDuration) + updated, err := s.entClient.PaymentOrder.Update(). + Where( + paymentorder.IDEQ(o.ID), + paymentorder.Or( + paymentorder.StatusIn(OrderStatusPaid, OrderStatusFailed), + paymentorder.And( + paymentorder.StatusEQ(OrderStatusRecharging), + paymentorder.UpdatedAtLTE(staleBefore), + ), + ), + ). + SetStatus(OrderStatusRecharging). + SetUpdatedAt(now). + ClearFailedAt(). + ClearFailedReason(). + Save(ctx) + if err != nil { + return nil, fmt.Errorf("acquire fulfillment lease: %w", err) + } + if updated == 0 { + current, getErr := s.entClient.PaymentOrder.Get(ctx, o.ID) + if getErr != nil { + return nil, fmt.Errorf("reload fulfillment lease: %w", getErr) + } + if current.Status == OrderStatusCompleted { + return nil, nil + } + if current.Status == OrderStatusRecharging { + return nil, infraerrors.Conflict("CONFLICT", "order is being processed") + } + return nil, infraerrors.Conflict("CONFLICT", "order status changed while acquiring fulfillment lease") + } + + // Reload the persisted timestamp instead of trusting application clock precision. + claimed, err := s.entClient.PaymentOrder.Get(ctx, o.ID) + if err != nil { + return nil, fmt.Errorf("reload acquired fulfillment lease: %w", err) + } + if claimed.Status != OrderStatusRecharging { + return nil, infraerrors.Conflict("CONFLICT", "fulfillment lease was lost") + } + return &paymentFulfillmentLease{version: claimed.UpdatedAt}, nil +} + // redeemAction represents the idempotency decision for balance fulfillment. type redeemAction int @@ -272,7 +327,7 @@ func resolveRedeemAction(existing *RedeemCode, lookupErr error) redeemAction { return redeemActionRedeem } -func (s *PaymentService) doBalance(ctx context.Context, o *dbent.PaymentOrder) error { +func (s *PaymentService) doBalance(ctx context.Context, o *dbent.PaymentOrder, lease *paymentFulfillmentLease) error { // Idempotency: check if redeem code already exists (from a previous partial run) existing, lookupErr := s.redeemService.GetByCode(ctx, o.RechargeCode) action := resolveRedeemAction(existing, lookupErr) @@ -283,7 +338,7 @@ func (s *PaymentService) doBalance(ctx context.Context, o *dbent.PaymentOrder) e return err } // Code already created and redeemed — just mark completed - return s.markCompleted(ctx, o, "RECHARGE_SUCCESS") + return s.markCompleted(ctx, o, lease, "RECHARGE_SUCCESS") case redeemActionCreate: rc := &RedeemCode{Code: o.RechargeCode, Type: RedeemTypeBalance, Value: o.Amount, Status: StatusUnused} if err := s.redeemService.CreateCode(ctx, rc); err != nil { @@ -298,21 +353,37 @@ func (s *PaymentService) doBalance(ctx context.Context, o *dbent.PaymentOrder) e if err := s.applyAffiliateRebateForOrder(ctx, o); err != nil { return err } - return s.markCompleted(ctx, o, "RECHARGE_SUCCESS") + return s.markCompleted(ctx, o, lease, "RECHARGE_SUCCESS") } -func (s *PaymentService) markCompleted(ctx context.Context, o *dbent.PaymentOrder, auditAction string) error { +func (s *PaymentService) markCompleted(ctx context.Context, o *dbent.PaymentOrder, lease *paymentFulfillmentLease, auditAction string) error { + if lease == nil { + return errors.New("missing payment fulfillment lease") + } now := time.Now() - _, err := s.entClient.PaymentOrder.Update().Where(paymentorder.IDEQ(o.ID), paymentorder.StatusEQ(OrderStatusRecharging)).SetStatus(OrderStatusCompleted).SetCompletedAt(now).Save(ctx) + updated, err := s.entClient.PaymentOrder.Update().Where( + paymentorder.IDEQ(o.ID), + paymentorder.StatusEQ(OrderStatusRecharging), + paymentorder.UpdatedAtEQ(lease.version), + ).SetStatus(OrderStatusCompleted).SetCompletedAt(now).Save(ctx) if err != nil { return fmt.Errorf("mark completed: %w", err) } - s.writeAuditLog(ctx, o.ID, auditAction, "system", map[string]any{ - "rechargeCode": o.RechargeCode, - "creditedAmount": o.Amount, - "payAmount": o.PayAmount, - }) - s.dispatchPaymentFulfillmentNotification(o, auditAction) + if updated == 0 { + current, getErr := s.entClient.PaymentOrder.Get(ctx, o.ID) + if getErr == nil && current.Status == OrderStatusCompleted { + return nil + } + return infraerrors.Conflict("CONFLICT", "fulfillment lease was lost before completion") + } + if !s.hasAuditLog(ctx, o.ID, auditAction) { + s.writeAuditLog(ctx, o.ID, auditAction, "system", map[string]any{ + "rechargeCode": o.RechargeCode, + "creditedAmount": o.Amount, + "payAmount": o.PayAmount, + }) + s.dispatchPaymentFulfillmentNotification(o, auditAction) + } return nil } @@ -404,51 +475,138 @@ func (s *PaymentService) ExecuteSubscriptionFulfillment(ctx context.Context, oid if psIsRefundStatus(o.Status) { return infraerrors.BadRequest("INVALID_STATUS", "refund-related order cannot fulfill") } - if o.Status != OrderStatusPaid && o.Status != OrderStatusFailed { + if o.Status != OrderStatusPaid && o.Status != OrderStatusFailed && o.Status != OrderStatusRecharging { return infraerrors.BadRequest("INVALID_STATUS", "order cannot fulfill in status "+o.Status) } if o.SubscriptionGroupID == nil || o.SubscriptionDays == nil { return infraerrors.BadRequest("INVALID_STATUS", "missing subscription info") } - c, err := s.entClient.PaymentOrder.Update().Where(paymentorder.IDEQ(oid), paymentorder.StatusIn(OrderStatusPaid, OrderStatusFailed)).SetStatus(OrderStatusRecharging).Save(ctx) + lease, err := s.acquirePaymentFulfillmentLease(ctx, o) if err != nil { - return fmt.Errorf("lock: %w", err) + return err } - if c == 0 { + if lease == nil { return nil } - if err := s.doSub(ctx, o); err != nil { - s.markFailed(ctx, oid, err) + if err := s.doSub(ctx, o, lease); err != nil { + s.markFailed(ctx, oid, lease, err) return err } return nil } -func (s *PaymentService) doSub(ctx context.Context, o *dbent.PaymentOrder) error { +func (s *PaymentService) doSub(ctx context.Context, o *dbent.PaymentOrder, lease *paymentFulfillmentLease) error { gid := *o.SubscriptionGroupID days := *o.SubscriptionDays g, err := s.groupRepo.GetByID(ctx, gid) if err != nil || g.Status != payment.EntityStatusActive { return fmt.Errorf("group %d no longer exists or inactive", gid) } - assigned := s.hasAuditLog(ctx, o.ID, "SUBSCRIPTION_ASSIGNED") || s.hasAuditLog(ctx, o.ID, "SUBSCRIPTION_SUCCESS") - if !assigned { - orderNote := fmt.Sprintf("payment order %d", o.ID) - _, _, err = s.subscriptionSvc.AssignOrExtendSubscription(ctx, &AssignSubscriptionInput{UserID: o.UserID, GroupID: gid, ValidityDays: days, AssignedBy: 0, Notes: orderNote}) - if err != nil { - return fmt.Errorf("assign subscription: %w", err) - } - s.writeAuditLog(ctx, o.ID, "SUBSCRIPTION_ASSIGNED", "system", map[string]any{ - "groupID": gid, - "validityDays": days, - }) - } else { - slog.Info("subscription already assigned for order, skipping", "orderID", o.ID, "groupID", gid) + if err := s.ensurePaymentSubscriptionAssigned(ctx, o, gid, days); err != nil { + return err } if err := s.applyAffiliateRebateForOrder(ctx, o); err != nil { return err } - return s.markCompleted(ctx, o, "SUBSCRIPTION_SUCCESS") + return s.markCompleted(ctx, o, lease, "SUBSCRIPTION_SUCCESS") +} + +func (s *PaymentService) ensurePaymentSubscriptionAssigned(ctx context.Context, o *dbent.PaymentOrder, groupID int64, days int) error { + if s.subscriptionSvc == nil { + return errors.New("subscription service is unavailable") + } + + tx, err := s.entClient.Tx(ctx) + if err != nil { + return fmt.Errorf("begin subscription fulfillment tx: %w", err) + } + defer func() { _ = tx.Rollback() }() + + txCtx := dbent.NewTxContext(ctx, tx) + txClient := tx.Client() + alreadyAssigned, err := hasPaymentSubscriptionAssignmentAudit(txCtx, txClient, o.ID) + if err != nil { + return fmt.Errorf("check subscription assignment audit: %w", err) + } + + recoveredFromNote := false + if !alreadyAssigned { + orderNote := paymentSubscriptionOrderNote(o.ID) + existing, lookupErr := s.subscriptionSvc.userSubRepo.GetByUserIDAndGroupID(txCtx, o.UserID, groupID) + switch { + case lookupErr == nil && existing != nil && hasPaymentSubscriptionOrderNote(existing.Notes, orderNote): + recoveredFromNote = true + case lookupErr != nil && !errors.Is(lookupErr, ErrSubscriptionNotFound): + return fmt.Errorf("check existing subscription assignment: %w", lookupErr) + default: + if _, _, err := s.subscriptionSvc.assignOrExtendSubscription(txCtx, &AssignSubscriptionInput{ + UserID: o.UserID, + GroupID: groupID, + ValidityDays: days, + AssignedBy: 0, + Notes: orderNote, + }, true); err != nil { + return fmt.Errorf("assign subscription: %w", err) + } + } + + detail, _ := json.Marshal(map[string]any{ + "groupID": groupID, + "validityDays": days, + "recoveredFromNote": recoveredFromNote, + }) + if _, err := txClient.PaymentAuditLog.Create(). + SetOrderID(strconv.FormatInt(o.ID, 10)). + SetAction("SUBSCRIPTION_ASSIGNED"). + SetDetail(string(detail)). + SetOperator("system"). + Save(txCtx); err != nil { + if dbent.IsConstraintError(err) { + _ = tx.Rollback() + claimed, checkErr := hasPaymentSubscriptionAssignmentAudit(ctx, s.entClient, o.ID) + if checkErr == nil && claimed { + return s.subscriptionSvc.invalidateSubscriptionCaches(o.UserID, groupID) + } + } + return fmt.Errorf("record subscription assignment audit: %w", err) + } + } else { + slog.Info("subscription already assigned for order, skipping", "orderID", o.ID, "groupID", groupID) + } + + if err := tx.Commit(); err != nil { + return fmt.Errorf("commit subscription fulfillment tx: %w", err) + } + // Assignment cache invalidation is deferred while this transaction is open, + // then performed synchronously against the committed subscription. + if err := s.subscriptionSvc.invalidateSubscriptionCaches(o.UserID, groupID); err != nil { + return fmt.Errorf("invalidate subscription cache after fulfillment: %w", err) + } + return nil +} + +func hasPaymentSubscriptionAssignmentAudit(ctx context.Context, client *dbent.Client, orderID int64) (bool, error) { + count, err := client.PaymentAuditLog.Query(). + Where( + paymentauditlog.OrderIDEQ(strconv.FormatInt(orderID, 10)), + paymentauditlog.ActionIn("SUBSCRIPTION_ASSIGNED", "SUBSCRIPTION_SUCCESS"), + ). + Limit(1). + Count(ctx) + return count > 0, err +} + +func paymentSubscriptionOrderNote(orderID int64) string { + return fmt.Sprintf("payment order %d", orderID) +} + +func hasPaymentSubscriptionOrderNote(notes string, orderNote string) bool { + for _, line := range strings.Split(strings.ReplaceAll(notes, "\r\n", "\n"), "\n") { + if strings.TrimSpace(line) == orderNote { + return true + } + } + return false } func (s *PaymentService) hasAuditLog(ctx context.Context, orderID int64, action string) bool { @@ -642,13 +800,20 @@ func (s *PaymentService) updateClaimedAffiliateRebateAudit(ctx context.Context, return nil } -func (s *PaymentService) markFailed(ctx context.Context, oid int64, cause error) { +func (s *PaymentService) markFailed(ctx context.Context, oid int64, lease *paymentFulfillmentLease, cause error) { + if lease == nil { + slog.Error("mark FAILED without fulfillment lease", "orderID", oid) + return + } now := time.Now() r := psErrMsg(cause) - // Only mark FAILED if still in RECHARGING state — prevents overwriting - // a COMPLETED order when markCompleted failed but fulfillment succeeded. + // The lease version prevents a stale worker from overwriting a newer owner. c, e := s.entClient.PaymentOrder.Update(). - Where(paymentorder.IDEQ(oid), paymentorder.StatusEQ(OrderStatusRecharging)). + Where( + paymentorder.IDEQ(oid), + paymentorder.StatusEQ(OrderStatusRecharging), + paymentorder.UpdatedAtEQ(lease.version), + ). SetStatus(OrderStatusFailed).SetFailedAt(now).SetFailedReason(r).Save(ctx) if e != nil { slog.Error("mark FAILED", "orderID", oid, "error", e) @@ -669,18 +834,11 @@ func (s *PaymentService) RetryFulfillment(ctx context.Context, oid int64) error if psIsRefundStatus(o.Status) { return infraerrors.BadRequest("INVALID_STATUS", "refund-related order cannot retry") } - if o.Status == OrderStatusRecharging { - return infraerrors.Conflict("CONFLICT", "order is being processed") - } if o.Status == OrderStatusCompleted { return infraerrors.BadRequest("INVALID_STATUS", "order already completed") } - if o.Status != OrderStatusFailed && o.Status != OrderStatusPaid { - return infraerrors.BadRequest("INVALID_STATUS", "only paid and failed orders can retry") - } - _, err = s.entClient.PaymentOrder.Update().Where(paymentorder.IDEQ(oid), paymentorder.StatusIn(OrderStatusFailed, OrderStatusPaid)).SetStatus(OrderStatusPaid).ClearFailedAt().ClearFailedReason().Save(ctx) - if err != nil { - return fmt.Errorf("reset for retry: %w", err) + if o.Status != OrderStatusFailed && o.Status != OrderStatusPaid && o.Status != OrderStatusRecharging { + return infraerrors.BadRequest("INVALID_STATUS", "only paid, failed, and recoverable recharging orders can retry") } s.writeAuditLog(ctx, oid, "RECHARGE_RETRY", "admin", map[string]any{"detail": "admin manual retry"}) return s.executeFulfillment(ctx, oid) diff --git a/backend/internal/service/payment_fulfillment_test.go b/backend/internal/service/payment_fulfillment_test.go index a8c78d713c..d040095d63 100644 --- a/backend/internal/service/payment_fulfillment_test.go +++ b/backend/internal/service/payment_fulfillment_test.go @@ -13,6 +13,7 @@ import ( dbent "github.com/Wei-Shaw/sub2api/ent" "github.com/Wei-Shaw/sub2api/ent/paymentauditlog" "github.com/Wei-Shaw/sub2api/internal/payment" + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -586,6 +587,238 @@ func TestPaymentAmountToleranceForThreeDecimalCurrency(t *testing.T) { assert.InDelta(t, 0.0005, paymentAmountToleranceForCurrency("KWD"), 1e-12) } +func TestRetryFulfillmentRejectsFreshRechargingLease(t *testing.T) { + ctx := context.Background() + client := newPaymentConfigServiceTestClient(t) + order := createPaymentFulfillmentSubscriptionOrder(t, ctx, client, OrderStatusRecharging, time.Now()) + + svc := &PaymentService{entClient: client} + err := svc.RetryFulfillment(ctx, order.ID) + require.Error(t, err) + require.Equal(t, "CONFLICT", infraerrors.Reason(err)) + + reloaded, getErr := client.PaymentOrder.Get(ctx, order.ID) + require.NoError(t, getErr) + require.Equal(t, OrderStatusRecharging, reloaded.Status) +} + +func TestAlreadyProcessedRecoversStaleRechargingLease(t *testing.T) { + ctx := context.Background() + client := newPaymentConfigServiceTestClient(t) + ensurePaymentAuditOrderActionUniqueIndex(t, ctx, client) + order := createPaymentFulfillmentSubscriptionOrder( + t, + ctx, + client, + OrderStatusRecharging, + time.Now().Add(-paymentFulfillmentLeaseDuration-time.Minute), + ) + _, err := client.PaymentAuditLog.Create(). + SetOrderID(strconv.FormatInt(order.ID, 10)). + SetAction("SUBSCRIPTION_ASSIGNED"). + SetDetail(`{"groupID":7,"validityDays":30}`). + SetOperator("system"). + Save(ctx) + require.NoError(t, err) + + groupRepo := &subscriptionGroupRepoStub{ + group: &Group{ID: 7, Status: payment.EntityStatusActive, SubscriptionType: SubscriptionTypeSubscription}, + } + svc := &PaymentService{ + entClient: client, + groupRepo: groupRepo, + subscriptionSvc: NewSubscriptionService(groupRepo, userSubRepoNoop{}, nil, nil, nil), + } + + require.NoError(t, svc.alreadyProcessed(ctx, order)) + reloaded, err := client.PaymentOrder.Get(ctx, order.ID) + require.NoError(t, err) + require.Equal(t, OrderStatusCompleted, reloaded.Status) +} + +func TestFulfillmentLeaseVersionRejectsStaleWorker(t *testing.T) { + ctx := context.Background() + client := newPaymentConfigServiceTestClient(t) + staleAt := time.Now().Add(-paymentFulfillmentLeaseDuration - time.Minute) + order := createPaymentFulfillmentSubscriptionOrder(t, ctx, client, OrderStatusRecharging, staleAt) + svc := &PaymentService{entClient: client} + + firstLease, err := svc.acquirePaymentFulfillmentLease(ctx, order) + require.NoError(t, err) + require.NotNil(t, firstLease) + + _, err = client.PaymentOrder.UpdateOneID(order.ID).SetUpdatedAt(staleAt).Save(ctx) + require.NoError(t, err) + time.Sleep(time.Millisecond) + staleOrder, err := client.PaymentOrder.Get(ctx, order.ID) + require.NoError(t, err) + secondLease, err := svc.acquirePaymentFulfillmentLease(ctx, staleOrder) + require.NoError(t, err) + require.NotNil(t, secondLease) + require.False(t, firstLease.version.Equal(secondLease.version)) + + err = svc.markCompleted(ctx, order, firstLease, "SUBSCRIPTION_SUCCESS") + require.Error(t, err) + require.Equal(t, "CONFLICT", infraerrors.Reason(err)) + svc.markFailed(ctx, order.ID, firstLease, errors.New("stale worker failure")) + + reloaded, err := client.PaymentOrder.Get(ctx, order.ID) + require.NoError(t, err) + require.Equal(t, OrderStatusRecharging, reloaded.Status) + require.NoError(t, svc.markCompleted(ctx, order, secondLease, "SUBSCRIPTION_SUCCESS")) +} + +func TestExecuteBalanceFulfillmentRecoversAfterRedeemWithoutCreditingAgain(t *testing.T) { + ctx := context.Background() + client := newPaymentConfigServiceTestClient(t) + ensurePaymentAuditOrderActionUniqueIndex(t, ctx, client) + staleAt := time.Now().Add(-paymentFulfillmentLeaseDuration - time.Minute) + order := createPaymentFulfillmentSubscriptionOrder(t, ctx, client, OrderStatusRecharging, staleAt) + order, err := client.PaymentOrder.UpdateOneID(order.ID). + SetOrderType(payment.OrderTypeBalance). + ClearPlanID(). + ClearSubscriptionGroupID(). + ClearSubscriptionDays(). + SetUpdatedAt(staleAt). + Save(ctx) + require.NoError(t, err) + + redeemRepo := &redeemCodeRepoStub{codesByCode: map[string]*RedeemCode{ + order.RechargeCode: { + ID: 101, + Code: order.RechargeCode, + Type: RedeemTypeBalance, + Value: order.Amount, + Status: StatusUsed, + }, + }} + svc := &PaymentService{ + entClient: client, + redeemService: &RedeemService{redeemRepo: redeemRepo}, + } + + require.NoError(t, svc.ExecuteBalanceFulfillment(ctx, order.ID)) + require.Empty(t, redeemRepo.useCalls, "an already-used order code must not be redeemed again") + reloaded, err := client.PaymentOrder.Get(ctx, order.ID) + require.NoError(t, err) + require.Equal(t, OrderStatusCompleted, reloaded.Status) +} + +func TestExecuteSubscriptionFulfillmentRecoversCommittedAssignmentWithoutExtendingAgain(t *testing.T) { + ctx := context.Background() + client := newPaymentConfigServiceTestClient(t) + ensurePaymentAuditOrderActionUniqueIndex(t, ctx, client) + staleAt := time.Now().Add(-paymentFulfillmentLeaseDuration - time.Minute) + order := createPaymentFulfillmentSubscriptionOrder(t, ctx, client, OrderStatusRecharging, staleAt) + + expiresAt := time.Now().Add(30 * 24 * time.Hour).Truncate(time.Second) + subRepo := newSubscriptionUserSubRepoStub() + subRepo.seed(&UserSubscription{ + ID: 99, + UserID: order.UserID, + GroupID: *order.SubscriptionGroupID, + StartsAt: time.Now().Add(-time.Hour), + ExpiresAt: expiresAt, + Status: SubscriptionStatusActive, + Notes: "manual note\n" + paymentSubscriptionOrderNote(order.ID) + "\nretained note", + }) + groupRepo := &subscriptionGroupRepoStub{ + group: &Group{ID: 7, Status: payment.EntityStatusActive, SubscriptionType: SubscriptionTypeSubscription}, + } + svc := &PaymentService{ + entClient: client, + groupRepo: groupRepo, + subscriptionSvc: NewSubscriptionService(groupRepo, subRepo, nil, nil, nil), + } + + require.NoError(t, svc.ExecuteSubscriptionFulfillment(ctx, order.ID)) + assertPaymentSubscriptionExpiry(t, subRepo, order, expiresAt) + + assignmentAuditCount, err := client.PaymentAuditLog.Query(). + Where( + paymentauditlog.OrderIDEQ(strconv.FormatInt(order.ID, 10)), + paymentauditlog.ActionEQ("SUBSCRIPTION_ASSIGNED"), + ). + Count(ctx) + require.NoError(t, err) + require.Equal(t, 1, assignmentAuditCount) + + // Simulate another stale recovery attempt after completion. The durable audit + // must make replay a no-op for the subscription entitlement. + _, err = client.PaymentOrder.UpdateOneID(order.ID). + SetStatus(OrderStatusRecharging). + SetUpdatedAt(staleAt). + ClearCompletedAt(). + Save(ctx) + require.NoError(t, err) + require.NoError(t, svc.ExecuteSubscriptionFulfillment(ctx, order.ID)) + assertPaymentSubscriptionExpiry(t, subRepo, order, expiresAt) + + assignmentAuditCount, err = client.PaymentAuditLog.Query(). + Where( + paymentauditlog.OrderIDEQ(strconv.FormatInt(order.ID, 10)), + paymentauditlog.ActionEQ("SUBSCRIPTION_ASSIGNED"), + ). + Count(ctx) + require.NoError(t, err) + require.Equal(t, 1, assignmentAuditCount) +} + +func TestHasPaymentSubscriptionOrderNoteRequiresIndependentExactLine(t *testing.T) { + t.Parallel() + require.True(t, hasPaymentSubscriptionOrderNote("before\r\npayment order 42\r\nafter", "payment order 42")) + require.False(t, hasPaymentSubscriptionOrderNote("payment order 420", "payment order 42")) + require.False(t, hasPaymentSubscriptionOrderNote("prefix payment order 42 suffix", "payment order 42")) +} + +func createPaymentFulfillmentSubscriptionOrder( + t *testing.T, + ctx context.Context, + client *dbent.Client, + status string, + updatedAt time.Time, +) *dbent.PaymentOrder { + t.Helper() + user, err := client.User.Create(). + SetEmail("fulfillment-" + strconv.FormatInt(time.Now().UnixNano(), 10) + "@example.com"). + SetPasswordHash("hash"). + SetUsername("payment-fulfillment-user"). + Save(ctx) + require.NoError(t, err) + + order, err := client.PaymentOrder.Create(). + SetUserID(user.ID). + SetUserEmail(user.Email). + SetUserName(user.Username). + SetAmount(80). + SetPayAmount(80). + SetFeeRate(0). + SetRechargeCode("PAY-SUB-" + strconv.FormatInt(time.Now().UnixNano(), 10)). + SetOutTradeNo("sub2_fulfillment_" + strconv.FormatInt(time.Now().UnixNano(), 10)). + SetPaymentType(payment.TypeAlipay). + SetPaymentTradeNo("trade-fulfillment"). + SetOrderType(payment.OrderTypeSubscription). + SetPlanID(100). + SetSubscriptionGroupID(7). + SetSubscriptionDays(30). + SetStatus(status). + SetPaidAt(time.Now().Add(-time.Hour)). + SetExpiresAt(time.Now().Add(time.Hour)). + SetClientIP("127.0.0.1"). + SetSrcHost("api.example.com"). + SetUpdatedAt(updatedAt). + Save(ctx) + require.NoError(t, err) + return order +} + +func assertPaymentSubscriptionExpiry(t *testing.T, repo *subscriptionUserSubRepoStub, order *dbent.PaymentOrder, expected time.Time) { + t.Helper() + sub, err := repo.GetByUserIDAndGroupID(context.Background(), order.UserID, *order.SubscriptionGroupID) + require.NoError(t, err) + require.True(t, sub.ExpiresAt.Equal(expected), "subscription expiry changed from %s to %s", expected, sub.ExpiresAt) +} + func TestExecuteSubscriptionFulfillmentAppliesAffiliateRebate(t *testing.T) { ctx := context.Background() client := newPaymentConfigServiceTestClient(t) diff --git a/backend/internal/service/redeem_service.go b/backend/internal/service/redeem_service.go index 2d1962dd3c..8794872e3d 100644 --- a/backend/internal/service/redeem_service.go +++ b/backend/internal/service/redeem_service.go @@ -135,6 +135,7 @@ type RedeemCodeBatchUpdateResult struct { type RedeemService struct { redeemRepo RedeemCodeRepository userRepo UserRepository + redeemUserRepo RedeemUserAdjustmentRepository subscriptionService *SubscriptionService cache RedeemCache billingCacheService *BillingCacheService @@ -154,9 +155,11 @@ func NewRedeemService( authCacheInvalidator APIKeyAuthCacheInvalidator, affiliateService *AffiliateService, ) *RedeemService { + redeemUserRepo, _ := userRepo.(RedeemUserAdjustmentRepository) return &RedeemService{ redeemRepo: redeemRepo, userRepo: userRepo, + redeemUserRepo: redeemUserRepo, subscriptionService: subscriptionService, cache: cache, billingCacheService: billingCacheService, @@ -426,7 +429,7 @@ func (s *RedeemService) Redeem(ctx context.Context, userID int64, code string) ( } // 获取用户信息 - user, err := s.userRepo.GetByID(ctx, userID) + _, err = s.userRepo.GetByID(ctx, userID) if err != nil { return nil, fmt.Errorf("get user: %w", err) } @@ -454,21 +457,27 @@ func (s *RedeemService) Redeem(ctx context.Context, userID int64, code string) ( switch redeemCode.Type { case RedeemTypeBalance: amount := redeemCode.Value - // 负数为退款扣减,余额最低为 0 - if amount < 0 && user.Balance+amount < 0 { - amount = -user.Balance - } - if err := s.userRepo.UpdateBalance(txCtx, userID, amount); err != nil { + if amount < 0 { + if s.redeemUserRepo == nil { + return nil, errors.New("user repository does not support atomic redeem balance adjustments") + } + if err := s.redeemUserRepo.ApplyRedeemBalanceAdjustment(txCtx, userID, amount); err != nil { + return nil, fmt.Errorf("update user balance: %w", err) + } + } else if err := s.userRepo.UpdateBalance(txCtx, userID, amount); err != nil { return nil, fmt.Errorf("update user balance: %w", err) } case RedeemTypeConcurrency: delta := int(redeemCode.Value) - // 负数为退款扣减,并发数最低为 0 - if delta < 0 && user.Concurrency+delta < 0 { - delta = -user.Concurrency - } - if err := s.userRepo.UpdateConcurrency(txCtx, userID, delta); err != nil { + if delta < 0 { + if s.redeemUserRepo == nil { + return nil, errors.New("user repository does not support atomic redeem concurrency adjustments") + } + if err := s.redeemUserRepo.ApplyRedeemConcurrencyAdjustment(txCtx, userID, delta); err != nil { + return nil, fmt.Errorf("update user concurrency: %w", err) + } + } else if err := s.userRepo.UpdateConcurrency(txCtx, userID, delta); err != nil { return nil, fmt.Errorf("update user concurrency: %w", err) } diff --git a/backend/internal/service/subscription_assign_idempotency_test.go b/backend/internal/service/subscription_assign_idempotency_test.go index 8e249af52b..d4913f8994 100644 --- a/backend/internal/service/subscription_assign_idempotency_test.go +++ b/backend/internal/service/subscription_assign_idempotency_test.go @@ -6,11 +6,49 @@ import ( "testing" "time" + dbent "github.com/Wei-Shaw/sub2api/ent" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/dgraph-io/ristretto" "github.com/stretchr/testify/require" ) +func TestWithSubscriptionUpdateTx_ReusesExistingTransaction(t *testing.T) { + existingTx := &dbent.Tx{} + ctx := dbent.NewTxContext(context.Background(), existingTx) + svc := &SubscriptionService{entClient: &dbent.Client{}} + + called := false + err := svc.withSubscriptionUpdateTx(ctx, func(txCtx context.Context) error { + called = true + require.Same(t, existingTx, dbent.TxFromContext(txCtx)) + return nil + }) + + require.NoError(t, err) + require.True(t, called) +} + +func TestMaybeInvalidateAssignmentCaches_DefersForOuterTransactionOwner(t *testing.T) { + cache, err := ristretto.NewCache(&ristretto.Config{NumCounters: 1_000, MaxCost: 100, BufferItems: 64}) + require.NoError(t, err) + t.Cleanup(cache.Close) + + svc := &SubscriptionService{subCacheL1: cache} + key := subCacheKey(7, 9) + require.True(t, cache.Set(key, &UserSubscription{ID: 42}, 1)) + cache.Wait() + + svc.maybeInvalidateAssignmentCaches(7, 9, true) + _, cachedBeforeCommit := cache.Get(key) + require.True(t, cachedBeforeCommit, "outer transaction must retain caches until its owner commits") + + svc.maybeInvalidateAssignmentCaches(7, 9, false) + cache.Wait() + _, cachedAfterCommit := cache.Get(key) + require.False(t, cachedAfterCommit, "post-commit invalidation must remove the cached subscription") +} + type groupRepoNoop struct{} func (groupRepoNoop) Create(context.Context, *Group) error { panic("unexpected Create call") } @@ -119,13 +157,16 @@ func (userSubRepoNoop) UpdateNotes(context.Context, int64, string) error { func (userSubRepoNoop) ActivateWindows(context.Context, int64, time.Time) error { panic("unexpected ActivateWindows call") } -func (userSubRepoNoop) ResetDailyUsage(context.Context, int64, time.Time) error { +func (userSubRepoNoop) ResetUsageWindows(context.Context, int64, bool, bool, bool, time.Time) error { + panic("unexpected ResetUsageWindows call") +} +func (userSubRepoNoop) ResetDailyUsage(context.Context, int64, *time.Time, time.Time) error { panic("unexpected ResetDailyUsage call") } -func (userSubRepoNoop) ResetWeeklyUsage(context.Context, int64, time.Time) error { +func (userSubRepoNoop) ResetWeeklyUsage(context.Context, int64, *time.Time, time.Time) error { panic("unexpected ResetWeeklyUsage call") } -func (userSubRepoNoop) ResetMonthlyUsage(context.Context, int64, time.Time) error { +func (userSubRepoNoop) ResetMonthlyUsage(context.Context, int64, *time.Time, time.Time) error { panic("unexpected ResetMonthlyUsage call") } func (userSubRepoNoop) IncrementUsage(context.Context, int64, float64) error { diff --git a/backend/internal/service/subscription_expiry_service_test.go b/backend/internal/service/subscription_expiry_service_test.go index 056315a289..7db642c076 100644 --- a/backend/internal/service/subscription_expiry_service_test.go +++ b/backend/internal/service/subscription_expiry_service_test.go @@ -87,15 +87,19 @@ func (r *subscriptionExpiryRepoStub) ActivateWindows(context.Context, int64, tim return nil } -func (r *subscriptionExpiryRepoStub) ResetDailyUsage(context.Context, int64, time.Time) error { +func (r *subscriptionExpiryRepoStub) ResetUsageWindows(context.Context, int64, bool, bool, bool, time.Time) error { return nil } -func (r *subscriptionExpiryRepoStub) ResetWeeklyUsage(context.Context, int64, time.Time) error { +func (r *subscriptionExpiryRepoStub) ResetDailyUsage(context.Context, int64, *time.Time, time.Time) error { return nil } -func (r *subscriptionExpiryRepoStub) ResetMonthlyUsage(context.Context, int64, time.Time) error { +func (r *subscriptionExpiryRepoStub) ResetWeeklyUsage(context.Context, int64, *time.Time, time.Time) error { + return nil +} + +func (r *subscriptionExpiryRepoStub) ResetMonthlyUsage(context.Context, int64, *time.Time, time.Time) error { return nil } diff --git a/backend/internal/service/subscription_reset_quota_test.go b/backend/internal/service/subscription_reset_quota_test.go index 3bbc217073..e4ed45ec45 100644 --- a/backend/internal/service/subscription_reset_quota_test.go +++ b/backend/internal/service/subscription_reset_quota_test.go @@ -11,7 +11,7 @@ import ( "github.com/stretchr/testify/require" ) -// resetQuotaUserSubRepoStub 支持 GetByID、ResetDailyUsage、ResetWeeklyUsage、ResetMonthlyUsage, +// resetQuotaUserSubRepoStub 支持 GetByID、ResetUsageWindows, // 其余方法继承 userSubRepoNoop(panic)。 type resetQuotaUserSubRepoStub struct { userSubRepoNoop @@ -34,7 +34,38 @@ func (r *resetQuotaUserSubRepoStub) GetByID(_ context.Context, id int64) (*UserS return &cp, nil } -func (r *resetQuotaUserSubRepoStub) ResetDailyUsage(_ context.Context, _ int64, windowStart time.Time) error { +func (r *resetQuotaUserSubRepoStub) ResetUsageWindows(_ context.Context, _ int64, resetDaily, resetWeekly, resetMonthly bool, windowStart time.Time) error { + r.resetDailyCalled = resetDaily + r.resetWeeklyCalled = resetWeekly + r.resetMonthlyCalled = resetMonthly + if resetDaily && r.resetDailyErr != nil { + return r.resetDailyErr + } + if resetWeekly && r.resetWeeklyErr != nil { + return r.resetWeeklyErr + } + if resetMonthly && r.resetMonthlyErr != nil { + return r.resetMonthlyErr + } + if r.sub == nil { + return nil + } + if resetDaily { + r.sub.DailyUsageUSD = 0 + r.sub.DailyWindowStart = &windowStart + } + if resetWeekly { + r.sub.WeeklyUsageUSD = 0 + r.sub.WeeklyWindowStart = &windowStart + } + if resetMonthly { + r.sub.MonthlyUsageUSD = 0 + r.sub.MonthlyWindowStart = &windowStart + } + return nil +} + +func (r *resetQuotaUserSubRepoStub) ResetDailyUsage(_ context.Context, _ int64, _ *time.Time, windowStart time.Time) error { r.resetDailyCalled = true if r.resetDailyErr == nil && r.sub != nil { r.sub.DailyUsageUSD = 0 @@ -43,12 +74,12 @@ func (r *resetQuotaUserSubRepoStub) ResetDailyUsage(_ context.Context, _ int64, return r.resetDailyErr } -func (r *resetQuotaUserSubRepoStub) ResetWeeklyUsage(_ context.Context, _ int64, _ time.Time) error { +func (r *resetQuotaUserSubRepoStub) ResetWeeklyUsage(_ context.Context, _ int64, _ *time.Time, _ time.Time) error { r.resetWeeklyCalled = true return r.resetWeeklyErr } -func (r *resetQuotaUserSubRepoStub) ResetMonthlyUsage(_ context.Context, _ int64, _ time.Time) error { +func (r *resetQuotaUserSubRepoStub) ResetMonthlyUsage(_ context.Context, _ int64, _ *time.Time, _ time.Time) error { r.resetMonthlyCalled = true return r.resetMonthlyErr } @@ -140,7 +171,7 @@ func TestAdminResetQuota_ResetDailyUsageError(t *testing.T) { require.ErrorIs(t, err, dbErr) require.True(t, stub.resetDailyCalled) - require.False(t, stub.resetWeeklyCalled, "daily 失败后不应继续调用 weekly") + require.True(t, stub.resetWeeklyCalled, "原子重置应在一次调用中提交所选窗口") } func TestAdminResetQuota_ResetWeeklyUsageError(t *testing.T) { @@ -200,7 +231,7 @@ func TestAdminResetQuota_ReturnsRefreshedSub(t *testing.T) { result, err := svc.AdminResetQuota(context.Background(), 6, true, false, false) require.NoError(t, err) - // ResetDailyUsage stub 会将 sub.DailyUsageUSD 归零, + // ResetUsageWindows stub 会将 sub.DailyUsageUSD 归零, // 服务应返回第二次 GetByID 的刷新值而非初始的 99.9 require.Equal(t, float64(0), result.DailyUsageUSD, "返回的订阅应反映已归零的用量") require.True(t, stub.resetDailyCalled) diff --git a/backend/internal/service/subscription_service.go b/backend/internal/service/subscription_service.go index 0a4fc7b757..ea1fd091d9 100644 --- a/backend/internal/service/subscription_service.go +++ b/backend/internal/service/subscription_service.go @@ -212,6 +212,10 @@ func (s *SubscriptionService) AssignSubscription(ctx context.Context, input *Ass // // 如果没有订阅:创建新订阅 func (s *SubscriptionService) AssignOrExtendSubscription(ctx context.Context, input *AssignSubscriptionInput) (*UserSubscription, bool, error) { + return s.assignOrExtendSubscription(ctx, input, false) +} + +func (s *SubscriptionService) assignOrExtendSubscription(ctx context.Context, input *AssignSubscriptionInput, deferCacheInvalidation bool) (*UserSubscription, bool, error) { // 检查分组是否存在且为订阅类型 group, err := s.groupRepo.GetByID(ctx, input.GroupID) if err != nil { @@ -260,15 +264,7 @@ func (s *SubscriptionService) AssignOrExtendSubscription(ctx context.Context, in } // 失效订阅缓存 - s.InvalidateSubCache(input.UserID, input.GroupID) - if s.billingCacheService != nil { - userID, groupID := input.UserID, input.GroupID - go func() { - cacheCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - _ = s.billingCacheService.InvalidateSubscription(cacheCtx, userID, groupID) - }() - } + s.maybeInvalidateAssignmentCaches(input.UserID, input.GroupID, deferCacheInvalidation) // 返回更新后的订阅 sub, err := s.userSubRepo.GetByID(ctx, existingSub.ID) @@ -282,17 +278,27 @@ func (s *SubscriptionService) AssignOrExtendSubscription(ctx context.Context, in } // 失效订阅缓存 - s.InvalidateSubCache(input.UserID, input.GroupID) + s.maybeInvalidateAssignmentCaches(input.UserID, input.GroupID, deferCacheInvalidation) + + return sub, false, nil // false 表示是新建 +} + +func (s *SubscriptionService) maybeInvalidateAssignmentCaches(userID, groupID int64, deferred bool) { + // Payment fulfillment owns an outer transaction and performs a synchronous + // invalidation after commit. Invalidating inside that transaction can reload + // the pre-commit subscription into cache. + if deferred { + return + } + + s.InvalidateSubCache(userID, groupID) if s.billingCacheService != nil { - userID, groupID := input.UserID, input.GroupID go func() { cacheCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() _ = s.billingCacheService.InvalidateSubscription(cacheCtx, userID, groupID) }() } - - return sub, false, nil // false 表示是新建 } func (s *SubscriptionService) updateExistingSubscriptionTerm( @@ -336,6 +342,9 @@ func (s *SubscriptionService) updateExistingSubscriptionTerm( } func (s *SubscriptionService) withSubscriptionUpdateTx(ctx context.Context, fn func(context.Context) error) error { + if dbent.TxFromContext(ctx) != nil { + return fn(ctx) + } if s.entClient == nil { return fn(ctx) } @@ -834,20 +843,8 @@ func (s *SubscriptionService) AdminResetQuota(ctx context.Context, subscriptionI return nil, err } windowStart := startOfDay(time.Now()) - if resetDaily { - if err := s.userSubRepo.ResetDailyUsage(ctx, sub.ID, windowStart); err != nil { - return nil, err - } - } - if resetWeekly { - if err := s.userSubRepo.ResetWeeklyUsage(ctx, sub.ID, windowStart); err != nil { - return nil, err - } - } - if resetMonthly { - if err := s.userSubRepo.ResetMonthlyUsage(ctx, sub.ID, windowStart); err != nil { - return nil, err - } + if err := s.userSubRepo.ResetUsageWindows(ctx, sub.ID, resetDaily, resetWeekly, resetMonthly, windowStart); err != nil { + return nil, err } // Invalidate L1 ristretto cache. Ristretto's Del() is asynchronous by design, // so call Wait() immediately after to flush pending operations and guarantee @@ -868,7 +865,8 @@ func (s *SubscriptionService) CheckAndResetWindows(ctx context.Context, sub *Use // 日窗口重置(24小时) if sub.NeedsDailyReset() { - if err := s.userSubRepo.ResetDailyUsage(ctx, sub.ID, windowStart); err != nil { + expectedWindowStart := sub.DailyWindowStart + if err := s.userSubRepo.ResetDailyUsage(ctx, sub.ID, expectedWindowStart, windowStart); err != nil { return err } sub.DailyWindowStart = &windowStart @@ -878,7 +876,8 @@ func (s *SubscriptionService) CheckAndResetWindows(ctx context.Context, sub *Use // 周窗口重置(7天) if sub.NeedsWeeklyReset() { - if err := s.userSubRepo.ResetWeeklyUsage(ctx, sub.ID, windowStart); err != nil { + expectedWindowStart := sub.WeeklyWindowStart + if err := s.userSubRepo.ResetWeeklyUsage(ctx, sub.ID, expectedWindowStart, windowStart); err != nil { return err } sub.WeeklyWindowStart = &windowStart @@ -888,7 +887,8 @@ func (s *SubscriptionService) CheckAndResetWindows(ctx context.Context, sub *Use // 月窗口重置(30天) if sub.NeedsMonthlyReset() { - if err := s.userSubRepo.ResetMonthlyUsage(ctx, sub.ID, windowStart); err != nil { + expectedWindowStart := sub.MonthlyWindowStart + if err := s.userSubRepo.ResetMonthlyUsage(ctx, sub.ID, expectedWindowStart, windowStart); err != nil { return err } sub.MonthlyWindowStart = &windowStart @@ -907,6 +907,32 @@ func (s *SubscriptionService) CheckAndResetWindows(ctx context.Context, sub *Use return nil } +// EnsureWindowMaintenance advances expired usage windows before a request is +// allowed to proceed. It returns a fresh database snapshot because a competing +// request may have won one of the conditional resets. +func (s *SubscriptionService) EnsureWindowMaintenance(ctx context.Context, sub *UserSubscription) (*UserSubscription, error) { + if sub == nil { + return nil, ErrSubscriptionNilInput + } + if !sub.IsWindowActivated() { + if err := s.CheckAndActivateWindow(ctx, sub); err != nil { + return nil, err + } + } + if err := s.CheckAndResetWindows(ctx, sub); err != nil { + return nil, err + } + + // GetByID bypasses the service caches. This prevents a stale loser of the + // CAS from validating limits against zeroed in-memory usage. + refreshed, err := s.userSubRepo.GetByID(ctx, sub.ID) + if err != nil { + return nil, err + } + s.InvalidateSubCacheSync(sub.UserID, sub.GroupID) + return refreshed, nil +} + // CheckUsageLimits 检查使用限额(返回错误如果超限) // 用于中间件的快速预检查,additionalCost 通常为 0 func (s *SubscriptionService) CheckUsageLimits(ctx context.Context, sub *UserSubscription, group *Group, additionalCost float64) error { @@ -923,8 +949,8 @@ func (s *SubscriptionService) CheckUsageLimits(ctx context.Context, sub *UserSub } // ValidateAndCheckLimits 合并验证+限额检查(中间件热路径专用) -// 仅做内存检查,不触发 DB 写入。窗口重置的 DB 写入由 DoWindowMaintenance 异步完成。 -// 返回 needsMaintenance 表示是否需要异步执行窗口维护。 +// 仅做内存检查,不触发 DB 写入。调用方必须在放行请求前同步完成窗口维护。 +// 返回 needsMaintenance 表示是否需要执行窗口维护并回读数据库快照。 func (s *SubscriptionService) ValidateAndCheckLimits(sub *UserSubscription, group *Group) (needsMaintenance bool, err error) { // 1. 验证订阅状态 if sub.Status == SubscriptionStatusExpired { @@ -937,8 +963,8 @@ func (s *SubscriptionService) ValidateAndCheckLimits(sub *UserSubscription, grou return false, ErrSubscriptionExpired } - // 2. 内存中修正过期窗口的用量,确保 CheckUsageLimits 不会误拒绝用户 - // 实际的 DB 窗口重置由 DoWindowMaintenance 异步完成 + // 2. 内存中修正过期窗口的用量,确保预检查不会误拒绝用户。 + // 调用方随后同步推进 DB 窗口,并用回读快照重新校验。 if sub.NeedsDailyReset() { sub.DailyUsageUSD = 0 needsMaintenance = true diff --git a/backend/internal/service/user_service.go b/backend/internal/service/user_service.go index 2c87221401..98b0c8f32b 100644 --- a/backend/internal/service/user_service.go +++ b/backend/internal/service/user_service.go @@ -122,6 +122,14 @@ type UserRepository interface { DisableTotp(ctx context.Context, userID int64) 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. +type RedeemUserAdjustmentRepository interface { + ApplyRedeemBalanceAdjustment(ctx context.Context, id int64, delta float64) error + ApplyRedeemConcurrencyAdjustment(ctx context.Context, id int64, delta int) error +} + type UserAuthIdentityRecord struct { ProviderType string ProviderKey string diff --git a/backend/internal/service/user_subscription_daily_quota_test.go b/backend/internal/service/user_subscription_daily_quota_test.go index 3738bdd698..bf58de7f4c 100644 --- a/backend/internal/service/user_subscription_daily_quota_test.go +++ b/backend/internal/service/user_subscription_daily_quota_test.go @@ -15,7 +15,7 @@ type dailyResetTrackingUserSubRepo struct { resetDailyCalled bool } -func (r *dailyResetTrackingUserSubRepo) ResetDailyUsage(context.Context, int64, time.Time) error { +func (r *dailyResetTrackingUserSubRepo) ResetDailyUsage(context.Context, int64, *time.Time, time.Time) error { r.resetDailyCalled = true return nil } diff --git a/backend/internal/service/user_subscription_port.go b/backend/internal/service/user_subscription_port.go index 43d41d6dd7..eeee0275f0 100644 --- a/backend/internal/service/user_subscription_port.go +++ b/backend/internal/service/user_subscription_port.go @@ -29,9 +29,10 @@ type UserSubscriptionRepository interface { UpdateNotes(ctx context.Context, subscriptionID int64, notes string) error ActivateWindows(ctx context.Context, id int64, start time.Time) error - ResetDailyUsage(ctx context.Context, id int64, newWindowStart time.Time) error - ResetWeeklyUsage(ctx context.Context, id int64, newWindowStart time.Time) error - ResetMonthlyUsage(ctx context.Context, id int64, newWindowStart time.Time) error + ResetUsageWindows(ctx context.Context, id int64, resetDaily, resetWeekly, resetMonthly bool, newWindowStart time.Time) error + ResetDailyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error + ResetWeeklyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error + ResetMonthlyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error IncrementUsage(ctx context.Context, id int64, costUSD float64) error BatchUpdateExpiredStatus(ctx context.Context) (int64, error) diff --git a/frontend/src/router/__tests__/feature-access.spec.ts b/frontend/src/router/__tests__/feature-access.spec.ts new file mode 100644 index 0000000000..3c98425f90 --- /dev/null +++ b/frontend/src/router/__tests__/feature-access.spec.ts @@ -0,0 +1,177 @@ +import { beforeAll, beforeEach, describe, expect, it, vi } from 'vitest' + +type NavigationGuard = ( + to: Record, + from: Record, + next: ReturnType +) => Promise + +const routerHarness = vi.hoisted(() => ({ + guard: null as NavigationGuard | null, +})) + +const authStore = vi.hoisted(() => ({ + checkAuth: vi.fn(), + isAuthenticated: true, + isAdmin: false, + isSimpleMode: false, + hasPendingAuthSession: false, +})) + +const appStore = vi.hoisted(() => ({ + siteName: 'Sub2API', + backendModeEnabled: false, + publicSettingsLoaded: false, + cachedPublicSettings: null as null | { + payment_enabled?: boolean + risk_control_enabled?: boolean + custom_menu_items?: [] + }, + fetchPublicSettings: vi.fn(), +})) + +vi.mock('vue-router', () => ({ + createWebHistory: vi.fn(() => ({})), + createRouter: vi.fn(() => ({ + beforeEach: vi.fn((guard: NavigationGuard) => { + routerHarness.guard = guard + }), + afterEach: vi.fn(), + onError: vi.fn(), + })), +})) + +vi.mock('@/stores/auth', () => ({ + useAuthStore: () => authStore, +})) + +vi.mock('@/stores/app', () => ({ + useAppStore: () => appStore, +})) + +vi.mock('@/stores/adminSettings', () => ({ + useAdminSettingsStore: () => ({ customMenuItems: [] }), +})) + +vi.mock('@/stores/adminCompliance', () => ({ + useAdminComplianceStore: () => ({ + initialized: true, + fetchStatus: vi.fn(), + requireAcknowledgement: vi.fn(), + }), +})) + +vi.mock('@/composables/useNavigationLoading', () => ({ + useNavigationLoadingState: () => ({ + startNavigation: vi.fn(), + endNavigation: vi.fn(), + isLoading: { value: false }, + }), +})) + +vi.mock('@/composables/useRoutePrefetch', () => ({ + useRoutePrefetch: () => ({ + triggerPrefetch: vi.fn(), + cancelPendingPrefetch: vi.fn(), + resetPrefetchState: vi.fn(), + }), +})) + +function createDeferred() { + let resolve!: (value: T | PromiseLike) => void + const promise = new Promise((resolvePromise) => { + resolve = resolvePromise + }) + return { promise, resolve } +} + +function runGuard(meta: Record, path: string) { + if (!routerHarness.guard) { + throw new Error('router guard was not registered') + } + + const next = vi.fn() + const navigation = routerHarness.guard( + { + path, + fullPath: path, + name: 'FeatureRoute', + params: {}, + meta: { requiresAuth: true, ...meta }, + }, + {}, + next + ) + return { navigation, next } +} + +describe('feature route guard', () => { + beforeAll(async () => { + await import('@/router') + }) + + beforeEach(() => { + authStore.isAuthenticated = true + authStore.isAdmin = false + authStore.isSimpleMode = false + appStore.publicSettingsLoaded = false + appStore.cachedPublicSettings = null + appStore.fetchPublicSettings.mockReset() + }) + + it('waits for the first public-settings request before deciding payment access', async () => { + const deferred = createDeferred<{ payment_enabled: boolean }>() + appStore.fetchPublicSettings.mockImplementation(async () => { + const settings = await deferred.promise + appStore.cachedPublicSettings = settings + appStore.publicSettingsLoaded = true + return settings + }) + + const { navigation, next } = runGuard({ requiresPayment: true }, '/purchase') + + await vi.waitFor(() => expect(appStore.fetchPublicSettings).toHaveBeenCalledTimes(1)) + expect(next).not.toHaveBeenCalled() + + deferred.resolve({ payment_enabled: true }) + await navigation + expect(next).toHaveBeenCalledOnce() + expect(next).toHaveBeenCalledWith() + }) + + it.each([ + ['payment', { requiresPayment: true }, '/purchase'], + ['risk control', { requiresRiskControl: true }, '/admin/risk-control'], + ])('does not treat a failed %s settings load as explicitly disabled', async (_name, meta, path) => { + authStore.isAdmin = meta.requiresRiskControl === true + appStore.fetchPublicSettings.mockResolvedValue(null) + + const { navigation, next } = runGuard(meta, path) + await navigation + + expect(appStore.publicSettingsLoaded).toBe(false) + expect(next).toHaveBeenCalledOnce() + expect(next).toHaveBeenCalledWith() + }) + + it.each([ + ['payment', { requiresPayment: true }, { payment_enabled: false }, '/dashboard'], + [ + 'risk control', + { requiresRiskControl: true }, + { risk_control_enabled: false }, + '/admin/settings', + ], + ])('redirects when loaded settings explicitly disable %s', async (_name, meta, settings, target) => { + authStore.isAdmin = meta.requiresRiskControl === true + appStore.cachedPublicSettings = settings + appStore.publicSettingsLoaded = true + + const { navigation, next } = runGuard(meta, '/feature') + await navigation + + expect(appStore.fetchPublicSettings).not.toHaveBeenCalled() + expect(next).toHaveBeenCalledOnce() + expect(next).toHaveBeenCalledWith(target) + }) +}) diff --git a/frontend/src/router/index.ts b/frontend/src/router/index.ts index 306a0eac30..e108d9d7b9 100644 --- a/frontend/src/router/index.ts +++ b/frontend/src/router/index.ts @@ -837,21 +837,24 @@ router.beforeEach(async (to, _from, next) => { } } - // Check payment requirement (internal payment system only) - if (to.meta.requiresPayment) { - const paymentEnabled = appStore.cachedPublicSettings?.payment_enabled - if (!paymentEnabled) { - next(authStore.isAdmin ? '/admin/dashboard' : '/dashboard') - return - } + // Only an explicit value from successfully loaded settings can disable a route. + // A transient settings failure is unknown state, not a confirmed feature toggle. + if ( + to.meta.requiresPayment && + appStore.publicSettingsLoaded && + appStore.cachedPublicSettings?.payment_enabled === false + ) { + next(authStore.isAdmin ? '/admin/dashboard' : '/dashboard') + return } - if (to.meta.requiresRiskControl) { - const riskControlEnabled = appStore.cachedPublicSettings?.risk_control_enabled === true - if (!riskControlEnabled) { - next(authStore.isAdmin ? '/admin/settings' : '/dashboard') - return - } + if ( + to.meta.requiresRiskControl && + appStore.publicSettingsLoaded && + appStore.cachedPublicSettings?.risk_control_enabled === false + ) { + next(authStore.isAdmin ? '/admin/settings' : '/dashboard') + return } // 简易模式下限制访问某些页面 diff --git a/frontend/src/stores/__tests__/app.spec.ts b/frontend/src/stores/__tests__/app.spec.ts index 803dad0e90..d9e7b17f57 100644 --- a/frontend/src/stores/__tests__/app.spec.ts +++ b/frontend/src/stores/__tests__/app.spec.ts @@ -2,6 +2,63 @@ import { describe, it, expect, vi, beforeEach, afterEach } from 'vitest' import { setActivePinia, createPinia } from 'pinia' import { useAppStore } from '@/stores/app' import { getPublicSettings } from '@/api/auth' +import type { PublicSettings } from '@/types' + +function createDeferred() { + let resolve!: (value: T | PromiseLike) => void + let reject!: (reason?: unknown) => void + const promise = new Promise((resolvePromise, rejectPromise) => { + resolve = resolvePromise + reject = rejectPromise + }) + + return { promise, resolve, reject } +} + +function createPublicSettings(overrides: Partial = {}): PublicSettings { + return { + registration_enabled: false, + email_verify_enabled: false, + force_email_on_third_party_signup: false, + registration_email_suffix_whitelist: [], + promo_code_enabled: true, + password_reset_enabled: false, + invitation_code_enabled: false, + turnstile_enabled: false, + turnstile_site_key: '', + site_name: 'Test Site', + site_logo: '', + site_subtitle: '', + api_base_url: '', + contact_info: '', + doc_url: '', + home_content: '', + hide_ccs_import_button: false, + payment_enabled: false, + risk_control_enabled: false, + table_default_page_size: 20, + table_page_size_options: [10, 20, 50, 100], + custom_menu_items: [], + custom_endpoints: [], + linuxdo_oauth_enabled: false, + wechat_oauth_enabled: false, + oidc_oauth_enabled: false, + oidc_oauth_provider_name: 'OIDC', + github_oauth_enabled: false, + google_oauth_enabled: false, + backend_mode_enabled: false, + version: '1.0.0', + balance_low_notify_enabled: false, + account_quota_notify_enabled: false, + balance_low_notify_threshold: 0, + channel_monitor_enabled: true, + channel_monitor_default_interval_seconds: 60, + available_channels_enabled: false, + service_quota_enabled: false, + affiliate_enabled: false, + ...overrides, + } +} // Mock API 模块 vi.mock('@/api/admin/system', () => ({ @@ -17,6 +74,7 @@ describe('useAppStore', () => { setActivePinia(createPinia()) vi.useFakeTimers() localStorage.clear() + vi.mocked(getPublicSettings).mockReset() // 清除 window.__APP_CONFIG__ delete (window as any).__APP_CONFIG__ }) @@ -263,6 +321,75 @@ describe('useAppStore', () => { // --- 公开设置 --- describe('公开设置加载', () => { + it('并发调用复用并等待同一个请求,包括 force 调用', async () => { + const deferred = createDeferred() + vi.mocked(getPublicSettings).mockReturnValue(deferred.promise) + const settings = createPublicSettings({ payment_enabled: true }) + const store = useAppStore() + + const first = store.fetchPublicSettings() + const second = store.fetchPublicSettings() + const forced = store.fetchPublicSettings(true) + + expect(getPublicSettings).toHaveBeenCalledTimes(1) + + const settled = vi.fn() + void first.then(settled) + void second.then(settled) + void forced.then(settled) + await Promise.resolve() + expect(settled).not.toHaveBeenCalled() + + deferred.resolve(settings) + await expect(Promise.all([first, second, forced])).resolves.toEqual([ + settings, + settings, + settings, + ]) + expect(store.publicSettingsLoaded).toBe(true) + expect(store.cachedPublicSettings?.payment_enabled).toBe(true) + }) + + it('force 在无活动请求时绕过缓存,刷新期间的普通调用等待刷新结果', async () => { + const initial = createPublicSettings({ site_name: 'Initial Site' }) + vi.mocked(getPublicSettings).mockResolvedValueOnce(initial) + const store = useAppStore() + await store.fetchPublicSettings() + + const deferred = createDeferred() + const updated = createPublicSettings({ site_name: 'Updated Site' }) + vi.mocked(getPublicSettings).mockReturnValueOnce(deferred.promise) + + const refresh = store.fetchPublicSettings(true) + const duringRefresh = store.fetchPublicSettings() + + expect(getPublicSettings).toHaveBeenCalledTimes(2) + + deferred.resolve(updated) + await expect(Promise.all([refresh, duringRefresh])).resolves.toEqual([updated, updated]) + expect(store.siteName).toBe('Updated Site') + + await expect(store.fetchPublicSettings()).resolves.toEqual(updated) + expect(getPublicSettings).toHaveBeenCalledTimes(2) + }) + + it('并发请求失败时所有调用得到 null,且不会标记设置已加载', async () => { + const deferred = createDeferred() + vi.mocked(getPublicSettings).mockReturnValue(deferred.promise) + const consoleError = vi.spyOn(console, 'error').mockImplementation(() => undefined) + const store = useAppStore() + + const first = store.fetchPublicSettings() + const second = store.fetchPublicSettings() + deferred.reject(new Error('network unavailable')) + + await expect(Promise.all([first, second])).resolves.toEqual([null, null]) + expect(getPublicSettings).toHaveBeenCalledTimes(1) + expect(store.publicSettingsLoaded).toBe(false) + expect(store.cachedPublicSettings).toBeNull() + consoleError.mockRestore() + }) + it('从 window.__APP_CONFIG__ 初始化', () => { const windowAny = window as any windowAny.__APP_CONFIG__ = { diff --git a/frontend/src/stores/app.ts b/frontend/src/stores/app.ts index 20d580f6b9..51fff5c0a6 100644 --- a/frontend/src/stores/app.ts +++ b/frontend/src/stores/app.ts @@ -33,6 +33,7 @@ export const useAppStore = defineStore('app', () => { const apiBaseUrl = ref('') const docUrl = ref('') const cachedPublicSettings = ref(null) + let publicSettingsRequest: Promise | null = null // Version cache state const versionLoaded = ref(false) @@ -306,19 +307,25 @@ export const useAppStore = defineStore('app', () => { * Fetch public settings (uses cache unless force=true) * @param force - Force refresh from API */ - async function fetchPublicSettings(force = false): Promise { + function fetchPublicSettings(force = false): Promise { + // An active request always wins over cache/force semantics so every caller observes + // the same refresh result and no older request can overwrite a newer one. + if (publicSettingsRequest) { + return publicSettingsRequest + } + // Check for injected config from server (eliminates flash) if (!publicSettingsLoaded.value && !force && window.__APP_CONFIG__) { applySettings(window.__APP_CONFIG__) - return window.__APP_CONFIG__ + return Promise.resolve(window.__APP_CONFIG__) } // Return cached data if available and not forcing refresh if (publicSettingsLoaded.value && !force) { if (cachedPublicSettings.value) { - return { ...cachedPublicSettings.value } + return Promise.resolve({ ...cachedPublicSettings.value }) } - return { + return Promise.resolve({ registration_enabled: false, email_verify_enabled: false, force_email_on_third_party_signup: false, @@ -362,25 +369,37 @@ export const useAppStore = defineStore('app', () => { service_quota_enabled: false, affiliate_enabled: false, allow_user_view_error_requests: false, - } - } - - // Prevent duplicate requests - if (publicSettingsLoading.value) { - return null + }) } publicSettingsLoading.value = true + let apiRequest: Promise try { - const data = await fetchPublicSettingsAPI() - applySettings(data) - return data + apiRequest = fetchPublicSettingsAPI() } catch (error) { console.error('Failed to fetch public settings:', error) - return null - } finally { publicSettingsLoading.value = false + return Promise.resolve(null) } + + const request = apiRequest + .then((data) => { + applySettings(data) + return data + }) + .catch((error) => { + console.error('Failed to fetch public settings:', error) + return null + }) + .finally(() => { + if (publicSettingsRequest === request) { + publicSettingsRequest = null + publicSettingsLoading.value = false + } + }) + + publicSettingsRequest = request + return request } /**