Fix refund pending finalization gaps

This commit is contained in:
wucm667
2026-06-30 10:19:50 +08:00
parent 7316d83027
commit 93a3bf3077
17 changed files with 591 additions and 39 deletions
@@ -257,6 +257,22 @@ func (h *PaymentHandler) ProcessRefund(c *gin.Context) {
response.Success(c, result)
}
// QueryAndFinalizeRefund queries the provider refund status and finalizes a pending refund.
// POST /api/v1/admin/payment/orders/:id/refund/query
func (h *PaymentHandler) QueryAndFinalizeRefund(c *gin.Context) {
orderID, ok := parseIDParam(c, "id")
if !ok {
return
}
result, err := h.paymentService.QueryAndFinalizeRefund(c.Request.Context(), orderID)
if err != nil {
response.ErrorFrom(c, err)
return
}
response.Success(c, result)
}
// --- Subscription Plans ---
// ListPlans returns all subscription plans.
@@ -300,6 +300,25 @@ func (a *Airwallex) Refund(ctx context.Context, req payment.RefundRequest) (*pay
return refundResp, nil
}
func (a *Airwallex) QueryRefund(ctx context.Context, req payment.RefundQueryRequest) (*payment.RefundResponse, error) {
refundID := strings.TrimSpace(req.RefundID)
if refundID == "" {
return nil, fmt.Errorf("airwallex query refund: missing refund id")
}
token, err := a.accessToken(ctx)
if err != nil {
return nil, fmt.Errorf("airwallex auth: %w", err)
}
var resp airwallexRefund
if err := a.doJSON(ctx, http.MethodGet, "/pa/refunds/"+url.PathEscape(refundID), token, nil, &resp); err != nil {
return nil, fmt.Errorf("airwallex query refund: %w", err)
}
if strings.TrimSpace(resp.ID) == "" {
resp.ID = refundID
}
return &payment.RefundResponse{RefundID: resp.ID, Status: airwallexRefundProviderStatus(resp.Status)}, nil
}
func (a *Airwallex) CancelPayment(ctx context.Context, tradeNo string) error {
intentID := strings.TrimSpace(tradeNo)
if intentID == "" {
@@ -248,6 +248,50 @@ func (s *Stripe) Refund(ctx context.Context, req payment.RefundRequest) (*paymen
}, nil
}
// QueryRefund retrieves a Stripe refund by refund ID when available, otherwise
// falls back to the latest refund for the PaymentIntent.
func (s *Stripe) QueryRefund(ctx context.Context, req payment.RefundQueryRequest) (*payment.RefundResponse, error) {
s.ensureInit()
var r *stripe.Refund
var err error
if refundID := strings.TrimSpace(req.RefundID); refundID != "" {
r, err = s.sc.V1Refunds.Retrieve(ctx, refundID, nil)
if err != nil {
return nil, fmt.Errorf("stripe query refund: %w", err)
}
} else {
tradeNo := strings.TrimSpace(req.TradeNo)
if tradeNo == "" {
return nil, fmt.Errorf("stripe query refund: missing payment intent id")
}
params := &stripe.RefundListParams{PaymentIntent: stripe.String(tradeNo)}
params.Limit = stripe.Int64(1)
list := s.sc.V1Refunds.List(ctx, params)
if list.Err() != nil {
return nil, fmt.Errorf("stripe query refund: %w", list.Err())
}
refunds := list.Data()
if len(refunds) == 0 {
return nil, fmt.Errorf("stripe query refund: no refund found")
}
r = refunds[0]
}
return &payment.RefundResponse{RefundID: r.ID, Status: stripeRefundProviderStatus(r.Status)}, nil
}
func stripeRefundProviderStatus(status stripe.RefundStatus) string {
switch status {
case stripe.RefundStatusSucceeded:
return payment.ProviderStatusSuccess
case stripe.RefundStatusFailed, stripe.RefundStatusCanceled:
return payment.ProviderStatusFailed
default:
return payment.ProviderStatusPending
}
}
func stripeIntentCurrency(raw stripe.Currency, fallback string) string {
currency, err := payment.NormalizePaymentCurrency(string(raw))
if err != nil || currency == payment.DefaultPaymentCurrency && strings.TrimSpace(string(raw)) == "" {
+48 -7
View File
@@ -10,7 +10,6 @@ import (
"strconv"
"strings"
"sync"
"time"
"github.com/Wei-Shaw/sub2api/internal/payment"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
@@ -471,24 +470,66 @@ func (w *Wxpay) Refund(ctx context.Context, req payment.RefundRequest) (*payment
}
rs := refunddomestic.RefundsApiService{Client: c}
cur := wxpayCurrency
outRefundNo := wxpayRefundID(req.OrderID, req.Amount)
res, _, err := rs.Create(ctx, refunddomestic.CreateRequest{
OutTradeNo: core.String(req.OrderID),
OutRefundNo: core.String(fmt.Sprintf("%s-refund-%d", req.OrderID, time.Now().UnixNano())),
OutRefundNo: core.String(outRefundNo),
Reason: core.String(req.Reason),
Amount: &refunddomestic.AmountReq{Refund: core.Int64(rf), Total: core.Int64(tf), Currency: &cur},
})
if err != nil {
return nil, fmt.Errorf("wxpay refund: %w", err)
}
rid := wxSV(res.RefundId)
if rid == "" {
rid = fmt.Sprintf("%s-refund", req.OrderID)
}
st := payment.ProviderStatusPending
if res.Status != nil && *res.Status == refunddomestic.STATUS_SUCCESS {
st = payment.ProviderStatusSuccess
}
return &payment.RefundResponse{RefundID: rid, Status: st}, nil
return &payment.RefundResponse{RefundID: outRefundNo, Status: st}, nil
}
func (w *Wxpay) QueryRefund(ctx context.Context, req payment.RefundQueryRequest) (*payment.RefundResponse, error) {
c, err := w.ensureClient()
if err != nil {
return nil, err
}
outRefundNo := strings.TrimSpace(req.RefundID)
if outRefundNo == "" {
outRefundNo = wxpayRefundID(req.OrderID, req.Amount)
}
if outRefundNo == "" {
return nil, fmt.Errorf("wxpay query refund: missing refund id")
}
rs := refunddomestic.RefundsApiService{Client: c}
res, _, err := rs.QueryByOutRefundNo(ctx, refunddomestic.QueryByOutRefundNoRequest{
OutRefundNo: core.String(outRefundNo),
})
if err != nil {
return nil, fmt.Errorf("wxpay query refund: %w", err)
}
status := payment.ProviderStatusPending
if res != nil && res.Status != nil {
switch *res.Status {
case refunddomestic.STATUS_SUCCESS:
status = payment.ProviderStatusSuccess
case refunddomestic.STATUS_CLOSED, refunddomestic.STATUS_ABNORMAL:
status = payment.ProviderStatusFailed
default:
status = payment.ProviderStatusPending
}
}
return &payment.RefundResponse{RefundID: outRefundNo, Status: status}, nil
}
func wxpayRefundID(orderID, amount string) string {
orderID = strings.TrimSpace(orderID)
if orderID == "" {
return ""
}
amount = strings.NewReplacer(".", "", "-", "").Replace(strings.TrimSpace(amount))
if amount == "" {
return orderID + "-refund"
}
return orderID + "-refund-" + amount
}
func (w *Wxpay) queryOrderTotalFen(ctx context.Context, c *core.Client, orderID string) (int64, error) {
+15
View File
@@ -182,6 +182,15 @@ type RefundRequest struct {
Reason string
}
// RefundQueryRequest contains identifiers needed to query a previously
// requested refund.
type RefundQueryRequest struct {
TradeNo string
OrderID string
RefundID string
Amount string
}
// RefundResponse is returned after a refund request.
type RefundResponse struct {
RefundID string
@@ -216,6 +225,12 @@ type Provider interface {
Refund(ctx context.Context, req RefundRequest) (*RefundResponse, error)
}
// RefundQueryProvider extends Provider with refund status querying.
type RefundQueryProvider interface {
Provider
QueryRefund(ctx context.Context, req RefundQueryRequest) (*RefundResponse, error)
}
// CancelableProvider extends Provider with the ability to cancel pending payments.
type CancelableProvider interface {
Provider
@@ -85,6 +85,7 @@ func RegisterPaymentRoutes(
adminOrders.POST("/:id/cancel", adminPaymentHandler.CancelOrder)
adminOrders.POST("/:id/retry", adminPaymentHandler.RetryFulfillment)
adminOrders.POST("/:id/refund", adminPaymentHandler.ProcessRefund)
adminOrders.POST("/:id/refund/query", adminPaymentHandler.QueryAndFinalizeRefund)
}
// Subscription Plans
@@ -12,7 +12,6 @@ import (
"github.com/Wei-Shaw/sub2api/ent/paymentauditlog"
"github.com/Wei-Shaw/sub2api/ent/paymentorder"
"github.com/Wei-Shaw/sub2api/internal/payment"
"github.com/Wei-Shaw/sub2api/internal/payment/provider"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
)
@@ -454,7 +453,7 @@ func (s *PaymentService) createProviderFromInstance(ctx context.Context, inst *d
}
instID := strconv.FormatInt(int64(inst.ID), 10)
prov, err := provider.CreateProvider(inst.ProviderKey, instID, cfg)
prov, err := createPaymentProviderFromInstance(inst.ProviderKey, instID, cfg)
if err != nil {
return nil, fmt.Errorf("create provider from instance: %w", err)
}
+148 -4
View File
@@ -2,6 +2,7 @@ package service
import (
"context"
"encoding/json"
"errors"
"fmt"
"log/slog"
@@ -10,15 +11,20 @@ import (
"strings"
"time"
"entgo.io/ent/dialect/sql"
dbent "github.com/Wei-Shaw/sub2api/ent"
"github.com/Wei-Shaw/sub2api/ent/paymentauditlog"
"github.com/Wei-Shaw/sub2api/ent/paymentorder"
"github.com/Wei-Shaw/sub2api/ent/paymentproviderinstance"
"github.com/Wei-Shaw/sub2api/internal/payment"
"github.com/Wei-Shaw/sub2api/internal/payment/provider"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
)
// --- Refund Flow ---
var createPaymentProviderFromInstance = provider.CreateProvider
// getOrderProviderInstance looks up the provider instance that processed this order.
// For legacy orders without provider_instance_id, it resolves only when the
// historical instance is uniquely identifiable from the stored order fields.
@@ -203,7 +209,7 @@ func (s *PaymentService) PrepareRefund(ctx context.Context, oid int64, amt float
if err != nil {
return nil, nil, infraerrors.NotFound("NOT_FOUND", "order not found")
}
ok := []string{OrderStatusCompleted, OrderStatusRefundRequested, OrderStatusRefundFailed}
ok := []string{OrderStatusCompleted, OrderStatusRefundRequested, OrderStatusRefundPending, OrderStatusRefundFailed}
if !psSliceContains(ok, o.Status) {
return nil, nil, infraerrors.BadRequest("INVALID_STATUS", "order status does not allow refund")
}
@@ -274,7 +280,7 @@ func (s *PaymentService) prepDeduct(ctx context.Context, o *dbent.PaymentOrder,
}
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, OrderStatusRefundFailed)).SetStatus(OrderStatusRefunding).Save(ctx)
c, err := s.entClient.PaymentOrder.Update().Where(paymentorder.IDEQ(p.OrderID), paymentorder.StatusIn(OrderStatusCompleted, OrderStatusRefundRequested, OrderStatusRefundPending, OrderStatusRefundFailed)).SetStatus(OrderStatusRefunding).Save(ctx)
if err != nil {
return nil, fmt.Errorf("lock: %w", err)
}
@@ -386,12 +392,142 @@ func (s *PaymentService) finishRefund(ctx context.Context, p *RefundPlan, resp *
case payment.ProviderStatusSuccess, payment.ProviderStatusRefunded:
return s.markRefundOk(ctx, p)
case payment.ProviderStatusPending:
return s.markRefundPending(ctx, p)
return s.markRefundPending(ctx, p, resp)
default:
return s.handleGwFail(ctx, p, fmt.Errorf("payment refund returned unknown status: %s", strings.TrimSpace(resp.Status)))
}
}
func (s *PaymentService) QueryAndFinalizeRefund(ctx context.Context, oid int64) (*RefundResult, error) {
o, err := s.entClient.PaymentOrder.Get(ctx, oid)
if err != nil {
return nil, infraerrors.NotFound("NOT_FOUND", "order not found")
}
if o.Status != OrderStatusRefundPending {
return nil, infraerrors.BadRequest("INVALID_STATUS", "only refund pending orders can be finalized")
}
prov, err := s.getRefundProvider(ctx, o)
if err != nil {
return nil, fmt.Errorf("get refund provider: %w", err)
}
queryProvider, ok := prov.(payment.RefundQueryProvider)
if !ok {
return nil, infraerrors.BadRequest("REFUND_QUERY_UNSUPPORTED", "this payment provider does not support refund status query; please verify manually")
}
pendingDetail := s.latestRefundPendingDetail(ctx, oid)
resp, err := queryProvider.QueryRefund(ctx, payment.RefundQueryRequest{
TradeNo: o.PaymentTradeNo,
OrderID: o.OutTradeNo,
RefundID: pendingDetail.RefundID,
Amount: formatGatewayRefundAmount(o.RefundAmount, o),
})
if err != nil {
return nil, fmt.Errorf("query refund: %w", err)
}
if err := validateRefundProviderResponse(resp); err != nil {
return s.finalizeRefundFailed(ctx, o, err)
}
plan := s.refundFinalizePlan(o)
if !pendingDetail.DeductionRollbackOK {
plan.BalanceToDeduct = 0
plan.SubDaysToDeduct = 0
} else if o.OrderType == payment.OrderTypeSubscription {
if early := s.prepDeduct(ctx, o, plan, true); early != nil {
return early, nil
}
}
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)
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
default:
return s.finalizeRefundFailed(ctx, o, fmt.Errorf("payment refund returned unknown status: %s", strings.TrimSpace(resp.Status)))
}
}
func (s *PaymentService) refundFinalizePlan(o *dbent.PaymentOrder) *RefundPlan {
refundAmount := o.RefundAmount
reason := strings.TrimSpace(psStringValue(o.RefundReason))
if reason == "" {
reason = fmt.Sprintf("refund order:%d", o.ID)
}
return &RefundPlan{
OrderID: o.ID,
Order: o,
RefundAmount: refundAmount,
GatewayAmount: calculateGatewayRefundAmount(o.Amount, o.PayAmount, refundAmount, PaymentOrderCurrency(o)),
Reason: reason,
Force: o.ForceRefund,
DeductBalance: true,
DeductionType: payment.DeductionTypeBalance,
BalanceToDeduct: func() float64 {
if o.OrderType == payment.OrderTypeBalance {
return refundAmount
}
return 0
}(),
}
}
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 {
return fmt.Errorf("deduction: %w", err)
}
}
if p.DeductionType == payment.DeductionTypeSubscription && p.SubDaysToDeduct > 0 && p.SubscriptionID > 0 {
if _, err := s.subscriptionSvc.ExtendSubscription(ctx, p.SubscriptionID, -p.SubDaysToDeduct); err != nil {
if errors.Is(err, ErrAdjustWouldExpire) {
if revokeErr := s.subscriptionSvc.RevokeSubscription(ctx, p.SubscriptionID); revokeErr != nil {
return fmt.Errorf("revoke subscription: %w", revokeErr)
}
} else {
return fmt.Errorf("deduct subscription days: %w", err)
}
}
}
return nil
}
func (s *PaymentService) finalizeRefundFailed(ctx context.Context, o *dbent.PaymentOrder, gErr error) (*RefundResult, error) {
now := time.Now()
_, _ = s.entClient.PaymentOrder.UpdateOneID(o.ID).SetStatus(OrderStatusRefundFailed).SetFailedAt(now).SetFailedReason(psErrMsg(gErr)).Save(ctx)
s.writeAuditLog(ctx, o.ID, "REFUND_FAILED", "admin", map[string]any{"detail": psErrMsg(gErr)})
return &RefundResult{Success: false, Warning: "gateway refund failed: " + psErrMsg(gErr)}, nil
}
type refundPendingAuditDetail struct {
RefundID string `json:"refundID"`
DeductionRollbackOK bool `json:"deductionRollbackOK"`
}
func (s *PaymentService) latestRefundPendingDetail(ctx context.Context, oid int64) refundPendingAuditDetail {
logEntry, err := s.entClient.PaymentAuditLog.Query().
Where(paymentauditlog.OrderIDEQ(strconv.FormatInt(oid, 10)), paymentauditlog.ActionEQ("REFUND_PENDING")).
Order(paymentauditlog.ByCreatedAt(sql.OrderDesc())).
First(ctx)
if err != nil || logEntry == nil {
return refundPendingAuditDetail{DeductionRollbackOK: true}
}
detail := refundPendingAuditDetail{DeductionRollbackOK: true}
_ = json.Unmarshal([]byte(logEntry.Detail), &detail)
detail.RefundID = strings.TrimSpace(detail.RefundID)
return detail
}
// getRefundProvider creates a provider using the order's original instance config.
// Delegates to getOrderProvider which handles instance lookup and fallback.
func (s *PaymentService) getRefundProvider(ctx context.Context, o *dbent.PaymentOrder) (payment.Provider, error) {
@@ -431,7 +567,7 @@ func (s *PaymentService) markRefundOk(ctx context.Context, p *RefundPlan) (*Refu
return &RefundResult{Success: true, BalanceDeducted: p.BalanceToDeduct, SubDaysDeducted: p.SubDaysToDeduct}, nil
}
func (s *PaymentService) markRefundPending(ctx context.Context, p *RefundPlan) (*RefundResult, error) {
func (s *PaymentService) markRefundPending(ctx context.Context, p *RefundPlan, resp *payment.RefundResponse) (*RefundResult, error) {
balanceDeducted := p.BalanceToDeduct
subDaysDeducted := p.SubDaysToDeduct
rollbackOK := s.RollbackRefund(ctx, p, nil)
@@ -454,6 +590,7 @@ func (s *PaymentService) markRefundPending(ctx context.Context, p *RefundPlan) (
}
detail := map[string]any{
"refundID": refundResponseID(resp),
"refundAmount": p.RefundAmount,
"reason": p.Reason,
"force": p.Force,
@@ -472,6 +609,13 @@ func (s *PaymentService) markRefundPending(ctx context.Context, p *RefundPlan) (
return &RefundResult{Success: false, Warning: warning}, nil
}
func refundResponseID(resp *payment.RefundResponse) string {
if resp == nil {
return ""
}
return strings.TrimSpace(resp.RefundID)
}
func (s *PaymentService) RollbackRefund(ctx context.Context, p *RefundPlan, gErr error) bool {
if p.DeductionType == payment.DeductionTypeBalance && p.BalanceToDeduct > 0 {
if err := s.userRepo.UpdateBalance(ctx, p.Order.UserID, p.BalanceToDeduct); err != nil {
@@ -359,3 +359,153 @@ func TestFinishRefundSuccessStatusesFinalize(t *testing.T) {
})
}
}
func TestQueryAndFinalizeRefundFinalizesProviderStatuses(t *testing.T) {
for _, tc := range []struct {
name string
status string
wantStatus string
wantDeduct float64
}{
{name: "success", status: payment.ProviderStatusSuccess, wantStatus: OrderStatusRefunded, wantDeduct: 100},
{name: "failed", status: payment.ProviderStatusFailed, wantStatus: OrderStatusRefundFailed},
{name: "pending", status: payment.ProviderStatusPending, wantStatus: OrderStatusRefundPending},
} {
t.Run(tc.name, func(t *testing.T) {
ctx := context.Background()
client := newPaymentConfigServiceTestClient(t)
order := createPendingRefundOrderForTest(t, ctx, client, "query-finalize-"+tc.name)
var deducted float64
svc := &PaymentService{
entClient: client,
loadBalancer: &captureLoadBalancer{},
userRepo: &mockUserRepo{deductBalanceFn: func(ctx context.Context, id int64, amount float64) error {
deducted += amount
return nil
}},
}
restore := replacePaymentProviderFactoryForTest(t, &refundQueryProviderTestDouble{
refundResponse: &payment.RefundResponse{RefundID: "rf_test", Status: tc.status},
})
defer restore()
result, err := svc.QueryAndFinalizeRefund(ctx, order.ID)
require.NoError(t, err)
require.NotNil(t, result)
require.Equal(t, tc.status == payment.ProviderStatusSuccess, result.Success)
require.Equal(t, tc.wantDeduct, deducted)
reloaded, err := client.PaymentOrder.Get(ctx, order.ID)
require.NoError(t, err)
require.Equal(t, tc.wantStatus, reloaded.Status)
})
}
}
func TestQueryAndFinalizeRefundUnsupportedProviderReturnsClearError(t *testing.T) {
ctx := context.Background()
client := newPaymentConfigServiceTestClient(t)
order := createPendingRefundOrderForTest(t, ctx, client, "query-finalize-unsupported")
svc := &PaymentService{entClient: client, loadBalancer: &captureLoadBalancer{}}
restore := replacePaymentProviderFactoryForTest(t, refundProviderTestDouble{})
defer restore()
result, err := svc.QueryAndFinalizeRefund(ctx, order.ID)
require.Nil(t, result)
require.Error(t, err)
require.Equal(t, "REFUND_QUERY_UNSUPPORTED", infraerrors.Reason(err))
}
func createPendingRefundOrderForTest(t *testing.T, ctx context.Context, client *dbent.Client, suffix string) *dbent.PaymentOrder {
t.Helper()
user, err := client.User.Create().
SetEmail(suffix + "@example.com").
SetPasswordHash("hash").
SetUsername(suffix).
Save(ctx)
require.NoError(t, err)
inst, err := client.PaymentProviderInstance.Create().
SetProviderKey(payment.TypeStripe).
SetName(suffix + "-provider").
SetConfig("{}").
SetSupportedTypes("stripe").
SetEnabled(true).
SetRefundEnabled(true).
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-" + suffix).
SetOutTradeNo("sub2_" + suffix).
SetPaymentType(payment.TypeStripe).
SetPaymentTradeNo("pi_" + suffix).
SetOrderType(payment.OrderTypeBalance).
SetStatus(OrderStatusRefundPending).
SetRefundAmount(100).
SetRefundReason("pending refund").
SetExpiresAt(time.Now().Add(time.Hour)).
SetPaidAt(time.Now()).
SetClientIP("127.0.0.1").
SetSrcHost("api.example.com").
SetProviderInstanceID(strconv.FormatInt(inst.ID, 10)).
Save(ctx)
require.NoError(t, err)
_, err = client.PaymentAuditLog.Create().
SetOrderID(strconv.FormatInt(order.ID, 10)).
SetAction("REFUND_PENDING").
SetOperator("admin").
SetDetail(`{"refundID":"rf_test","deductionRollbackOK":true}`).
Save(ctx)
require.NoError(t, err)
return order
}
func replacePaymentProviderFactoryForTest(t *testing.T, prov payment.Provider) func() {
t.Helper()
original := createPaymentProviderFromInstance
createPaymentProviderFromInstance = func(providerKey, instanceID string, config map[string]string) (payment.Provider, error) {
return prov, nil
}
return func() { createPaymentProviderFromInstance = original }
}
type refundProviderTestDouble struct{}
func (refundProviderTestDouble) Name() string { return "refund-test" }
func (refundProviderTestDouble) ProviderKey() string {
return payment.TypeStripe
}
func (refundProviderTestDouble) SupportedTypes() []payment.PaymentType {
return []payment.PaymentType{payment.TypeStripe}
}
func (refundProviderTestDouble) CreatePayment(context.Context, payment.CreatePaymentRequest) (*payment.CreatePaymentResponse, error) {
return nil, nil
}
func (refundProviderTestDouble) QueryOrder(context.Context, string) (*payment.QueryOrderResponse, error) {
return nil, nil
}
func (refundProviderTestDouble) VerifyNotification(context.Context, string, map[string]string) (*payment.PaymentNotification, error) {
return nil, nil
}
func (refundProviderTestDouble) Refund(context.Context, payment.RefundRequest) (*payment.RefundResponse, error) {
return nil, nil
}
type refundQueryProviderTestDouble struct {
refundProviderTestDouble
refundResponse *payment.RefundResponse
}
func (p *refundQueryProviderTestDouble) QueryRefund(context.Context, payment.RefundQueryRequest) (*payment.RefundResponse, error) {
return p.refundResponse, nil
}
@@ -25,6 +25,7 @@ import (
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
@@ -193,7 +194,12 @@ func (m *mockUserRepo) UpdateUserLastActiveAt(_ context.Context, userID int64, a
m.updateLastActiveAt = append(m.updateLastActiveAt, activeAt)
return m.updateLastActiveErr
}
func (m *mockUserRepo) DeductBalance(context.Context, int64, float64) error { return nil }
func (m *mockUserRepo) DeductBalance(ctx context.Context, id int64, amount float64) error {
if m.deductBalanceFn != nil {
return m.deductBalanceFn(ctx, id, amount)
}
return nil
}
func (m *mockUserRepo) UpdateConcurrency(context.Context, int64, int) error { return nil }
func (m *mockUserRepo) ExistsByEmail(context.Context, string) (bool, error) { return false, nil }
func (m *mockUserRepo) RemoveGroupFromAllowedGroups(context.Context, int64) (int64, error) {