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

404 lines
15 KiB
Go

package payment
import (
"bytes"
"context"
"crypto/hmac"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
bizpayment "kra/internal/biz/payment"
"net/http"
"strings"
"time"
"kra/internal/paymentkit"
)
type vendorProfile string
const (
vendorChinaums vendorProfile = "chinaums"
vendorSFT vendorProfile = "sft"
vendorSupperPay vendorProfile = "supper-pay"
vendorWechatGame vendorProfile = "wechat-game"
vendorDouyinGame vendorProfile = "douyin-game"
)
type vendorPaymentAdapter struct {
provider string
profile vendorProfile
}
func newVendorAdapter(provider string, profile vendorProfile) bizpayment.PaymentAdapter {
return &vendorPaymentAdapter{provider: provider, profile: profile}
}
func (a *vendorPaymentAdapter) Create(ctx context.Context, req *bizpayment.PaymentRequest, c map[string]any) (*bizpayment.PaymentResult, error) {
if req == nil {
return nil, errors.New("配置驱动支付下单请求为空")
}
payload := map[string]any{"merchant_id": text(c, "merchant_id"), "app_id": text(c, "app_id"), "trade_no": req.TradeNo, "subject": req.Subject, "amount": req.Amount, "currency": req.Currency, "notify_url": req.NotifyURL, "client_ip": req.ClientIP, "timestamp": time.Now().Unix(), "nonce": nonce()}
mergeMap(payload, req.Extra, "merchant_id", "app_id", "trade_no", "subject", "amount", "currency", "notify_url", "client_ip", "timestamp", "nonce")
return a.call(ctx, "create_url", payload, req.TradeNo, c)
}
func (a *vendorPaymentAdapter) Query(ctx context.Context, tradeNo string, c map[string]any) (*bizpayment.PaymentResult, error) {
return a.call(ctx, "query_url", map[string]any{"merchant_id": text(c, "merchant_id"), "app_id": text(c, "app_id"), "trade_no": tradeNo, "timestamp": time.Now().Unix(), "nonce": nonce()}, tradeNo, c)
}
func (a *vendorPaymentAdapter) Refund(ctx context.Context, req *bizpayment.PaymentRefundRequest, c map[string]any) (*bizpayment.PaymentResult, error) {
if req == nil {
return nil, errors.New("支付退款请求为空")
}
return a.call(ctx, "refund_url", map[string]any{"merchant_id": text(c, "merchant_id"), "app_id": text(c, "app_id"), "trade_no": req.TradeNo, "refund_no": req.RefundNo, "amount": req.Amount, "total_amount": req.TotalAmount, "currency": req.Currency, "timestamp": time.Now().Unix(), "nonce": nonce()}, req.TradeNo, c)
}
func (a *vendorPaymentAdapter) Callback(_ context.Context, callback *bizpayment.PaymentCallback, c map[string]any) (*bizpayment.PaymentResult, error) {
if callback == nil || len(callback.Body) == 0 {
return nil, errors.New("配置驱动支付回调为空")
}
fields := callbackFields(callback)
if err := a.verify(fields, callback.Body, callback.Headers, c); err != nil {
return nil, err
}
object := jsonObject(callback.Body)
statusField := text(c, "callback_status_field")
if statusField == "" {
statusField = text(c, "query_status_field")
}
tradeNoField := text(c, "callback_trade_no_field")
if tradeNoField == "" {
tradeNoField = text(c, "query_trade_no_field")
}
providerTradeNoField := text(c, "callback_provider_trade_no_field")
if providerTradeNoField == "" {
providerTradeNoField = text(c, "query_provider_trade_no_field")
}
state := first(fields, statusField, "status", "trade_status", "order_status", "pay_status")
tradeNo := first(fields, tradeNoField, "trade_no", "out_trade_no", "merchant_order_no", "cp_order_id")
providerTradeNo := first(fields, providerTradeNoField, "transaction_id", "platform_trade_no", "order_no")
if object != nil {
if value := stringAtPath(object, statusField); value != "" {
state = value
}
if value := stringAtPath(object, tradeNoField); value != "" {
tradeNo = value
}
if value := stringAtPath(object, providerTradeNoField); value != "" {
providerTradeNo = value
}
}
status := "pending"
successValues := configuredValues(c, "callback_success_values")
if len(successValues) == 0 {
successValues = configuredValues(c, "query_success_values")
}
if containsFold(successValues, state) {
status = "success"
}
payload, _ := json.Marshal(fields)
return &bizpayment.PaymentResult{Provider: a.provider, Status: status, TradeNo: tradeNo, ProviderTradeNo: providerTradeNo, Payload: payload}, nil
}
func (a *vendorPaymentAdapter) call(ctx context.Context, endpointKey string, payload map[string]any, tradeNo string, c map[string]any) (*bizpayment.PaymentResult, error) {
endpoint := text(c, endpointKey)
if endpoint == "" {
return nil, fmt.Errorf("%s 未配置 %s", a.provider, endpointKey)
}
secret := firstAny(c, "app_key", "merchant_key", "signing_secret", "token")
if secret == "" {
return nil, fmt.Errorf("%s 未配置签名密钥", a.provider)
}
raw, err := json.Marshal(payload)
if err != nil {
return nil, err
}
request, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(raw))
if err != nil {
return nil, err
}
request.Header.Set("Content-Type", "application/json")
a.signRequest(request, payload, raw, secret, c)
response, err := (&http.Client{Timeout: 20 * time.Second}).Do(request)
if err != nil {
return nil, err
}
defer response.Body.Close()
body, err := io.ReadAll(io.LimitReader(response.Body, 8<<20))
if err != nil {
return nil, err
}
if response.StatusCode >= 300 {
return nil, fmt.Errorf("%s HTTP %d", a.provider, response.StatusCode)
}
result := &bizpayment.PaymentResult{Provider: a.provider, Status: "created", TradeNo: tradeNo, Payload: ensureJSON(body)}
if endpointKey != "query_url" {
operation := strings.TrimSuffix(endpointKey, "_url")
object := jsonObject(body)
if object == nil {
if operation == "refund" {
return nil, fmt.Errorf("%s 退款响应不是 JSON 对象", a.provider)
}
} else {
statusField := text(c, operation+"_status_field")
if statusField == "" {
statusField = text(c, "query_status_field")
}
state := ""
if statusField != "" {
state = stringAtPath(object, statusField)
}
if state == "" {
for _, key := range []string{"status", "trade_status", "order_status", "pay_status", "refund_status"} {
if value := stringAtPath(object, key); value != "" {
state = value
break
}
}
}
if state != "" {
switch {
case containsFold(configuredValues(c, operation+"_failure_values"), state):
result.Status = "failed"
case containsFold(configuredValues(c, operation+"_success_values"), state):
result.Status = "created"
default:
if operation == "refund" {
normalized := normalizeRefundStatus(state, "")
switch normalized {
case "success", "pending", "failed":
result.Status = normalized
default:
return nil, fmt.Errorf("%s 退款响应状态无法识别: %s", a.provider, state)
}
} else {
result.Status = normalizePaymentStatus(state, "created")
}
}
} else if operation == "refund" {
return nil, fmt.Errorf("%s 退款响应缺少业务状态", a.provider)
}
}
return result, nil
}
object := jsonObject(body)
if object == nil {
return nil, fmt.Errorf("%s 查询响应不是 JSON 对象", a.provider)
}
state := stringAtPath(object, text(c, "query_status_field"))
if !containsFold(configuredValues(c, "query_success_values"), state) {
result.Status = normalizePaymentStatus(state, "pending")
return result, nil
}
result.Status = "success"
result.TradeNo = stringAtPath(object, text(c, "query_trade_no_field"))
result.ProviderTradeNo = stringAtPath(object, text(c, "query_provider_trade_no_field"))
result.Currency = strings.ToUpper(stringAtPath(object, text(c, "query_currency_field")))
amountText := stringAtPath(object, text(c, "query_amount_field"))
amountScale := configuredInt64(c, "query_amount_scale", 0)
if amountScale == 1 {
result.Amount, err = parseIntegerAmount(amountText)
} else {
result.Amount, err = parseDecimalAmount(amountText, amountScale)
}
if err != nil {
return nil, fmt.Errorf("解析 %s 查询金额: %w", a.provider, err)
}
if result.TradeNo == "" || result.ProviderTradeNo == "" || result.Currency == "" {
return nil, fmt.Errorf("%s 查询响应缺少订单号、平台单号或币种", a.provider)
}
if err = populateVendorBreakdown(result, object, c); err != nil {
return nil, fmt.Errorf("解析 %s 查询金额拆分: %w", a.provider, err)
}
return result, nil
}
func populateVendorBreakdown(result *bizpayment.PaymentResult, object map[string]any, c map[string]any) error {
if result == nil || result.Status != "success" {
return nil
}
defaultScale := configuredInt64(c, "query_amount_scale", 0)
payer, payerOK, err := parseConfiguredAmount(object, c, "query_payer_paid_amount_field", "query_payer_paid_amount_scale", defaultScale)
if err != nil {
return fmt.Errorf("payer paid amount: %w", err)
}
if !payerOK {
return nil
}
cash, cashOK, err := parseConfiguredAmount(object, c, "query_cash_paid_amount_field", "query_cash_paid_amount_scale", defaultScale)
if err != nil {
return fmt.Errorf("cash paid amount: %w", err)
}
point, pointOK, err := parseConfiguredAmount(object, c, "query_point_paid_amount_field", "query_point_paid_amount_scale", defaultScale)
if err != nil {
return fmt.Errorf("point paid amount: %w", err)
}
discount, discountOK, err := parseConfiguredAmount(object, c, "query_discount_amount_field", "query_discount_amount_scale", defaultScale)
if err != nil {
return fmt.Errorf("discount amount: %w", err)
}
providerDiscount, _, err := parseConfiguredAmount(object, c, "query_provider_discount_amount_field", "query_provider_discount_amount_scale", defaultScale)
if err != nil {
return fmt.Errorf("provider discount amount: %w", err)
}
merchantDiscount, _, err := parseConfiguredAmount(object, c, "query_merchant_discount_amount_field", "query_merchant_discount_amount_scale", defaultScale)
if err != nil {
return fmt.Errorf("merchant discount amount: %w", err)
}
settlement, _, err := parseConfiguredAmount(object, c, "query_settlement_amount_field", "query_settlement_amount_scale", defaultScale)
if err != nil {
return fmt.Errorf("settlement amount: %w", err)
}
if !cashOK && !pointOK {
cash = payer
cashOK = true
}
if !cashOK {
cash = payer - point
cashOK = cash >= 0
}
if !pointOK {
point = payer - cash
pointOK = point >= 0
}
if !discountOK {
discount = result.Amount - payer
discountOK = discount >= 0
}
if !cashOK || !pointOK || !discountOK || payer > result.Amount || cash+point != payer || payer+discount != result.Amount {
return errors.New("总额、用户实付、现金/积分和优惠金额不守恒")
}
result.PayerPaidAmount = payer
result.CashPaidAmount = cash
result.PointPaidAmount = point
result.DiscountAmount = discount
result.ProviderDiscountAmount = providerDiscount
result.MerchantDiscountAmount = merchantDiscount
result.SettlementAmount = settlement
result.PayerCurrency = strings.ToUpper(stringAtPath(object, text(c, "query_payer_currency_field")))
if result.PayerCurrency == "" {
result.PayerCurrency = result.Currency
}
result.AmountBreakdownKnown = true
return nil
}
func (a *vendorPaymentAdapter) signRequest(request *http.Request, payload map[string]any, raw []byte, secret string, c map[string]any) {
switch a.profile {
case vendorChinaums:
timestamp, nonceValue := fmt.Sprint(payload["timestamp"]), fmt.Sprint(payload["nonce"])
appID := text(c, "app_id")
digest := sha256.Sum256([]byte(appID + timestamp + nonceValue + string(raw) + secret))
request.Header.Set("Authorization", "OPEN-BODY-SIG AppId="+appID+", Timestamp="+timestamp+", Nonce="+nonceValue+", Signature="+hex.EncodeToString(digest[:]))
case vendorSFT:
request.Header.Set("X-SFT-Sign", paymentkit.MD5Canonical(payload, secret))
case vendorWechatGame:
request.Header.Set("X-Wechat-Game-Sign", paymentkit.HMACSHA256Hex(raw, secret, false))
if token := text(c, "access_token"); token != "" {
query := request.URL.Query()
query.Set("access_token", token)
request.URL.RawQuery = query.Encode()
}
case vendorDouyinGame:
request.Header.Set("X-TT-Pay-Sign", paymentkit.MD5Canonical(payload, secret))
default:
request.Header.Set("X-Payment-Sign", paymentkit.HMACSHA256Hex(raw, secret, false))
}
}
func (a *vendorPaymentAdapter) verify(fields map[string]string, raw []byte, headers map[string]string, c map[string]any) error {
secret := firstAny(c, "app_key", "merchant_key", "signing_secret", "token")
if secret == "" {
return errors.New("支付回调未配置签名密钥")
}
expected := first(fields, "sign", "signature")
if expected == "" {
expected = firstVendorHeader(headers, "X-Payment-Sign", "X-SFT-Sign", "X-Wechat-Game-Sign", "X-TT-Pay-Sign")
}
var actual string
switch a.profile {
case vendorChinaums:
// Chinaums uses the OPEN-BODY-SIG Authorization contract rather than
// the HMAC header used by the other configurable profiles. Keep the
// header parsing case-insensitive because net/http canonicalizes names
// but test/proxy callers do not necessarily do so.
auth := parseChinaumsAuthorization(firstVendorHeader(headers, "Authorization"))
if len(auth) == 0 {
return errors.New("Chinaums 回调缺少有效 Authorization")
}
expected = auth["signature"]
appID := strings.TrimSpace(auth["appid"])
configuredAppID := strings.TrimSpace(text(c, "app_id"))
if appID == "" || configuredAppID == "" || appID != configuredAppID {
return errors.New("Chinaums 回调 AppId 不匹配")
}
timestamp := strings.TrimSpace(auth["timestamp"])
nonceValue := strings.TrimSpace(auth["nonce"])
if expected == "" || timestamp == "" || nonceValue == "" {
return errors.New("Chinaums 回调签名信息不完整")
}
digest := sha256.Sum256([]byte(configuredAppID + timestamp + nonceValue + string(raw) + secret))
actual = hex.EncodeToString(digest[:])
case vendorSFT, vendorDouyinGame:
values := map[string]any{}
for key, value := range fields {
if key != "sign" && key != "signature" {
values[key] = value
}
}
actual = paymentkit.MD5Canonical(values, secret)
default:
actual = paymentkit.HMACSHA256Hex(raw, secret, false)
}
if !hmac.Equal([]byte(strings.ToLower(expected)), []byte(strings.ToLower(actual))) {
return errors.New("支付回调签名校验失败")
}
return nil
}
func firstVendorHeader(headers map[string]string, names ...string) string {
for _, name := range names {
for key, value := range headers {
if strings.EqualFold(strings.TrimSpace(key), name) {
if value = strings.TrimSpace(value); value != "" {
return value
}
}
}
}
return ""
}
func parseChinaumsAuthorization(value string) map[string]string {
result := map[string]string{}
value = strings.TrimSpace(value)
if value == "" {
return result
}
parts := strings.Fields(value)
if len(parts) < 2 || !strings.EqualFold(parts[0], "OPEN-BODY-SIG") {
return result
}
if index := strings.IndexAny(value, " \t"); index >= 0 {
value = strings.TrimSpace(value[index+1:])
}
for _, item := range strings.Split(value, ",") {
key, rawValue, ok := strings.Cut(strings.TrimSpace(item), "=")
if !ok {
continue
}
key = strings.ToLower(strings.TrimSpace(key))
rawValue = strings.Trim(strings.TrimSpace(rawValue), "\"")
if key != "" && rawValue != "" {
result[key] = rawValue
}
}
return result
}
func ensureJSON(raw []byte) []byte {
if json.Valid(raw) {
return raw
}
encoded, _ := json.Marshal(string(raw))
return encoded
}