fix: harden billing concurrency and payment recovery

This commit is contained in:
superman2003
2026-07-10 10:56:49 +08:00
parent 12d811bd76
commit fc66a30ffc
25 changed files with 1353 additions and 198 deletions
+44
View File
@@ -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)
+6 -3
View File
@@ -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 = &current
fresh.WeeklyWindowStart = &current
fresh.MonthlyWindowStart = &current
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)
}
+212 -54
View File
@@ -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)
+20 -11
View File
@@ -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
+8
View File
@@ -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)