package payment import ( "bytes" "context" "crypto/hmac" "crypto/sha256" "encoding/hex" "encoding/json" "errors" "fmt" "io" "net/http" "strings" "time" "kra/internal/biz" "kra/internal/utils/paymentutil" ) 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) Adapter { return &vendorPaymentAdapter{provider: provider, profile: profile} } func (a *vendorPaymentAdapter) Create(ctx context.Context, req *biz.PaymentRequest, c map[string]any) (*biz.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) (*biz.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 *biz.PaymentRefundRequest, c map[string]any) (*biz.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 *biz.PaymentCallback, c map[string]any) (*biz.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 &biz.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) (*biz.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 := &biz.PaymentResult{Provider: a.provider, Status: "created", TradeNo: tradeNo, Payload: ensureJSON(body)} if endpointKey != "query_url" { 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 *biz.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", paymentutil.MD5Canonical(payload, secret)) case vendorWechatGame: request.Header.Set("X-Wechat-Game-Sign", paymentutil.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", paymentutil.MD5Canonical(payload, secret)) default: request.Header.Set("X-Payment-Sign", paymentutil.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 = paymentutil.MD5Canonical(values, secret) default: actual = paymentutil.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 }