mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
fix: harden billing concurrency and payment recovery
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user