feat(v2): gateway.Refund 编排(部分/多次累计 + 超退守卫 + crypto 人工向闭环)+ RefundingProvider 加 refundID 幂等键

This commit is contained in:
wangjia
2026-07-10 17:23:14 +08:00
parent aa14204640
commit 2016c5d16e
10 changed files with 367 additions and 14 deletions
+2 -1
View File
@@ -44,6 +44,7 @@ func TestE2ECryptoQuerySettles(t *testing.T) {
db := model.OpenTestDB(t)
orders := store.NewOrderStore(db)
refunds := store.NewRefundStore(db)
acctReg := accounts.New([]config.AccountConfig{
{AccountID: "e2e-1", Channel: "crypto", Enabled: true, Region: "global", CredentialEnvPrefix: "e2e"},
})
@@ -55,7 +56,7 @@ func TestE2ECryptoQuerySettles(t *testing.T) {
picker := accounts.NewRouter(acctReg, nil, nil)
spy := &spyEnqueuer{}
g := gateway.New(orders, preg, picker, cryptoResolver{}, spy, "global")
g := gateway.New(orders, refunds, preg, picker, cryptoResolver{}, spy, "global")
// 下单 → 从 session payload 拿到期望链上金额(base+唯一尾数),喂给假 TronGrid。
res, err := g.CreateOrder(context.Background(), gateway.CreateOrderInput{
+3 -2
View File
@@ -40,6 +40,7 @@ type WebhookEnqueuer interface {
type Gateway struct {
orders *store.OrderStore
refunds *store.RefundStore
providers *provider.Registry
picker accounts.Picker
products ProductResolver
@@ -47,9 +48,9 @@ type Gateway struct {
region string
}
func New(orders *store.OrderStore, providers *provider.Registry, picker accounts.Picker,
func New(orders *store.OrderStore, refunds *store.RefundStore, providers *provider.Registry, picker accounts.Picker,
products ProductResolver, webhook WebhookEnqueuer, region string) *Gateway {
return &Gateway{orders: orders, providers: providers, picker: picker,
return &Gateway{orders: orders, refunds: refunds, providers: providers, picker: picker,
products: products, webhook: webhook, region: region}
}
+4 -2
View File
@@ -49,7 +49,9 @@ func (s *spyEnqueuer) Enqueue(outTradeNo, bizSystem, eventType, refundID string,
func newGateway(t *testing.T) (*gateway.Gateway, *fake.Provider, *spyEnqueuer, *store.OrderStore) {
t.Helper()
orders := store.NewOrderStore(model.OpenTestDB(t))
db := model.OpenTestDB(t)
orders := store.NewOrderStore(db)
refunds := store.NewRefundStore(db)
preg := provider.NewRegistry()
fp := fake.New()
preg.Register(fp)
@@ -60,7 +62,7 @@ func newGateway(t *testing.T) (*gateway.Gateway, *fake.Provider, *spyEnqueuer, *
})
picker := accounts.NewRouter(areg, nil, nil) // 默认 round_robin
spy := &spyEnqueuer{}
g := gateway.New(orders, preg, picker, stubResolver{}, spy, "global")
g := gateway.New(orders, refunds, preg, picker, stubResolver{}, spy, "global")
return g, fp, spy, orders
}
+199
View File
@@ -0,0 +1,199 @@
package gateway
import (
"context"
"errors"
"time"
"github.com/wangjia/pay/internal/model"
"github.com/wangjia/pay/internal/provider"
"github.com/wangjia/pay/internal/util"
)
var (
ErrOrderNotRefundable = errors.New("gateway: order not refundable (not settled or fully refunded)")
ErrRefundAmountInvalid = errors.New("gateway: refund amount invalid or exceeds refundable balance")
ErrRefundNotManual = errors.New("gateway: refund is not awaiting manual settlement")
)
type RefundInput struct {
OutTradeNo string
AmountMinor int64
Reason string
BizSystem string // 非空则须等于 order.BizSystem(归属校验);平台发起可空
InitiatedBy string // business/platform;空默认 business
}
type RefundResult struct {
RefundID string `json:"refund_id"`
Status string `json:"status"`
}
// Refund 业务发起退款:校验 → 经 store.CreateRefundGuarded 单事务(锁订单行 + SUM +
// 超退守卫 + 插入)建单 → 渠道退款(或 crypto 人工向)→ 幂等落状态 + 态机 + 事件。
//
// 超退守卫铁律:插入 Refund 行必须走 CreateRefundGuarded 的单事务 锁单+SUM+guard+insert,
// 绝不能用 RefundSum(...)-then-CreateRefund 两步序列——并发同单退款下两步之间会竞态
// (都读到 reserved=0、都通过校验、都插入 → 超退)。guard 用的"已付金额"取自本次已读到
// 的订单行 o.AmountMinor(而非调用方传入的任意数),防止上游拼错/被篡改的金额把守卫绕过。
func (g *Gateway) Refund(ctx context.Context, in RefundInput) (*RefundResult, error) {
o, err := g.orders.GetOrder(in.OutTradeNo)
if err != nil {
return nil, err // ErrOrderNotFound
}
if in.BizSystem != "" && o.BizSystem != in.BizSystem {
return nil, ErrOrderNotRefundable // 非本业务的单
}
if !o.Status.Settled() || o.Status == model.OrderRefundedV2 {
return nil, ErrOrderNotRefundable
}
if in.AmountMinor <= 0 {
return nil, ErrRefundAmountInvalid
}
att, err := g.orders.PaidAttempt(in.OutTradeNo)
if err != nil {
return nil, err // ErrAttemptNotFound
}
if in.InitiatedBy == "" {
in.InitiatedBy = "business"
}
refundID := util.NewOutTradeNo("rf")
prov, err := g.providers.Get(att.Channel)
if err != nil {
return nil, err
}
rp, canRefund := prov.(provider.RefundingProvider)
channelRefundable := canRefund && prov.Capabilities().SupportsRefund
// 建单态:渠道可退 → processing(即将调渠道 API);crypto 自托管等不可退渠道 →
// manual_pending(待运营人工向,钱未退)。两者都占用退款额度,都必须过守卫。
status := model.RefundManualPending
if channelRefundable {
status = model.RefundProcessing
}
r := &model.Refund{
RefundID: refundID, OutTradeNo: in.OutTradeNo, AttemptProviderRef: att.ProviderRef,
AmountMinor: in.AmountMinor, Currency: o.Currency, Reason: in.Reason,
Status: status, InitiatedBy: in.InitiatedBy,
}
ok, err := g.refunds.CreateRefundGuarded(r, o.AmountMinor) // paid amount = 订单行读到的金额,非调用方输入
if err != nil {
return nil, err
}
if !ok {
return nil, ErrRefundAmountInvalid
}
if !channelRefundable {
// crypto 自托管:无退款 API → 建"待人工"单即返回(钱未退,不动 order/不发事件)。
return &RefundResult{RefundID: refundID, Status: string(model.RefundManualPending)}, nil
}
// 渠道退款:refundID 作渠道幂等键。
refundRef, pstatus, rerr := rp.Refund(ctx, att.ProviderRef, refundID, in.AmountMinor, in.Reason)
if rerr != nil {
_, _ = g.refunds.MarkRefundStatus(refundID, model.RefundProcessing, model.RefundFailed, "", time.Now())
_ = g.enqueueRefundEvent(o, att, r, "refund.failed", "")
return &RefundResult{RefundID: refundID, Status: string(model.RefundFailed)}, rerr
}
switch pstatus {
case provider.PaidSucceeded:
if err := g.settleRefundSucceeded(o, att, r, refundRef); err != nil {
return nil, err
}
return &RefundResult{RefundID: refundID, Status: string(model.RefundSucceeded)}, nil
case provider.PaidFailed:
_, _ = g.refunds.MarkRefundStatus(refundID, model.RefundProcessing, model.RefundFailed, refundRef, time.Now())
_ = g.enqueueRefundEvent(o, att, r, "refund.failed", refundRef)
return &RefundResult{RefundID: refundID, Status: string(model.RefundFailed)}, nil
default: // PaidPending:stripe requires_action 等异步 → 保持 processing,交人工/P6 收敛
return &RefundResult{RefundID: refundID, Status: string(model.RefundProcessing)}, nil
}
}
// settleRefundSucceeded 复刻 settle 的「先入队后翻转」:先幂等入队 refund.succeeded
// (dedupe=refund_id,同单多次部分退不撞键),再翻转 refund 单 + 推进 order 退款态机。
// 顺序不变量:refund 为 succeeded ⇒ outbox 行必已存在。
func (g *Gateway) settleRefundSucceeded(o *model.OrderV2, att *model.Attempt, r *model.Refund, refundRef string) error {
if err := g.enqueueRefundEvent(o, att, r, "refund.succeeded", refundRef); err != nil {
return err
}
if _, err := g.refunds.MarkRefundStatus(r.RefundID, r.Status, model.RefundSucceeded, refundRef, time.Now()); err != nil {
return err
}
succeeded, err := g.refunds.RefundSum(o.OutTradeNo, model.RefundSucceeded)
if err != nil {
return err
}
_, err = g.orders.ApplyRefundToOrder(o.OutTradeNo, succeeded >= o.AmountMinor)
return err
}
func (g *Gateway) enqueueRefundEvent(o *model.OrderV2, att *model.Attempt, r *model.Refund, eventType, refundRef string) error {
if o.BizSystem == "" {
return nil // 独立收款无业务方回调
}
data := map[string]any{
"event_type": eventType,
"out_trade_no": o.OutTradeNo,
"biz_system": o.BizSystem,
"biz_ref": o.BizRef,
"product_biz_code": o.BizCode,
"refund_id": r.RefundID,
"provider_refund_ref": refundRef,
"amount_minor": r.AmountMinor,
"currency": r.Currency,
"channel": att.Channel,
"reason": r.Reason,
}
return g.webhook.Enqueue(o.OutTradeNo, o.BizSystem, eventType, r.RefundID, data)
}
// CompleteManualRefund 收口 crypto 人工退款:运营链上转账后回填 tx,走与渠道成功同一落地路径。
func (g *Gateway) CompleteManualRefund(_ context.Context, refundID, providerRefundRef string) (*RefundResult, error) {
r, err := g.refunds.GetRefund(refundID)
if err != nil {
return nil, err // ErrRefundNotFound
}
if r.Status != model.RefundManualPending {
return nil, ErrRefundNotManual
}
o, err := g.orders.GetOrder(r.OutTradeNo)
if err != nil {
return nil, err
}
att, err := g.orders.PaidAttempt(r.OutTradeNo)
if err != nil {
return nil, err
}
if err := g.settleRefundSucceeded(o, att, r, providerRefundRef); err != nil {
return nil, err
}
return &RefundResult{RefundID: refundID, Status: string(model.RefundSucceeded)}, nil
}
func (g *Gateway) ListManualPendingRefunds(limit int) ([]model.Refund, error) {
return g.refunds.ListManualPending(limit)
}
// RefundStatusView / GetRefund 供 HTTP 查询退款状态(Task 6)。
type RefundStatusView struct {
RefundID string `json:"refund_id"`
OutTradeNo string `json:"out_trade_no"`
AmountMinor int64 `json:"amount_minor"`
Currency string `json:"currency"`
Status string `json:"status"`
ProviderRefundRef string `json:"provider_refund_ref,omitempty"`
}
func (g *Gateway) GetRefund(refundID string) (*RefundStatusView, error) {
r, err := g.refunds.GetRefund(refundID)
if err != nil {
return nil, err
}
return &RefundStatusView{
RefundID: r.RefundID, OutTradeNo: r.OutTradeNo, AmountMinor: r.AmountMinor,
Currency: r.Currency, Status: string(r.Status), ProviderRefundRef: r.ProviderRefundRef,
}, nil
}
+117
View File
@@ -0,0 +1,117 @@
package gateway_test
import (
"context"
"errors"
"testing"
"github.com/wangjia/pay/internal/gateway"
"github.com/wangjia/pay/internal/model"
"github.com/wangjia/pay/internal/provider"
)
// 把 fake 订单推到 paid(复用 P2/P3 的 SyncPendingAttempts + SetQueryResult 路径)。
func createAndPay(t *testing.T, g *gateway.Gateway, fp interface {
SetQueryResult(string, provider.PaidEvent)
}, orders interface {
ListAttemptsByStatus(model.AttemptStatus, int) ([]model.Attempt, error)
}) string {
t.Helper()
res, err := g.CreateOrder(context.Background(), gateway.CreateOrderInput{
SKU: "pro_year", Method: "fake", BizSystem: "pangolin", BizRef: "u-1",
})
if err != nil {
t.Fatalf("create: %v", err)
}
atts, _ := orders.ListAttemptsByStatus(model.AttemptPending, 10)
fp.SetQueryResult(atts[0].ProviderRef, provider.PaidEvent{
ProviderRef: atts[0].ProviderRef, Status: provider.PaidSucceeded,
PaidAmountMinor: 29990000, PaidCurrency: "USDT",
})
if _, err := g.SyncPendingAttempts(context.Background(), 10); err != nil {
t.Fatalf("sync: %v", err)
}
return res.OrderNo
}
func TestRefundPartialThenFull(t *testing.T) {
g, fp, spy, orders := newGateway(t)
fp.EnableRefund("fake-refund-ref", provider.PaidSucceeded, nil) // fake 变可退渠道,同步成功
no := createAndPay(t, g, fp, orders)
// 部分退 1/3(总价 29990000)
r1, err := g.Refund(context.Background(), gateway.RefundInput{
OutTradeNo: no, AmountMinor: 9990000, Reason: "test", BizSystem: "pangolin",
})
if err != nil || r1.Status != string(model.RefundSucceeded) {
t.Fatalf("refund1 = %+v, %v", r1, err)
}
if o, _ := orders.GetOrder(no); o.Status != model.OrderPartRefundedV2 {
t.Fatalf("after partial: order = %s want partially_refunded", o.Status)
}
// 退款事件已入队(refund.succeeded,带 refund_id)
last := spy.calls[len(spy.calls)-1]
if last["event_type"] != "refund.succeeded" || last["refund_id"] != r1.RefundID {
t.Fatalf("refund webhook = %+v", last)
}
// 超退守卫:再退 25000000 > 剩余 20000000 → 拒
if _, err := g.Refund(context.Background(), gateway.RefundInput{
OutTradeNo: no, AmountMinor: 25000000, BizSystem: "pangolin",
}); !errors.Is(err, gateway.ErrRefundAmountInvalid) {
t.Fatalf("over-refund err = %v want ErrRefundAmountInvalid", err)
}
// 退剩余 → refunded
r2, err := g.Refund(context.Background(), gateway.RefundInput{
OutTradeNo: no, AmountMinor: 20000000, BizSystem: "pangolin",
})
if err != nil || r2.Status != string(model.RefundSucceeded) {
t.Fatalf("refund2 = %+v, %v", r2, err)
}
if o, _ := orders.GetOrder(no); o.Status != model.OrderRefundedV2 {
t.Fatalf("after full: order = %s want refunded", o.Status)
}
}
func TestRefundNotRefundableWhenPending(t *testing.T) {
g, _, _, _ := newGateway(t)
res, _ := g.CreateOrder(context.Background(), gateway.CreateOrderInput{SKU: "pro_year", Method: "fake"})
if _, err := g.Refund(context.Background(), gateway.RefundInput{OutTradeNo: res.OrderNo, AmountMinor: 1}); !errors.Is(err, gateway.ErrOrderNotRefundable) {
t.Fatalf("err = %v want ErrOrderNotRefundable", err)
}
}
func TestRefundManualForCryptoLikeChannel(t *testing.T) {
g, fp, spy, orders := newGateway(t)
// fp 不 EnableRefund → SupportsRefund=false → 建 manual_pending,不动 order、不发事件
no := createAndPay(t, g, fp, orders)
nCalls := len(spy.calls)
r, err := g.Refund(context.Background(), gateway.RefundInput{OutTradeNo: no, AmountMinor: 9990000, BizSystem: "pangolin"})
if err != nil || r.Status != string(model.RefundManualPending) {
t.Fatalf("manual refund = %+v, %v", r, err)
}
if o, _ := orders.GetOrder(no); o.Status != model.OrderPaidV2 {
t.Fatalf("manual pending 不应动 order,得 %s", o.Status)
}
if len(spy.calls) != nCalls {
t.Fatal("manual pending 不应发退款事件")
}
// 待办可捞
list, _ := g.ListManualPendingRefunds(10)
if len(list) != 1 || list[0].RefundID != r.RefundID {
t.Fatalf("manual list = %+v", list)
}
// 运营回填完成 → 成功落地 + 事件 + 态机
done, err := g.CompleteManualRefund(context.Background(), r.RefundID, "tron-tx-hash")
if err != nil || done.Status != string(model.RefundSucceeded) {
t.Fatalf("complete = %+v, %v", done, err)
}
if o, _ := orders.GetOrder(no); o.Status != model.OrderPartRefundedV2 {
t.Fatalf("after complete: order = %s want partially_refunded", o.Status)
}
last := spy.calls[len(spy.calls)-1]
if last["event_type"] != "refund.succeeded" || last["provider_refund_ref"] != "tron-tx-hash" {
t.Fatalf("complete webhook = %+v", last)
}
}
+2 -1
View File
@@ -137,6 +137,7 @@ func TestSettleEnqueueFailureKeepsOrderPending(t *testing.T) {
func TestSettleTransientReadErrorIsFailed(t *testing.T) {
db := model.OpenTestDB(t)
orders := store.NewOrderStore(db)
refunds := store.NewRefundStore(db)
preg := provider.NewRegistry()
fp := fake.New()
preg.Register(fp)
@@ -145,7 +146,7 @@ func TestSettleTransientReadErrorIsFailed(t *testing.T) {
})
picker := accounts.NewRouter(areg, nil, nil)
spy := &spyEnqueuer{}
g := gateway.New(orders, preg, picker, stubResolver{}, spy, "global")
g := gateway.New(orders, refunds, preg, picker, stubResolver{}, spy, "global")
ctx := context.Background()
g.CreateOrder(ctx, gateway.CreateOrderInput{SKU: "pro_year", Method: "fake", BizSystem: "pangolin", BizRef: "u-1"})
+8 -4
View File
@@ -41,14 +41,16 @@ func buildEngine(t *testing.T) *gin.Engine {
func buildEngineWithStore(t *testing.T) (*gin.Engine, *store.OrderStore) {
t.Helper()
gin.SetMode(gin.TestMode)
orders := store.NewOrderStore(model.OpenTestDB(t))
db := model.OpenTestDB(t)
orders := store.NewOrderStore(db)
refunds := store.NewRefundStore(db)
preg := provider.NewRegistry()
preg.Register(fake.New())
areg := accounts.New([]config.AccountConfig{
{AccountID: "fake-a1", Channel: "fake", Region: "global", Enabled: true, Weight: 1},
})
picker := accounts.NewRouter(areg, nil, nil)
g := gateway.New(orders, preg, picker, oneResolver{}, nopEnqueuer{}, "global")
g := gateway.New(orders, refunds, preg, picker, oneResolver{}, nopEnqueuer{}, "global")
r := gin.New()
router.SetupV2(r, g)
return r, orders
@@ -172,7 +174,9 @@ func TestV2CallbackMalformedBodyIs400Transient(t *testing.T) {
func TestV2RetryCurrencyMismatch409(t *testing.T) {
t.Helper()
gin.SetMode(gin.TestMode)
orders := store.NewOrderStore(model.OpenTestDB(t))
db := model.OpenTestDB(t)
orders := store.NewOrderStore(db)
refunds := store.NewRefundStore(db)
preg := provider.NewRegistry()
preg.Register(fake.New()) // method="fake", settles in "USDT"
@@ -185,7 +189,7 @@ func TestV2RetryCurrencyMismatch409(t *testing.T) {
{AccountID: "fake-a2", Channel: "fake_eur", Region: "global", Enabled: true, Weight: 1},
})
picker := accounts.NewRouter(areg, nil, nil)
g := gateway.New(orders, preg, picker, oneResolver{}, nopEnqueuer{}, "global")
g := gateway.New(orders, refunds, preg, picker, oneResolver{}, nopEnqueuer{}, "global")
r := gin.New()
router.SetupV2(r, g)
+26 -1
View File
@@ -20,6 +20,11 @@ import (
type Provider struct {
mu sync.Mutex
queryResults map[string]provider.PaidEvent
supportsRefund bool
refundRef string
refundStatus provider.PaidStatus
refundErr error
}
func New() *Provider { return &Provider{queryResults: map[string]provider.PaidEvent{}} }
@@ -29,7 +34,7 @@ func (p *Provider) Method() string { return "fake" }
func (p *Provider) Capabilities() provider.Capabilities {
return provider.Capabilities{
RenderTypes: []provider.RenderType{provider.RenderCryptoAddress},
SupportsRefund: false,
SupportsRefund: p.supportsRefund,
SettleCurrencies: []string{"USDT"},
Regions: []string{"global"},
}
@@ -89,3 +94,23 @@ func (p *Provider) SetQueryResult(providerRef string, ev provider.PaidEvent) {
defer p.mu.Unlock()
p.queryResults[providerRef] = ev
}
// EnableRefund 令 fake 表现为可退渠道并预置一次退款结果(测试缝)。
func (p *Provider) EnableRefund(refundRef string, status provider.PaidStatus, err error) {
p.mu.Lock()
defer p.mu.Unlock()
p.supportsRefund = true
p.refundRef = refundRef
p.refundStatus = status
p.refundErr = err
}
// Refund implements provider.RefundingProvider (测试缝:忽略入参,回放预置结果)。
func (p *Provider) Refund(_ context.Context, _, _ string, _ int64, _ string) (string, provider.PaidStatus, error) {
p.mu.Lock()
defer p.mu.Unlock()
if p.refundErr != nil {
return "", provider.PaidFailed, p.refundErr
}
return p.refundRef, p.refundStatus, nil
}
+4 -2
View File
@@ -109,10 +109,12 @@ type Provider interface {
Query(ctx context.Context, req QueryRequest) (*PaidEvent, error)
}
// RefundingProvider — 可选:支持渠道退款的 Provider 额外实现(P4;不支持则 capabilities=false)。
// RefundingProvider — 可选:支持渠道退款的 Provider 额外实现(P4)。refundID 为 pay 侧
// 退款单号,作渠道幂等键(alipay out_request_no / stripe Idempotency-Key)——同单多次
// 部分退款靠它去重,重试不重复退。不支持退款的渠道不实现本接口(capabilities=false)。
type RefundingProvider interface {
Provider
Refund(ctx context.Context, providerRef string, amountMinor int64, reason string) (refundRef string, status PaidStatus, err error)
Refund(ctx context.Context, providerRef, refundID string, amountMinor int64, reason string) (refundRef string, status PaidStatus, err error)
}
// RecurringProvider — 可选:支持自动续订(P8,设计 §5.1 4 类 kind)。
+2 -1
View File
@@ -40,6 +40,7 @@ func main() {
// v2 统一网关装配(P2 骨架,P3 配置驱动注册接线):provider 注册表 + gateway + webhook notifier。
orderStore := store.NewOrderStore(db)
refundStore := store.NewRefundStore(db)
webhookStore := store.NewWebhookStore(db)
notifier := webhook.NewNotifier(webhookStore, config.C.BizByName, func(no string) (bool, error) {
o, err := orderStore.GetOrder(no) // 投递门禁:订单已结算(paid 及退款态)才放行
@@ -55,7 +56,7 @@ func main() {
// P5 多账户路由:按 config.routing.<channel> 选策略(缺省 round_robin)。
// limit_aware 用量数据源 P6 对账就绪前用空源(NopUsage,退化为 round_robin)。
acctPicker := accounts.NewRouter(acctReg, config.C.Routing, accounts.NopUsage{})
gw := gateway.New(orderStore, pReg, acctPicker, productResolver, notifier, "cn")
gw := gateway.New(orderStore, refundStore, pReg, acctPicker, productResolver, notifier, "cn")
router.SetupV2(r, gw)
if config.C.QuerySync.Enabled {