mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:08:03 +08:00
fix: harden billing concurrency and payment recovery
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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 原子性地累加订阅用量。
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -0,0 +1,177 @@
|
||||
import { beforeAll, beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
|
||||
type NavigationGuard = (
|
||||
to: Record<string, any>,
|
||||
from: Record<string, any>,
|
||||
next: ReturnType<typeof vi.fn>
|
||||
) => Promise<void>
|
||||
|
||||
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<T>() {
|
||||
let resolve!: (value: T | PromiseLike<T>) => void
|
||||
const promise = new Promise<T>((resolvePromise) => {
|
||||
resolve = resolvePromise
|
||||
})
|
||||
return { promise, resolve }
|
||||
}
|
||||
|
||||
function runGuard(meta: Record<string, unknown>, 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)
|
||||
})
|
||||
})
|
||||
@@ -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
|
||||
}
|
||||
|
||||
// 简易模式下限制访问某些页面
|
||||
|
||||
@@ -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<T>() {
|
||||
let resolve!: (value: T | PromiseLike<T>) => void
|
||||
let reject!: (reason?: unknown) => void
|
||||
const promise = new Promise<T>((resolvePromise, rejectPromise) => {
|
||||
resolve = resolvePromise
|
||||
reject = rejectPromise
|
||||
})
|
||||
|
||||
return { promise, resolve, reject }
|
||||
}
|
||||
|
||||
function createPublicSettings(overrides: Partial<PublicSettings> = {}): 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<PublicSettings>()
|
||||
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<PublicSettings>()
|
||||
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<PublicSettings>()
|
||||
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__ = {
|
||||
|
||||
+34
-15
@@ -33,6 +33,7 @@ export const useAppStore = defineStore('app', () => {
|
||||
const apiBaseUrl = ref<string>('')
|
||||
const docUrl = ref<string>('')
|
||||
const cachedPublicSettings = ref<PublicSettings | null>(null)
|
||||
let publicSettingsRequest: Promise<PublicSettings | null> | null = null
|
||||
|
||||
// Version cache state
|
||||
const versionLoaded = ref<boolean>(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<PublicSettings | null> {
|
||||
function fetchPublicSettings(force = false): Promise<PublicSettings | null> {
|
||||
// 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<PublicSettings>
|
||||
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
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
Reference in New Issue
Block a user