378 lines
13 KiB
Go
378 lines
13 KiB
Go
package payment
|
||
|
||
import (
|
||
"bytes"
|
||
"context"
|
||
"crypto"
|
||
"crypto/rsa"
|
||
"crypto/sha1"
|
||
"encoding/base64"
|
||
"encoding/json"
|
||
"errors"
|
||
"fmt"
|
||
"io"
|
||
"net/http"
|
||
"strings"
|
||
"sync"
|
||
|
||
"kra/internal/biz"
|
||
|
||
"github.com/go-pay/crypto/xpem"
|
||
"github.com/go-pay/crypto/xrsa"
|
||
"github.com/go-pay/gopay"
|
||
"github.com/go-pay/gopay/allinpay"
|
||
)
|
||
|
||
type allinpayAdapter struct{}
|
||
|
||
type allinpayRoundTripperFunc func(*http.Request) (*http.Response, error)
|
||
|
||
func (f allinpayRoundTripperFunc) RoundTrip(req *http.Request) (*http.Response, error) {
|
||
return f(req)
|
||
}
|
||
|
||
type allinpayResponseCapture struct {
|
||
mu sync.Mutex
|
||
body []byte
|
||
}
|
||
|
||
func (c *allinpayResponseCapture) set(body []byte) {
|
||
c.mu.Lock()
|
||
c.body = append(c.body[:0], body...)
|
||
c.mu.Unlock()
|
||
}
|
||
|
||
func (c *allinpayResponseCapture) get() []byte {
|
||
c.mu.Lock()
|
||
defer c.mu.Unlock()
|
||
return append([]byte(nil), c.body...)
|
||
}
|
||
|
||
// captureAllinpayResponse preserves the exact JSON returned by the gateway.
|
||
// Query and Refund in GoPay v1.5.122 unmarshal and return without invoking its
|
||
// internal verifySign, so adapter-level verification must use the raw bytes
|
||
// rather than re-marshalling a response struct.
|
||
func captureAllinpayResponse(client *allinpay.Client) (*allinpayResponseCapture, error) {
|
||
if client == nil || client.GetHttpClient() == nil || client.GetHttpClient().HttpClient == nil {
|
||
return nil, errors.New("通联支付 HTTP 客户端为空")
|
||
}
|
||
hc := client.GetHttpClient()
|
||
base := hc.HttpClient.Transport
|
||
if base == nil {
|
||
base = http.DefaultTransport
|
||
}
|
||
capture := &allinpayResponseCapture{}
|
||
hc.SetTransport(allinpayRoundTripperFunc(func(req *http.Request) (*http.Response, error) {
|
||
res, err := base.RoundTrip(req)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
if res == nil || res.Body == nil {
|
||
return res, nil
|
||
}
|
||
body, readErr := io.ReadAll(res.Body)
|
||
_ = res.Body.Close()
|
||
if readErr != nil {
|
||
return nil, readErr
|
||
}
|
||
capture.set(body)
|
||
res.Body = io.NopCloser(bytes.NewReader(body))
|
||
return res, nil
|
||
}))
|
||
return capture, nil
|
||
}
|
||
|
||
func verifyAllinpayResponse(publicKey string, raw []byte) error {
|
||
if len(bytes.TrimSpace(raw)) == 0 {
|
||
return errors.New("通联支付响应签名校验失败: 响应为空")
|
||
}
|
||
values := gopay.BodyMap{}
|
||
if err := json.Unmarshal(raw, &values); err != nil {
|
||
return fmt.Errorf("通联支付响应签名校验失败: %w", err)
|
||
}
|
||
sign := strings.TrimSpace(values.GetString("sign"))
|
||
if sign == "" {
|
||
return errors.New("通联支付响应签名校验失败: 缺少 sign")
|
||
}
|
||
values.Remove("sign")
|
||
key, err := xpem.DecodePublicKey([]byte(xrsa.FormatAlipayPublicKey(publicKey)))
|
||
if err != nil {
|
||
return fmt.Errorf("通联支付公钥解析失败: %w", err)
|
||
}
|
||
signature, err := base64.StdEncoding.DecodeString(sign)
|
||
if err != nil {
|
||
return fmt.Errorf("通联支付响应签名编码无效: %w", err)
|
||
}
|
||
digest := sha1.Sum([]byte(values.EncodeAliPaySignParams()))
|
||
if err = rsa.VerifyPKCS1v15(key, crypto.SHA1, digest[:], signature); err != nil {
|
||
return fmt.Errorf("通联支付响应签名校验失败: %w", err)
|
||
}
|
||
return nil
|
||
}
|
||
|
||
func (a *allinpayAdapter) client(c map[string]any) (*allinpay.Client, error) {
|
||
client, err := allinpay.NewClient(text(c, "cus_id"), text(c, "app_id"), text(c, "private_key"), text(c, "public_key"), !strings.EqualFold(text(c, "environment"), "sandbox"))
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
if org := text(c, "org_id"); org != "" {
|
||
client.SetOrgId(org)
|
||
}
|
||
return client, nil
|
||
}
|
||
|
||
func (a *allinpayAdapter) Create(ctx context.Context, req *biz.PaymentRequest, c map[string]any) (*biz.PaymentResult, error) {
|
||
if req == nil {
|
||
return nil, errors.New("通联支付下单请求为空")
|
||
}
|
||
method, err := allinpayCreateMethod(req.Extra, c)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
orderType, err := allinpayOrderType(c)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
client, err := a.client(c)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
bm := gopay.BodyMap{"reqsn": req.TradeNo, "trxamt": req.Amount, "body": req.Subject}
|
||
mergeGoPayExtras(bm, req.Extra, "reqsn", "trxamt", "body", "method", "pay_method", "trade_type")
|
||
if method == "scan" {
|
||
bm.Set("authcode", firstAny(req.Extra, "authcode", "auth_code"))
|
||
bm.SetBodyMap("terminfo", func(info gopay.BodyMap) {
|
||
info.Set("devicetype", firstAnyOr(req.Extra, "10", "device_type", "devicetype"))
|
||
info.Set("termno", firstAnyOr(req.Extra, "00000001", "termno", "terminal_no"))
|
||
})
|
||
rsp, callErr := client.ScanPay(ctx, bm)
|
||
if callErr != nil {
|
||
return nil, callErr
|
||
}
|
||
if rsp == nil {
|
||
return nil, errors.New("通联支付扫码下单响应为空")
|
||
}
|
||
queryID, err := allinpayCreateQueryID(orderType, rsp.Trxid)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
if returnedTradeNo := strings.TrimSpace(rsp.Reqsn); returnedTradeNo != "" && returnedTradeNo != strings.TrimSpace(req.TradeNo) {
|
||
return nil, errors.New("通联支付扫码响应的 reqsn 不匹配")
|
||
}
|
||
return &biz.PaymentResult{Provider: biz.PaymentAllinPay, Status: normalizeAllinState(rsp.TrxStatus), TradeNo: req.TradeNo, ProviderTradeNo: strings.TrimSpace(rsp.Trxid), QueryID: queryID, Payload: mustJSON(rsp)}, nil
|
||
}
|
||
if method == "native" {
|
||
if orderType == allinpay.OrderTypeTrxId {
|
||
return nil, errors.New("通联支付 Native 下单不支持 query_order_type=trxid")
|
||
}
|
||
bm.Set("expiretime", firstAny(req.Extra, "expiretime", "expire_time"))
|
||
rsp, callErr := client.NativePay(ctx, bm)
|
||
if callErr != nil {
|
||
return nil, callErr
|
||
}
|
||
if rsp == nil {
|
||
return nil, errors.New("通联支付下单响应为空")
|
||
}
|
||
if returnedTradeNo := strings.TrimSpace(rsp.ReqSn); returnedTradeNo != "" && returnedTradeNo != strings.TrimSpace(req.TradeNo) {
|
||
return nil, errors.New("通联支付 Native 响应的 reqsn 不匹配")
|
||
}
|
||
return &biz.PaymentResult{Provider: biz.PaymentAllinPay, Status: "created", TradeNo: req.TradeNo, Payload: mustJSON(rsp)}, nil
|
||
}
|
||
payType := firstAny(req.Extra, "paytype", "pay_type")
|
||
if payType == "" {
|
||
payType = firstAny(c, "paytype", "pay_type")
|
||
}
|
||
if payType == "" {
|
||
payType = allinpay.PayTypeWXJS
|
||
}
|
||
bm.Set("paytype", payType)
|
||
rsp, err := client.Pay(ctx, bm)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
if rsp == nil {
|
||
return nil, errors.New("通联支付下单响应为空")
|
||
}
|
||
if returnedTradeNo := strings.TrimSpace(rsp.Reqsn); returnedTradeNo != "" && returnedTradeNo != strings.TrimSpace(req.TradeNo) {
|
||
return nil, errors.New("通联支付下单响应的 reqsn 不匹配")
|
||
}
|
||
queryID, err := allinpayCreateQueryID(orderType, rsp.Trxid)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
return &biz.PaymentResult{Provider: biz.PaymentAllinPay, Status: "created", TradeNo: req.TradeNo, ProviderTradeNo: strings.TrimSpace(rsp.Trxid), QueryID: queryID, Payload: mustJSON(rsp)}, nil
|
||
}
|
||
|
||
func allinpayCreateMethod(extra, config map[string]any) (string, error) {
|
||
value := firstAny(extra, "method", "pay_method", "trade_type")
|
||
if value == "" {
|
||
value = firstAny(config, "method", "pay_method", "trade_type")
|
||
}
|
||
normalized := strings.NewReplacer(".", "_", "-", "_", " ", "_").Replace(strings.ToLower(strings.TrimSpace(value)))
|
||
switch normalized {
|
||
case "", "pay", "unified", "unified_pay":
|
||
return "pay", nil
|
||
case "scan", "scan_pay", "micropay", "micro_pay", "barcode", "barcode_pay":
|
||
return "scan", nil
|
||
case "native", "native_pay", "qr", "qrcode":
|
||
return "native", nil
|
||
default:
|
||
return "", fmt.Errorf("通联支付不支持的下单方式: %s", value)
|
||
}
|
||
}
|
||
|
||
func (a *allinpayAdapter) Query(ctx context.Context, tradeNo string, c map[string]any) (*biz.PaymentResult, error) {
|
||
orderType, err := allinpayOrderType(c)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
tradeNo = strings.TrimSpace(tradeNo)
|
||
if tradeNo == "" {
|
||
return nil, errors.New("通联支付查单缺少订单号")
|
||
}
|
||
client, err := a.client(c)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
capture, err := captureAllinpayResponse(client)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
rsp, err := client.Query(ctx, orderType, tradeNo)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
if err = verifyAllinpayResponse(text(c, "public_key"), capture.get()); err != nil {
|
||
return nil, err
|
||
}
|
||
if rsp == nil {
|
||
return nil, errors.New("通联支付查单响应为空")
|
||
}
|
||
result := &biz.PaymentResult{Provider: biz.PaymentAllinPay, Status: normalizeAllinState(rsp.TrxStatus), TradeNo: strings.TrimSpace(rsp.Reqsn), ProviderTradeNo: strings.TrimSpace(rsp.Trxid), Currency: strings.ToUpper(firstAny(c, "currency")), Payload: mustJSON(rsp)}
|
||
if strings.EqualFold(orderType, allinpay.OrderTypeTrxId) {
|
||
if result.ProviderTradeNo == "" {
|
||
return nil, errors.New("通联支付查单响应缺少 trxid")
|
||
}
|
||
if result.ProviderTradeNo != tradeNo {
|
||
return nil, errors.New("通联支付查单响应的 trxid 不匹配")
|
||
}
|
||
result.QueryID = tradeNo
|
||
if strings.TrimSpace(result.TradeNo) == "" {
|
||
return nil, errors.New("通联支付查单响应缺少商户订单号")
|
||
}
|
||
} else if result.TradeNo == "" {
|
||
result.TradeNo = tradeNo
|
||
} else if result.TradeNo != tradeNo {
|
||
return nil, errors.New("通联支付查单响应的 reqsn 不匹配")
|
||
}
|
||
if rsp.TrxAmt != "" {
|
||
result.Amount, err = parseIntegerAmount(rsp.TrxAmt)
|
||
if err != nil {
|
||
if scale := configuredInt64(c, "amount_scale", 100); scale != 1 {
|
||
result.Amount, err = parseDecimalAmount(rsp.TrxAmt, scale)
|
||
}
|
||
}
|
||
}
|
||
if result.Currency == "" {
|
||
result.Currency = "CNY"
|
||
}
|
||
return result, err
|
||
}
|
||
|
||
func allinpayOrderType(c map[string]any) (string, error) {
|
||
orderType := strings.ToLower(strings.TrimSpace(firstAny(c, "query_order_type", "order_type")))
|
||
if orderType == "" {
|
||
return allinpay.OrderTypeReqSN, nil
|
||
}
|
||
if orderType != allinpay.OrderTypeReqSN && orderType != allinpay.OrderTypeTrxId {
|
||
return "", fmt.Errorf("通联支付 query_order_type 必须是 %s 或 %s", allinpay.OrderTypeReqSN, allinpay.OrderTypeTrxId)
|
||
}
|
||
return orderType, nil
|
||
}
|
||
|
||
func allinpayCreateQueryID(orderType, transactionID string) (string, error) {
|
||
if orderType == allinpay.OrderTypeTrxId {
|
||
queryID := strings.TrimSpace(transactionID)
|
||
if queryID == "" {
|
||
return "", errors.New("通联支付下单响应缺少 trxid,无法按 trxid 查单")
|
||
}
|
||
return queryID, nil
|
||
}
|
||
return "", nil
|
||
}
|
||
|
||
func (a *allinpayAdapter) Refund(ctx context.Context, req *biz.PaymentRefundRequest, c map[string]any) (*biz.PaymentResult, error) {
|
||
if req == nil {
|
||
return nil, errors.New("通联支付退款请求为空")
|
||
}
|
||
tradeNo := strings.TrimSpace(req.TradeNo)
|
||
refundNo := strings.TrimSpace(req.RefundNo)
|
||
if tradeNo == "" || refundNo == "" || req.Amount <= 0 {
|
||
return nil, errors.New("通联支付退款缺少商户订单号、退款单号或有效金额")
|
||
}
|
||
orderType, err := allinpayOrderType(c)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
queryID := strings.TrimSpace(req.QueryID)
|
||
if orderType == allinpay.OrderTypeTrxId && queryID == "" {
|
||
return nil, errors.New("通联支付按 trxid 退款缺少持久化 QueryID")
|
||
}
|
||
client, err := a.client(c)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
capture, err := captureAllinpayResponse(client)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
bm := gopay.BodyMap{"reqsn": refundNo, "trxamt": req.Amount, "remark": firstAny(c, "refund_remark", "remark")}
|
||
if orderType == allinpay.OrderTypeTrxId {
|
||
bm.Set("oldtrxid", queryID)
|
||
} else {
|
||
bm.Set("oldreqsn", tradeNo)
|
||
}
|
||
mergeGoPayConfigExtras(bm, c, "refund_extra", "reqsn", "trxamt", "oldreqsn", "oldtrxid", "remark")
|
||
rsp, err := client.Refund(ctx, bm)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
if err = verifyAllinpayResponse(text(c, "public_key"), capture.get()); err != nil {
|
||
return nil, err
|
||
}
|
||
effective := *req
|
||
effective.TradeNo = tradeNo
|
||
effective.RefundNo = refundNo
|
||
return allinpayRefundResult(&effective, rsp, firstAny(c, "currency"))
|
||
}
|
||
|
||
func allinpayRefundResult(req *biz.PaymentRefundRequest, rsp *allinpay.RefundRsp, currency string) (*biz.PaymentResult, error) {
|
||
if rsp == nil {
|
||
return nil, errors.New("通联支付退款响应为空")
|
||
}
|
||
if returnedRefundNo := strings.TrimSpace(rsp.Reqsn); returnedRefundNo == "" {
|
||
return nil, errors.New("通联支付退款响应缺少 reqsn")
|
||
} else if returnedRefundNo != req.RefundNo {
|
||
return nil, errors.New("通联支付退款响应的 reqsn 不匹配")
|
||
}
|
||
providerRefundID := strings.TrimSpace(rsp.Trxid)
|
||
if providerRefundID == "" {
|
||
return nil, errors.New("通联支付退款响应缺少 trxid")
|
||
}
|
||
if strings.TrimSpace(rsp.Fee) == "" {
|
||
return nil, errors.New("通联支付退款响应缺少 fee")
|
||
}
|
||
refundAmount, err := parseIntegerAmount(rsp.Fee)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("解析通联支付退款金额: %w", err)
|
||
}
|
||
if refundAmount != req.Amount {
|
||
return nil, errors.New("通联支付退款响应金额不匹配")
|
||
}
|
||
return &biz.PaymentResult{Provider: biz.PaymentAllinPay, Status: normalizeRefundState(rsp.TrxStatus), TradeNo: req.TradeNo, ProviderTradeNo: providerRefundID, Amount: refundAmount, Currency: strings.ToUpper(strings.TrimSpace(currency)), Payload: mustJSON(rsp)}, nil
|
||
}
|
||
|
||
func (a *allinpayAdapter) Callback(_ context.Context, _ *biz.PaymentCallback, _ map[string]any) (*biz.PaymentResult, error) {
|
||
return nil, errors.New("通联支付回调没有可复用的 GoPay 验签器,请改用主动查单")
|
||
}
|