mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 15:38:22 +08:00
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:
@@ -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() {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user