kra-new/internal/data/payment/payment.go

519 lines
20 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 (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
integrationbiz "kra/internal/biz/integration"
bizpayment "kra/internal/biz/payment"
"kra/internal/paymentkit"
"strconv"
"strings"
"time"
"github.com/google/uuid"
)
type paymentRepo struct {
data Provider
config integrationbiz.PaymentConfigReader
factory bizpayment.PaymentAdapterFactory
orders bizpayment.PaymentOrderRepo
}
func NewPaymentRepo(data Provider, config integrationbiz.PaymentConfigReader, factory bizpayment.PaymentAdapterFactory, orders bizpayment.PaymentOrderRepo) bizpayment.PaymentRepo {
return &paymentRepo{data: data, config: config, factory: factory, orders: orders}
}
func (r *paymentRepo) values(ctx context.Context, provider string) (map[string]any, error) {
if r == nil || r.config == nil {
return nil, errors.New("支付配置仓储未接入")
}
config, err := r.config.ReadPaymentConfig(ctx, provider)
if err != nil {
if errors.Is(err, integrationbiz.ErrPaymentConfigNotFound) {
return nil, bizpayment.ErrPaymentProviderNotFound
}
return nil, err
}
if config == nil || !config.Enabled {
return nil, fmt.Errorf("支付渠道 %s 未启用", provider)
}
values := map[string]any{}
if err := json.Unmarshal(config.Values, &values); err != nil {
return nil, fmt.Errorf("支付配置格式错误: %w", err)
}
return values, nil
}
func (r *paymentRepo) adapter(ctx context.Context, provider string) (bizpayment.PaymentAdapter, map[string]any, error) {
values, err := r.values(ctx, provider)
if err != nil {
return nil, nil, err
}
if r.factory == nil {
return nil, nil, errors.New("支付渠道适配器未接入")
}
adapter, err := r.factory.New(provider)
return adapter, values, err
}
func (r *paymentRepo) TestProvider(ctx context.Context, provider string) (*bizpayment.PaymentTestResult, error) {
started := time.Now()
test := &bizpayment.PaymentTestResult{Provider: provider, TradeNo: "", Passed: false, Stages: []bizpayment.PaymentTestStage{}}
add := func(name, status, message, tradeNo string, since time.Time) {
test.Stages = append(test.Stages, bizpayment.PaymentTestStage{Name: name, Status: status, Message: message, TradeNo: tradeNo, Duration: time.Since(since).Milliseconds()})
}
values, err := r.testRow(ctx, provider)
if err != nil {
add("config", "failed", err.Error(), "", started)
return test, err
}
test.Mode = strings.ToLower(strings.TrimSpace(paymentkit.Text(values, "environment")))
configStart := time.Now()
if err = integrationbiz.ValidateIntegrationConfig(integrationbiz.IntegrationKindPayment, provider, values); err != nil {
add("config", "failed", err.Error(), "", configStart)
return test, err
}
add("config", "passed", "支付配置校验通过", "", configStart)
if err = validatePaymentTestSettings(provider, values); err != nil {
add("test_settings", "failed", err.Error(), "", time.Now())
return test, err
}
if r.factory == nil {
err = errors.New("支付渠道适配器未接入")
}
var adapter bizpayment.PaymentAdapter
if err == nil {
adapter, err = r.factory.New(provider)
}
if err != nil {
add("adapter", "failed", err.Error(), "", time.Now())
return test, err
}
if r.orders == nil {
err = errors.New("支付订单仓储未接入")
add("local_order", "failed", err.Error(), "", time.Now())
return test, err
}
req := paymentTestRequest(provider, values)
test.TradeNo = req.TradeNo
extra, _ := json.Marshal(req.Extra)
localStart := time.Now()
order, _, err := r.orders.CreatePaymentOrder(ctx, &bizpayment.PaymentOrder{
TradeNo: req.TradeNo, Provider: provider, BusinessType: req.BusinessType, BusinessID: req.BusinessID,
Subject: req.Subject, PaymentMode: bizpayment.PaymentModeExternal, OriginalAmount: req.Amount, Amount: req.Amount,
Currency: req.Currency, PaymentStatus: bizpayment.PaymentStatusInitialized, FulfillmentStatus: bizpayment.FulfillmentStatusPending,
RefundStatus: bizpayment.RefundStatusNone, ConfirmationID: uuid.NewString(), RequestFingerprint: paymentTestFingerprint(req), Extra: extra,
})
if err != nil {
add("local_order", "failed", err.Error(), req.TradeNo, localStart)
return test, err
}
add("local_order", "passed", "本地测试订单已创建", req.TradeNo, localStart)
createStart := time.Now()
created, err := adapter.Create(ctx, req, values)
if err != nil {
recordPaymentTestError(ctx, r.data, provider, req.TradeNo, err)
add("create", "failed", err.Error(), req.TradeNo, createStart)
return test, err
}
if created == nil {
err = errors.New("测试下单响应为空")
add("create", "failed", err.Error(), req.TradeNo, createStart)
return test, err
}
if err = validatePaymentTestResult(provider, req.TradeNo, created); err != nil {
recordPaymentTestError(ctx, r.data, provider, req.TradeNo, err)
add("create", "failed", err.Error(), req.TradeNo, createStart)
return test, err
}
if order, err = r.orders.RecordPaymentCreate(ctx, provider, req.TradeNo, paymentTestProviderUpdate(created)); err != nil {
add("local_order", "failed", "记录第三方下单结果失败: "+err.Error(), req.TradeNo, createStart)
return test, err
}
add("create", "passed", "测试订单已提交", req.TradeNo, createStart)
test.Result = created
queryID := strings.TrimSpace(created.QueryID)
if queryID == "" {
queryID = strings.TrimSpace(created.ProviderTradeNo)
}
if queryID == "" {
queryID = req.TradeNo
}
if provider == bizpayment.PaymentApple {
queryID = strings.TrimSpace(paymentkit.Text(values, "test_transaction_id"))
if queryID == "" {
err = errors.New("Apple 连通性测试需要配置 test_transaction_id沙箱交易 ID")
add("query", "failed", err.Error(), req.TradeNo, time.Now())
return test, err
}
}
queryStart := time.Now()
queried, queryErr := queryPaymentTest(ctx, adapter, queryID, values)
if queryErr != nil {
recordPaymentTestError(ctx, r.data, provider, req.TradeNo, queryErr)
add("query", "failed", queryErr.Error(), req.TradeNo, queryStart)
return test, queryErr
}
if err = validatePaymentTestResult(provider, req.TradeNo, queried); err != nil {
recordPaymentTestError(ctx, r.data, provider, req.TradeNo, err)
add("query", "failed", err.Error(), req.TradeNo, queryStart)
return test, err
}
if provider != bizpayment.PaymentApple {
if order, err = r.orders.ApplyPaymentResult(ctx, provider, req.TradeNo, paymentTestProviderUpdate(queried)); err != nil {
add("local_order", "failed", "回写测试查单结果失败: "+err.Error(), req.TradeNo, queryStart)
return test, err
}
}
if queried == nil {
err = errors.New("测试查单响应为空")
add("query", "failed", err.Error(), req.TradeNo, queryStart)
return test, err
}
test.Result = queried
add("query", "passed", "测试订单查询成功,状态: "+queried.Status, req.TradeNo, queryStart)
if queried.Status != "success" || provider == bizpayment.PaymentApple {
message := "订单尚未支付成功,已完成配置、下单和查单连通性测试;请在沙箱完成付款后重试"
if provider == bizpayment.PaymentApple {
message = "Apple 退款由 App Store 管理,已完成配置、下单和交易查询测试"
}
add("refund", "skipped", message, req.TradeNo, time.Now())
test.Passed = true
return test, nil
}
refundStart := time.Now()
order, refundToken, beginErr := r.orders.BeginPaymentRefund(ctx, provider, req.TradeNo, req.Amount, time.Minute)
if beginErr != nil {
add("refund", "failed", beginErr.Error(), req.TradeNo, refundStart)
return test, beginErr
}
refund, refundErr := adapter.Refund(ctx, &bizpayment.PaymentRefundRequest{Provider: provider, TradeNo: req.TradeNo, ProviderTradeNo: order.ProviderTradeNo, QueryID: order.QueryID, RefundNo: order.RefundNo, Amount: req.Amount, TotalAmount: req.Amount, Currency: req.Currency}, values)
if refundErr != nil {
recordPaymentTestError(ctx, r.data, provider, req.TradeNo, refundErr)
add("refund", "failed", refundErr.Error(), req.TradeNo, refundStart)
return test, refundErr
}
if refund == nil {
err = errors.New("测试退款响应为空")
add("refund", "failed", err.Error(), req.TradeNo, refundStart)
return test, err
}
if _, err = r.orders.CompletePaymentRefundRequest(ctx, provider, req.TradeNo, refundToken, true, ""); err != nil {
add("local_order", "failed", "回写测试退款结果失败: "+err.Error(), req.TradeNo, refundStart)
return test, err
}
test.Result = refund
add("refund", "passed", "测试退款申请已被渠道接受", req.TradeNo, refundStart)
test.Passed = true
test.FullFlow = true
return test, nil
}
func paymentTestFingerprint(req *bizpayment.PaymentRequest) string {
raw, _ := json.Marshal(req)
hash := sha256.Sum256(raw)
return hex.EncodeToString(hash[:])
}
func paymentTestProviderUpdate(result *bizpayment.PaymentResult) *bizpayment.PaymentProviderUpdate {
if result == nil {
return nil
}
return &bizpayment.PaymentProviderUpdate{
Status: result.Status, ProviderStatus: result.Status, ProviderTradeNo: result.ProviderTradeNo, QueryID: result.QueryID,
Amount: result.Amount, PayerPaidAmount: result.PayerPaidAmount, CashPaidAmount: result.CashPaidAmount,
PointPaidAmount: result.PointPaidAmount, DiscountAmount: result.DiscountAmount,
ProviderDiscountAmount: result.ProviderDiscountAmount, MerchantDiscountAmount: result.MerchantDiscountAmount,
SettlementAmount: result.SettlementAmount, Currency: result.Currency, PayerCurrency: result.PayerCurrency,
AmountBreakdownKnown: result.AmountBreakdownKnown, CreatePayload: result.Payload,
}
}
func validatePaymentTestResult(provider, tradeNo string, result *bizpayment.PaymentResult) error {
if result == nil {
return errors.New("支付渠道响应为空")
}
if strings.TrimSpace(result.Provider) != provider {
return errors.New("支付渠道响应的 provider 不匹配")
}
if value := strings.TrimSpace(result.TradeNo); provider != bizpayment.PaymentApple && value != "" && value != tradeNo {
return errors.New("支付渠道响应的商户订单号不匹配")
}
return nil
}
func queryPaymentTest(ctx context.Context, adapter bizpayment.PaymentAdapter, queryID string, values map[string]any) (*bizpayment.PaymentResult, error) {
var result *bizpayment.PaymentResult
var err error
for attempt := 0; attempt < 3; attempt++ {
result, err = adapter.Query(ctx, queryID, values)
if err == nil {
return result, nil
}
if attempt == 2 {
break
}
timer := time.NewTimer(500 * time.Millisecond)
select {
case <-ctx.Done():
timer.Stop()
return nil, ctx.Err()
case <-timer.C:
}
}
return nil, err
}
func recordPaymentTestError(ctx context.Context, data Provider, provider, tradeNo string, err error) {
if data == nil || data.DB() == nil || err == nil {
return
}
message := err.Error()
if len(message) > 512 {
message = message[:512]
}
_ = data.DB().WithContext(ctx).Model(&paymentOrderPO{}).Where("provider = ? AND trade_no = ?", provider, tradeNo).Update("last_error", message).Error
}
func (r *paymentRepo) testRow(ctx context.Context, provider string) (map[string]any, error) {
if r == nil || r.config == nil {
return nil, errors.New("支付配置仓储未接入")
}
config, err := r.config.ReadPaymentConfig(ctx, provider)
if err != nil {
if errors.Is(err, integrationbiz.ErrPaymentConfigNotFound) {
return nil, bizpayment.ErrPaymentProviderNotFound
}
return nil, err
}
if config == nil {
return nil, errors.New("支付配置为空")
}
values := map[string]any{}
if err := json.Unmarshal(config.Values, &values); err != nil || values == nil {
return nil, errors.New("支付配置格式错误")
}
return values, nil
}
func paymentTestRequest(provider string, values map[string]any) *bizpayment.PaymentRequest {
tradeNo := "kra-test-" + time.Now().UTC().Format("20060102150405.000000000")
amount := paymentkit.ConfiguredInt64(values, "test_amount", 1)
if amount <= 0 {
amount = 1
}
req := &bizpayment.PaymentRequest{Provider: provider, TradeNo: strings.ReplaceAll(tradeNo, ".", ""), Subject: "Kra 支付渠道连通性测试", Amount: amount, Currency: strings.ToUpper(paymentkit.FirstText(values, "test_currency", "currency", "fee_type")), NotifyURL: paymentkit.Text(values, "notify_url"), ReturnURL: paymentkit.Text(values, "return_url"), BusinessType: "system_payment_test", BusinessID: uuid.NewString(), Extra: map[string]any{}}
if req.Currency == "" {
req.Currency = "CNY"
}
if provider == bizpayment.PaymentApple {
req.TradeNo = uuid.NewString()
req.Extra["product_id"] = paymentkit.FirstText(values, "product_id", "test_product_id")
}
if raw := strings.TrimSpace(paymentkit.Text(values, "test_extra")); raw != "" {
var extra map[string]any
if json.Unmarshal([]byte(raw), &extra) == nil {
for key, value := range extra {
req.Extra[key] = value
}
}
}
for _, key := range []string{"trade_type", "method", "pay_type", "channel", "openid", "open_id", "auth_code", "authcode", "barcode"} {
if value := paymentkit.Text(values, key); value != "" {
req.Extra[key] = value
}
}
return req
}
func validatePaymentTestSettings(provider string, values map[string]any) error {
if !testModeEnabled(values) {
return errors.New("请先打开 test_mode允许执行渠道测试")
}
if raw := strings.TrimSpace(paymentkit.Text(values, "test_extra")); raw != "" {
var extra map[string]any
if err := json.Unmarshal([]byte(raw), &extra); err != nil {
return fmt.Errorf("test_extra 必须是 JSON 对象: %w", err)
}
}
if provider == bizpayment.PaymentApple && strings.TrimSpace(paymentkit.Text(values, "test_transaction_id")) == "" {
return errors.New("Apple 测试需要 test_transaction_id沙箱交易 ID")
}
if provider == bizpayment.PaymentApple && strings.TrimSpace(paymentkit.FirstText(values, "test_product_id", "product_id")) == "" {
return errors.New("Apple 测试需要 test_product_id沙箱商品 ID")
}
return nil
}
func testModeEnabled(values map[string]any) bool {
value, exists := values["test_mode"]
if !exists {
return false
}
switch typed := value.(type) {
case bool:
return typed
case string:
return strings.EqualFold(strings.TrimSpace(typed), "true") || typed == "1"
case float64:
return typed == 1
default:
return false
}
}
func (r *paymentRepo) Create(ctx context.Context, req *bizpayment.PaymentRequest) (*bizpayment.PaymentResult, error) {
if req == nil {
return nil, errors.New("支付下单请求为空")
}
a, c, err := r.adapter(ctx, req.Provider)
if err != nil {
return nil, err
}
effective := *req
effective.NotifyURL = strings.TrimSpace(paymentkit.Text(c, "notify_url"))
effective.ReturnURL = strings.TrimSpace(paymentkit.Text(c, "return_url"))
if paymentCreateRequiresNotifyURL(req.Provider, req.Extra, c) && effective.NotifyURL == "" {
return nil, fmt.Errorf("支付渠道 %s 未配置服务端 notify_url", req.Provider)
}
return a.Create(ctx, &effective, c)
}
func paymentProviderRequiresNotifyURL(provider string) bool {
switch provider {
case bizpayment.PaymentApple, bizpayment.PaymentAllinPay, bizpayment.PaymentSaobei, bizpayment.PaymentPayPal:
return false
default:
return true
}
}
// paymentCreateRequiresNotifyURL keeps the repository-level URL guard aligned
// with the selected provider operation. Synchronous barcode/retail APIs return
// their execution result directly and do not consume notify_url; redirect and
// client-side prepay APIs still require the configured server callback URL.
func paymentCreateRequiresNotifyURL(provider string, extra, config map[string]any) bool {
if !paymentProviderRequiresNotifyURL(provider) {
return false
}
keys := []string{"method", "pay_method", "trade_type", "pay_type", "channel"}
switch provider {
case bizpayment.PaymentAlipay, bizpayment.PaymentAlipayV3:
keys = []string{"method", "pay_method", "trade_type", "channel"}
case bizpayment.PaymentWechatV2:
keys = []string{"trade_type", "pay_type", "method", "pay_method", "channel"}
case bizpayment.PaymentWechatV3:
keys = []string{"trade_type", "pay_type", "method"}
case bizpayment.PaymentQQ:
keys = []string{"trade_type", "pay_type", "method", "pay_method"}
case bizpayment.PaymentLakala:
keys = []string{"method", "pay_method", "trade_type"}
}
value := paymentkit.FirstText(extra, keys...)
if value == "" {
value = paymentkit.FirstText(config, keys...)
}
normalized := strings.NewReplacer(".", "_", "-", "_", " ", "_").Replace(strings.ToLower(strings.TrimSpace(value)))
switch provider {
case bizpayment.PaymentAlipay, bizpayment.PaymentAlipayV3:
return !paymentkit.ContainsFold([]string{"pay", "trade_pay", "alipay_trade_pay", "barcode", "barcode_pay", "micropay", "face_to_face"}, normalized)
case bizpayment.PaymentWechatV2:
return !paymentkit.ContainsFold([]string{"micropay", "micro_pay", "barcode", "barcode_pay", "pay_code", "payment_code"}, normalized)
case bizpayment.PaymentWechatV3:
return !paymentkit.ContainsFold([]string{"micropay", "micro_pay", "codepay", "code_pay", "barcode", "barcode_pay", "facepay", "face_pay"}, normalized)
case bizpayment.PaymentQQ:
return !paymentkit.ContainsFold([]string{"micropay", "micro_pay", "barcode", "barcode_pay"}, normalized)
case bizpayment.PaymentLakala:
return !paymentkit.ContainsFold([]string{"retail", "retail_pay", "micropay", "barcode"}, normalized)
default:
return true
}
}
func (r *paymentRepo) Query(ctx context.Context, provider, tradeNo string) (*bizpayment.PaymentResult, error) {
a, c, err := r.adapter(ctx, provider)
if err != nil {
return nil, err
}
return a.Query(ctx, tradeNo, c)
}
func (r *paymentRepo) Refund(ctx context.Context, req *bizpayment.PaymentRefundRequest) (*bizpayment.PaymentResult, error) {
if req == nil {
return nil, errors.New("支付退款请求为空")
}
a, c, err := r.adapter(ctx, req.Provider)
if err != nil {
return nil, err
}
return a.Refund(ctx, req, c)
}
func (r *paymentRepo) HandleCallback(ctx context.Context, callback *bizpayment.PaymentCallback) (*bizpayment.PaymentResult, error) {
if callback == nil {
return nil, errors.New("支付回调为空")
}
a, c, err := r.adapter(ctx, callback.Provider)
if err != nil {
return nil, err
}
result, err := a.Callback(ctx, callback, c)
if err != nil {
return nil, &bizpayment.PaymentCallbackError{Cause: err, Ack: paymentCallbackAck(callback.Provider, c, false)}
}
if result == nil {
err = errors.New("支付回调解析结果为空")
return nil, &bizpayment.PaymentCallbackError{Cause: err, Ack: paymentCallbackAck(callback.Provider, c, false)}
}
result.SuccessAck = paymentCallbackAck(callback.Provider, c, true)
result.FailureAck = paymentCallbackAck(callback.Provider, c, false)
if result.Provider != callback.Provider {
err = errors.New("支付回调渠道不匹配")
return nil, &bizpayment.PaymentCallbackError{Cause: err, Ack: result.FailureAck}
}
result.EventID = paymentCallbackEventID(callback, result)
return result, nil
}
func paymentCallbackEventID(callback *bizpayment.PaymentCallback, result *bizpayment.PaymentResult) string {
if result != nil {
if eventID := strings.TrimSpace(result.EventID); eventID != "" {
return eventID
}
}
if callback == nil {
return ""
}
fields := paymentkit.CallbackFields(callback.Query, callback.Headers, callback.Body)
if eventID := strings.TrimSpace(paymentkit.FirstString(fields, "event_id", "notify_id", "notificationUUID", "id")); eventID != "" {
return eventID
}
hash := sha256.Sum256(append([]byte(callback.Provider+"\x00"), callback.Body...))
return hex.EncodeToString(hash[:])
}
func paymentCallbackAck(provider string, values map[string]any, success bool) bizpayment.PaymentCallbackAck {
ack := bizpayment.DefaultPaymentCallbackAck(provider, success)
prefix := "callback_success_"
if !success {
prefix = "callback_failure_"
}
if configured := strings.TrimSpace(paymentkit.Text(values, prefix+"status")); configured != "" {
if status, err := strconv.Atoi(configured); err == nil && status >= 200 && status <= 599 {
ack.StatusCode = status
}
}
if contentType := strings.TrimSpace(paymentkit.Text(values, prefix+"content_type")); contentType != "" {
ack.ContentType = contentType
}
if body, exists := values[prefix+"body"]; exists {
ack.Body = []byte(fmt.Sprint(body))
}
return ack
}