Merge pull request #5199 from wucm667/fix/issue-5191-refund-balance-force

fix(payment): require force for insufficient refund balance
This commit is contained in:
Wesley Liddick
2026-08-03 11:27:52 +08:00
committed by GitHub
5 changed files with 351 additions and 36 deletions
+43
View File
@@ -859,6 +859,49 @@ func (r *userRepository) DeductBalance(ctx context.Context, id int64, amount flo
return nil
}
// DeductAvailableBalance atomically deducts min(amount, max(balance, 0)).
// Unlike DeductBalance, this refund-specific operation never increases an
// existing deficit or permits a concurrent deduction to cause an overdraft.
func (r *userRepository) DeductAvailableBalance(ctx context.Context, id int64, amount float64) (deducted float64, err error) {
if amount < 0 {
return 0, fmt.Errorf("deduction amount must be nonnegative")
}
const updateSQL = `
WITH target AS (
SELECT id, balance
FROM users
WHERE id = $2 AND deleted_at IS NULL
FOR UPDATE
), updated AS (
UPDATE users AS u
SET balance = target.balance - LEAST($1, GREATEST(target.balance, 0)), updated_at = NOW()
FROM target
WHERE u.id = target.id AND u.deleted_at IS NULL
RETURNING target.balance - u.balance AS deducted
)
SELECT deducted FROM updated
`
rows, err := clientFromContext(ctx, r.client).QueryContext(ctx, updateSQL, amount, id)
if err != nil {
return 0, err
}
defer func() {
if closeErr := rows.Close(); closeErr != nil && err == nil {
err = closeErr
}
}()
if !rows.Next() {
if rowsErr := rows.Err(); rowsErr != nil {
return 0, rowsErr
}
return 0, service.ErrUserNotFound
}
if err := rows.Scan(&deducted); err != nil {
return 0, err
}
return deducted, rows.Err()
}
// AdjustBalance 原子地把 delta 累加到余额上,结果为负时整条语句不生效。
// 相比"读余额 → 算新值 → 整行写回",这里把读与写压进同一条 UPDATE,
// 并发的计费扣款不会被旧快照覆盖。
@@ -4,6 +4,7 @@ package repository
import (
"context"
"strings"
"sync"
"testing"
"time"
@@ -479,6 +480,30 @@ func (s *UserRepoSuite) TestDeductBalance_AllowsOverdraft() {
s.Require().InDelta(-5.0, got.Balance, 1e-6, "Balance should be -5.0 after overdraft")
}
func (s *UserRepoSuite) TestDeductAvailableBalance_ClampsToNonnegativeBalance() {
for _, tc := range []struct {
name string
balance float64
requested float64
wantDeduct float64
wantBalance float64
}{
{name: "enough balance", balance: 10, requested: 4, wantDeduct: 4, wantBalance: 6},
{name: "insufficient balance", balance: 5, requested: 10, wantDeduct: 5, wantBalance: 0},
{name: "negative balance unchanged", balance: -3, requested: 10, wantDeduct: 0, wantBalance: -3},
} {
s.Run(tc.name, func() {
user := s.mustCreateUser(&service.User{Email: "available-" + strings.ReplaceAll(tc.name, " ", "-") + "@test.com", Balance: tc.balance})
deducted, err := s.repo.DeductAvailableBalance(s.ctx, user.ID, tc.requested)
s.Require().NoError(err)
s.Require().InDelta(tc.wantDeduct, deducted, 1e-6)
got, err := s.repo.GetByID(s.ctx, user.ID)
s.Require().NoError(err)
s.Require().InDelta(tc.wantBalance, got.Balance, 1e-6)
})
}
}
// --- Concurrency ---
func (s *UserRepoSuite) TestUpdateConcurrency() {
+84 -12
View File
@@ -276,10 +276,25 @@ func (s *PaymentService) prepDeduct(ctx context.Context, o *dbent.PaymentOrder,
return nil
}
p.DeductionType = payment.DeductionTypeBalance
p.BalanceToDeduct = math.Min(p.RefundAmount, u.Balance)
if u.Balance < p.RefundAmount && !force {
return &RefundResult{Success: false, Warning: "user balance is insufficient for deduction, use force", RequireForce: true}
}
p.BalanceToDeduct = math.Max(0, math.Min(p.RefundAmount, u.Balance))
return nil
}
type availableBalanceDeductor interface {
DeductAvailableBalance(ctx context.Context, id int64, amount float64) (float64, error)
}
func (s *PaymentService) deductAvailableBalance(ctx context.Context, userID int64, amount float64) (float64, error) {
repo, ok := s.userRepo.(availableBalanceDeductor)
if !ok {
return 0, errors.New("user repository does not support available balance deduction")
}
return repo.DeductAvailableBalance(ctx, userID, amount)
}
func (s *PaymentService) ExecuteRefund(ctx context.Context, p *RefundPlan) (*RefundResult, error) {
c, err := s.entClient.PaymentOrder.Update().Where(paymentorder.IDEQ(p.OrderID), paymentorder.StatusIn(OrderStatusCompleted, OrderStatusRefundRequested, OrderStatusRefundPending, OrderStatusRefundFailed)).SetStatus(OrderStatusRefunding).Save(ctx)
if err != nil {
@@ -292,10 +307,12 @@ func (s *PaymentService) ExecuteRefund(ctx context.Context, p *RefundPlan) (*Ref
// Skip balance deduction on retry if previous attempt already deducted
// but failed to roll back (REFUND_ROLLBACK_FAILED in audit log).
if !s.hasAuditLog(ctx, p.OrderID, "REFUND_ROLLBACK_FAILED") {
if err := s.userRepo.DeductBalance(ctx, p.Order.UserID, p.BalanceToDeduct); err != nil {
deducted, err := s.deductAvailableBalance(ctx, p.Order.UserID, p.BalanceToDeduct)
if err != nil {
s.restoreStatus(ctx, p)
return nil, fmt.Errorf("deduction: %w", err)
}
p.BalanceToDeduct = deducted
} else {
slog.Warn("skipping balance deduction on retry (previous rollback failed)", "orderID", p.OrderID)
p.BalanceToDeduct = 0
@@ -446,10 +463,7 @@ func (s *PaymentService) QueryAndFinalizeRefund(ctx context.Context, oid int64)
}
switch strings.TrimSpace(resp.Status) {
case payment.ProviderStatusSuccess, payment.ProviderStatusRefunded:
if err := s.applyRefundFinalDeduction(ctx, plan); err != nil {
return nil, err
}
return s.markRefundOk(ctx, plan)
return s.finalizePendingRefundSuccess(ctx, plan)
case payment.ProviderStatusPending:
s.writeAuditLog(ctx, oid, "REFUND_QUERY_PENDING", "admin", map[string]any{"refundID": resp.RefundID})
return &RefundResult{Success: false, Warning: "gateway refund is still pending confirmation"}, nil
@@ -458,6 +472,42 @@ func (s *PaymentService) QueryAndFinalizeRefund(ctx context.Context, oid int64)
}
}
func (s *PaymentService) finalizePendingRefundSuccess(ctx context.Context, p *RefundPlan) (_ *RefundResult, err error) {
tx, err := s.entClient.Tx(ctx)
if err != nil {
return nil, fmt.Errorf("begin refund finalization: %w", err)
}
defer func() {
if err != nil {
_ = tx.Rollback()
}
}()
txCtx := dbent.NewTxContext(ctx, tx)
claimed, err := tx.PaymentOrder.Update().
Where(paymentorder.IDEQ(p.OrderID), paymentorder.StatusEQ(OrderStatusRefundPending)).
SetStatus(OrderStatusRefunding).
Save(txCtx)
if err != nil {
return nil, fmt.Errorf("claim pending refund: %w", err)
}
if claimed == 0 {
return nil, infraerrors.Conflict("CONFLICT", "order status changed")
}
if err := s.applyRefundFinalDeduction(txCtx, p); err != nil {
return nil, err
}
result, err := s.markRefundOkTx(txCtx, tx.Client(), p)
if err != nil {
return nil, err
}
if err = tx.Commit(); err != nil {
return nil, fmt.Errorf("commit refund finalization: %w", err)
}
return result, nil
}
func (s *PaymentService) refundFinalizePlan(o *dbent.PaymentOrder) *RefundPlan {
refundAmount := o.RefundAmount
reason := strings.TrimSpace(psStringValue(o.RefundReason))
@@ -483,15 +533,12 @@ func (s *PaymentService) refundFinalizePlan(o *dbent.PaymentOrder) *RefundPlan {
}
func (s *PaymentService) applyRefundFinalDeduction(ctx context.Context, p *RefundPlan) error {
if s.hasAuditLog(ctx, p.OrderID, "REFUND_SUCCESS") {
p.BalanceToDeduct = 0
p.SubDaysToDeduct = 0
return nil
}
if p.DeductionType == payment.DeductionTypeBalance && p.BalanceToDeduct > 0 {
if err := s.userRepo.DeductBalance(ctx, p.Order.UserID, p.BalanceToDeduct); err != nil {
deducted, err := s.deductAvailableBalance(ctx, p.Order.UserID, p.BalanceToDeduct)
if err != nil {
return fmt.Errorf("deduction: %w", err)
}
p.BalanceToDeduct = deducted
}
if p.DeductionType == payment.DeductionTypeSubscription && p.SubDaysToDeduct > 0 && p.SubscriptionID > 0 {
if _, err := s.subscriptionSvc.ExtendSubscription(ctx, p.SubscriptionID, -p.SubDaysToDeduct); err != nil {
@@ -572,6 +619,31 @@ func (s *PaymentService) markRefundOk(ctx context.Context, p *RefundPlan) (*Refu
return &RefundResult{Success: true, BalanceDeducted: p.BalanceToDeduct, SubDaysDeducted: p.SubDaysToDeduct}, nil
}
func (s *PaymentService) markRefundOkTx(ctx context.Context, client *dbent.Client, p *RefundPlan) (*RefundResult, error) {
fs := OrderStatusRefunded
if p.RefundAmount < p.Order.Amount {
fs = OrderStatusPartiallyRefunded
}
now := time.Now()
_, err := client.PaymentOrder.UpdateOneID(p.OrderID).SetStatus(fs).SetRefundAmount(p.RefundAmount).SetRefundReason(p.Reason).SetRefundAt(now).SetForceRefund(p.Force).Save(ctx)
if err != nil {
return nil, fmt.Errorf("mark refund: %w", err)
}
detail, err := json.Marshal(map[string]any{"refundAmount": p.RefundAmount, "reason": p.Reason, "balanceDeducted": p.BalanceToDeduct, "force": p.Force})
if err != nil {
return nil, fmt.Errorf("marshal refund audit: %w", err)
}
if _, err := client.PaymentAuditLog.Create().
SetOrderID(strconv.FormatInt(p.OrderID, 10)).
SetAction("REFUND_SUCCESS").
SetDetail(string(detail)).
SetOperator("admin").
Save(ctx); err != nil {
return nil, fmt.Errorf("write refund audit: %w", err)
}
return &RefundResult{Success: true, BalanceDeducted: p.BalanceToDeduct, SubDaysDeducted: p.SubDaysToDeduct}, nil
}
func (s *PaymentService) markRefundPending(ctx context.Context, p *RefundPlan, resp *payment.RefundResponse) (*RefundResult, error) {
balanceDeducted := p.BalanceToDeduct
subDaysDeducted := p.SubDaysToDeduct
+171 -4
View File
@@ -4,6 +4,8 @@ package service
import (
"context"
"errors"
"fmt"
"strconv"
"testing"
"time"
@@ -119,6 +121,93 @@ func TestPrepareRefundRejectsLegacyGuessedProviderInstance(t *testing.T) {
require.Equal(t, "REFUND_DISABLED", infraerrors.Reason(err))
}
func TestPrepDeductBalanceRequiresForceWhenBalanceIsInsufficient(t *testing.T) {
for _, tc := range []struct {
name string
balance float64
force bool
wantDeduct float64
wantWarning bool
}{
{name: "insufficient balance", balance: 40, wantWarning: true},
{name: "forced insufficient balance", balance: 40, force: true, wantDeduct: 40},
{name: "equal balance", balance: 100, wantDeduct: 100},
} {
t.Run(tc.name, func(t *testing.T) {
plan := &RefundPlan{RefundAmount: 100}
svc := &PaymentService{userRepo: &mockUserRepo{getByIDUser: &User{Balance: tc.balance}}}
result := svc.prepDeduct(context.Background(), &dbent.PaymentOrder{
UserID: 1,
OrderType: payment.OrderTypeBalance,
}, plan, tc.force)
if tc.wantWarning {
require.NotNil(t, result)
require.False(t, result.Success)
require.True(t, result.RequireForce)
require.Equal(t, "user balance is insufficient for deduction, use force", result.Warning)
require.Zero(t, plan.BalanceToDeduct)
return
}
require.Nil(t, result)
require.Equal(t, payment.DeductionTypeBalance, plan.DeductionType)
require.Equal(t, tc.wantDeduct, plan.BalanceToDeduct)
})
}
}
func TestExecuteRefundUsesActualAvailableBalanceDeduction(t *testing.T) {
ctx := context.Background()
client := newPaymentConfigServiceTestClient(t)
user, err := client.User.Create().
SetEmail("refund-execute-clamp@example.com").
SetPasswordHash("hash").
SetUsername("refund-execute-clamp").
Save(ctx)
require.NoError(t, err)
order, err := client.PaymentOrder.Create().
SetUserID(user.ID).
SetUserEmail(user.Email).
SetUserName(user.Username).
SetAmount(100).
SetPayAmount(100).
SetFeeRate(0).
SetRechargeCode("REFUND-EXECUTE-CLAMP").
SetOutTradeNo("refund_execute_clamp").
SetPaymentType(payment.TypeStripe).
SetPaymentTradeNo("").
SetOrderType(payment.OrderTypeBalance).
SetStatus(OrderStatusCompleted).
SetExpiresAt(time.Now().Add(time.Hour)).
SetPaidAt(time.Now()).
SetClientIP("127.0.0.1").
SetSrcHost("api.example.com").
Save(ctx)
require.NoError(t, err)
repo := &mockUserRepo{deductAvailableBalanceFn: func(_ context.Context, id int64, amount float64) (float64, error) {
require.Equal(t, user.ID, id)
require.Equal(t, 100.0, amount)
return 25, nil
}}
plan := &RefundPlan{
OrderID: order.ID, Order: order, RefundAmount: 100, GatewayAmount: 100,
Reason: "concurrent spend", Force: true, DeductionType: payment.DeductionTypeBalance, BalanceToDeduct: 100,
}
result, err := (&PaymentService{entClient: client, userRepo: repo}).ExecuteRefund(ctx, plan)
require.NoError(t, err)
require.True(t, result.Success)
require.Equal(t, 25.0, plan.BalanceToDeduct)
require.Equal(t, 25.0, result.BalanceDeducted)
audit, err := client.PaymentAuditLog.Query().
Where(paymentauditlog.OrderIDEQ(strconv.FormatInt(order.ID, 10)), paymentauditlog.ActionEQ("REFUND_SUCCESS")).
Only(ctx)
require.NoError(t, err)
require.Contains(t, audit.Detail, `"balanceDeducted":25`)
}
func TestGwRefundRejectsAlipayMerchantIdentitySnapshotMismatch(t *testing.T) {
ctx := context.Background()
client := newPaymentConfigServiceTestClient(t)
@@ -366,8 +455,10 @@ func TestQueryAndFinalizeRefundFinalizesProviderStatuses(t *testing.T) {
status string
wantStatus string
wantDeduct float64
available float64
}{
{name: "success", status: payment.ProviderStatusSuccess, wantStatus: OrderStatusRefunded, wantDeduct: 100},
{name: "success", status: payment.ProviderStatusSuccess, wantStatus: OrderStatusRefunded, wantDeduct: 100, available: 100},
{name: "success clamps current balance", status: payment.ProviderStatusSuccess, wantStatus: OrderStatusRefunded, wantDeduct: 35, available: 35},
{name: "failed", status: payment.ProviderStatusFailed, wantStatus: OrderStatusRefundFailed},
{name: "pending", status: payment.ProviderStatusPending, wantStatus: OrderStatusRefundPending},
} {
@@ -380,9 +471,9 @@ func TestQueryAndFinalizeRefundFinalizesProviderStatuses(t *testing.T) {
svc := &PaymentService{
entClient: client,
loadBalancer: &captureLoadBalancer{},
userRepo: &mockUserRepo{deductBalanceFn: func(ctx context.Context, id int64, amount float64) error {
deducted += amount
return nil
userRepo: &mockUserRepo{deductAvailableBalanceFn: func(ctx context.Context, id int64, amount float64) (float64, error) {
deducted += tc.available
return tc.available, nil
}},
}
restore := replacePaymentProviderFactoryForTest(t, &refundQueryProviderTestDouble{
@@ -395,6 +486,14 @@ func TestQueryAndFinalizeRefundFinalizesProviderStatuses(t *testing.T) {
require.NotNil(t, result)
require.Equal(t, tc.status == payment.ProviderStatusSuccess, result.Success)
require.Equal(t, tc.wantDeduct, deducted)
if tc.status == payment.ProviderStatusSuccess {
require.Equal(t, tc.wantDeduct, result.BalanceDeducted)
audit, err := client.PaymentAuditLog.Query().
Where(paymentauditlog.OrderIDEQ(strconv.FormatInt(order.ID, 10)), paymentauditlog.ActionEQ("REFUND_SUCCESS")).
Only(ctx)
require.NoError(t, err)
require.Contains(t, audit.Detail, fmt.Sprintf(`"balanceDeducted":%v`, tc.wantDeduct))
}
reloaded, err := client.PaymentOrder.Get(ctx, order.ID)
require.NoError(t, err)
@@ -403,6 +502,74 @@ func TestQueryAndFinalizeRefundFinalizesProviderStatuses(t *testing.T) {
}
}
func TestFinalizePendingRefundSuccessRejectsStaleCallerBeforeSecondDeduction(t *testing.T) {
ctx := context.Background()
client := newPaymentConfigServiceTestClient(t)
order := createPendingRefundOrderForTest(t, ctx, client, "finalize-stale")
deductions := 0
svc := &PaymentService{
entClient: client,
userRepo: &mockUserRepo{deductAvailableBalanceFn: func(ctx context.Context, id int64, amount float64) (float64, error) {
require.NotNil(t, dbent.TxFromContext(ctx))
deductions++
return amount, nil
}},
}
first, err := svc.finalizePendingRefundSuccess(ctx, svc.refundFinalizePlan(order))
require.NoError(t, err)
require.True(t, first.Success)
second, err := svc.finalizePendingRefundSuccess(ctx, svc.refundFinalizePlan(order))
require.Nil(t, second)
require.Error(t, err)
require.Equal(t, "CONFLICT", infraerrors.Reason(err))
require.Equal(t, 1, deductions)
successAudits, err := client.PaymentAuditLog.Query().
Where(paymentauditlog.OrderIDEQ(strconv.FormatInt(order.ID, 10)), paymentauditlog.ActionEQ("REFUND_SUCCESS")).
Count(ctx)
require.NoError(t, err)
require.Equal(t, 1, successAudits)
}
func TestFinalizePendingRefundSuccessRollsBackPostDeductionFailure(t *testing.T) {
ctx := context.Background()
client := newPaymentConfigServiceTestClient(t)
order := createPendingRefundOrderForTest(t, ctx, client, "finalize-rollback")
_, err := client.User.UpdateOneID(order.UserID).SetBalance(100).Save(ctx)
require.NoError(t, err)
svc := &PaymentService{
entClient: client,
userRepo: &mockUserRepo{deductAvailableBalanceFn: func(ctx context.Context, id int64, amount float64) (float64, error) {
tx := dbent.TxFromContext(ctx)
require.NotNil(t, tx)
if _, updateErr := tx.Client().User.UpdateOneID(id).AddBalance(-amount).Save(ctx); updateErr != nil {
return 0, updateErr
}
return 0, errors.New("injected failure after deduction")
}},
}
result, err := svc.finalizePendingRefundSuccess(ctx, svc.refundFinalizePlan(order))
require.Nil(t, result)
require.ErrorContains(t, err, "injected failure after deduction")
user, err := client.User.Get(ctx, order.UserID)
require.NoError(t, err)
require.Equal(t, 100.0, user.Balance)
reloaded, err := client.PaymentOrder.Get(ctx, order.ID)
require.NoError(t, err)
require.Equal(t, OrderStatusRefundPending, reloaded.Status)
successAudits, err := client.PaymentAuditLog.Query().
Where(paymentauditlog.OrderIDEQ(strconv.FormatInt(order.ID, 10)), paymentauditlog.ActionEQ("REFUND_SUCCESS")).
Count(ctx)
require.NoError(t, err)
require.Zero(t, successAudits)
}
func TestQueryAndFinalizeRefundUnsupportedProviderReturnsClearError(t *testing.T) {
ctx := context.Background()
client := newPaymentConfigServiceTestClient(t)
+28 -20
View File
@@ -23,26 +23,27 @@ import (
// --- mock: UserRepository ---
type mockUserRepo struct {
updateBalanceErr error
updateBalanceFn func(ctx context.Context, id int64, amount float64) error
deductBalanceFn func(ctx context.Context, id int64, amount float64) error
getByIDUser *User
getByIDErr error
identities []UserAuthIdentityRecord
unbindIdentityErr error
unboundProviders []string
updateLastActiveErr error
updateLastActiveUserIDs []int64
updateLastActiveAt []time.Time
updateFn func(ctx context.Context, user *User) error
updateCalls int
updateFields []UserUpdateFields
upsertAvatarFn func(ctx context.Context, userID int64, input UpsertUserAvatarInput) (*UserAvatar, error)
upsertAvatarArgs []UpsertUserAvatarInput
deleteAvatarFn func(ctx context.Context, userID int64) error
deleteAvatarIDs []int64
getAvatarFn func(ctx context.Context, userID int64) (*UserAvatar, error)
txCalls int
updateBalanceErr error
updateBalanceFn func(ctx context.Context, id int64, amount float64) error
deductBalanceFn func(ctx context.Context, id int64, amount float64) error
deductAvailableBalanceFn func(ctx context.Context, id int64, amount float64) (float64, error)
getByIDUser *User
getByIDErr error
identities []UserAuthIdentityRecord
unbindIdentityErr error
unboundProviders []string
updateLastActiveErr error
updateLastActiveUserIDs []int64
updateLastActiveAt []time.Time
updateFn func(ctx context.Context, user *User) error
updateCalls int
updateFields []UserUpdateFields
upsertAvatarFn func(ctx context.Context, userID int64, input UpsertUserAvatarInput) (*UserAvatar, error)
upsertAvatarArgs []UpsertUserAvatarInput
deleteAvatarFn func(ctx context.Context, userID int64) error
deleteAvatarIDs []int64
getAvatarFn func(ctx context.Context, userID int64) (*UserAvatar, error)
txCalls int
}
type mockUserRepoTxKey struct{}
@@ -204,6 +205,13 @@ func (m *mockUserRepo) DeductBalance(ctx context.Context, id int64, amount float
return nil
}
func (m *mockUserRepo) DeductAvailableBalance(ctx context.Context, id int64, amount float64) (float64, error) {
if m.deductAvailableBalanceFn != nil {
return m.deductAvailableBalanceFn(ctx, id, amount)
}
return amount, nil
}
func (m *mockUserRepo) AdjustBalance(ctx context.Context, id int64, delta float64) (BalanceChange, error) {
panic("unexpected AdjustBalance call")
}