kra-new/internal/integration/payment/allinpay.go

378 lines
13 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package payment
import (
"bytes"
"context"
"crypto"
"crypto/rsa"
"crypto/sha1"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"io"
bizpayment "kra/internal/biz/payment"
"kra/internal/paymentkit"
"net/http"
"strings"
"sync"
"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 *bizpayment.PaymentRequest, c map[string]any) (*bizpayment.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 &bizpayment.PaymentResult{Provider: bizpayment.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 &bizpayment.PaymentResult{Provider: bizpayment.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 &bizpayment.PaymentResult{Provider: bizpayment.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 := paymentkit.NormalizePaymentMethod(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) (*bizpayment.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 := &bizpayment.PaymentResult{Provider: bizpayment.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 *bizpayment.PaymentRefundRequest, c map[string]any) (*bizpayment.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 *bizpayment.PaymentRefundRequest, rsp *allinpay.RefundRsp, currency string) (*bizpayment.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 &bizpayment.PaymentResult{Provider: bizpayment.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, _ *bizpayment.PaymentCallback, _ map[string]any) (*bizpayment.PaymentResult, error) {
return nil, errors.New("通联支付回调没有可复用的 GoPay 验签器,请改用主动查单")
}