mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
Fix refund pending finalization gaps
This commit is contained in:
@@ -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)) == "" {
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
Reference in New Issue
Block a user