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

279 lines
9.0 KiB
Go

package payment
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"net/url"
"strconv"
"strings"
"kra/app/system/internal/biz"
datapayment "kra/app/system/internal/integration/payment"
"gorm.io/gorm"
)
type paymentRepo struct{ data Provider }
func NewPaymentRepo(data Provider) biz.PaymentRepo { return &paymentRepo{data: data} }
func ensurePaymentIntegrationConfigs(db *gorm.DB) error {
for _, provider := range biz.SupportedPaymentProviders {
var row integrationConfigPO
err := db.Where("kind = ? AND provider = ?", integrationKindPayment, provider).First(&row).Error
defaults := biz.DefaultIntegrationConfig(integrationKindPayment, provider)
if errors.Is(err, gorm.ErrRecordNotFound) {
encoded, _ := json.Marshal(defaults)
if err := db.Create(&integrationConfigPO{Kind: integrationKindPayment, Provider: provider, Enabled: false, Config: string(encoded)}).Error; err != nil {
return err
}
continue
}
if err != nil {
return err
}
values := map[string]any{}
_ = json.Unmarshal([]byte(row.Config), &values)
changed := false
for key, value := range defaults {
if _, exists := values[key]; !exists {
values[key] = value
changed = true
}
}
if changed {
encoded, _ := json.Marshal(values)
if err := db.Model(&row).Update("config", string(encoded)).Error; err != nil {
return err
}
}
}
return nil
}
func (r *paymentRepo) row(ctx context.Context, provider string) (*integrationConfigPO, map[string]any, error) {
var row integrationConfigPO
if err := r.data.DB().WithContext(ctx).Where("kind = ? AND provider = ?", integrationKindPayment, provider).First(&row).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil, biz.ErrPaymentProviderNotFound
}
return nil, nil, err
}
if !row.Enabled {
return nil, nil, fmt.Errorf("支付渠道 %s 未启用", provider)
}
values := map[string]any{}
if err := json.Unmarshal([]byte(row.Config), &values); err != nil {
return nil, nil, fmt.Errorf("支付配置格式错误: %w", err)
}
return &row, values, nil
}
func (r *paymentRepo) adapter(ctx context.Context, provider string) (datapayment.Adapter, map[string]any, error) {
_, values, err := r.row(ctx, provider)
if err != nil {
return nil, nil, err
}
adapter, err := datapayment.New(provider)
return adapter, values, err
}
func (r *paymentRepo) Create(ctx context.Context, req *biz.PaymentRequest) (*biz.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(text(c, "notify_url"))
effective.ReturnURL = strings.TrimSpace(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 biz.PaymentApple, biz.PaymentAllinPay, biz.PaymentSaobei, biz.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 biz.PaymentAlipay, biz.PaymentAlipayV3:
keys = []string{"method", "pay_method", "trade_type", "channel"}
case biz.PaymentWechatV2:
keys = []string{"trade_type", "pay_type", "method", "pay_method", "channel"}
case biz.PaymentWechatV3:
keys = []string{"trade_type", "pay_type", "method"}
case biz.PaymentQQ:
keys = []string{"trade_type", "pay_type", "method", "pay_method"}
case biz.PaymentLakala:
keys = []string{"method", "pay_method", "trade_type"}
}
value := firstAny(extra, keys...)
if value == "" {
value = firstAny(config, keys...)
}
normalized := strings.NewReplacer(".", "_", "-", "_", " ", "_").Replace(strings.ToLower(strings.TrimSpace(value)))
switch provider {
case biz.PaymentAlipay, biz.PaymentAlipayV3:
return !contains([]string{"pay", "trade_pay", "alipay_trade_pay", "barcode", "barcode_pay", "micropay", "face_to_face"}, normalized)
case biz.PaymentWechatV2:
return !contains([]string{"micropay", "micro_pay", "barcode", "barcode_pay", "pay_code", "payment_code"}, normalized)
case biz.PaymentWechatV3:
return !contains([]string{"micropay", "micro_pay", "codepay", "code_pay", "barcode", "barcode_pay", "facepay", "face_pay"}, normalized)
case biz.PaymentQQ:
return !contains([]string{"micropay", "micro_pay", "barcode", "barcode_pay"}, normalized)
case biz.PaymentLakala:
return !contains([]string{"retail", "retail_pay", "micropay", "barcode"}, normalized)
default:
return true
}
}
func (r *paymentRepo) Query(ctx context.Context, provider, tradeNo string) (*biz.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 *biz.PaymentRefundRequest) (*biz.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 *biz.PaymentCallback) (*biz.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, &biz.PaymentCallbackError{Cause: err, Ack: paymentCallbackAck(callback.Provider, c, false)}
}
if result == nil {
err = errors.New("支付回调解析结果为空")
return nil, &biz.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, &biz.PaymentCallbackError{Cause: err, Ack: result.FailureAck}
}
result.EventID = paymentCallbackEventID(callback, result)
return result, nil
}
func paymentCallbackEventID(callback *biz.PaymentCallback, result *biz.PaymentResult) string {
if result != nil {
if eventID := strings.TrimSpace(result.EventID); eventID != "" {
return eventID
}
}
if callback == nil {
return ""
}
fields := callbackFields(callback)
if eventID := strings.TrimSpace(first(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) biz.PaymentCallbackAck {
ack := biz.DefaultPaymentCallbackAck(provider, success)
prefix := "callback_success_"
if !success {
prefix = "callback_failure_"
}
if configured := strings.TrimSpace(text(values, prefix+"status")); configured != "" {
if status, err := strconv.Atoi(configured); err == nil && status >= 200 && status <= 599 {
ack.StatusCode = status
}
}
if contentType := strings.TrimSpace(text(values, prefix+"content_type")); contentType != "" {
ack.ContentType = contentType
}
if body, exists := values[prefix+"body"]; exists {
ack.Body = []byte(fmt.Sprint(body))
}
return ack
}
func callbackFields(callback *biz.PaymentCallback) map[string]string {
fields := map[string]string{}
for key, value := range callback.Query {
fields[key] = value
}
contentType := strings.ToLower(callback.Headers["Content-Type"])
if strings.Contains(contentType, "application/x-www-form-urlencoded") {
if values, err := url.ParseQuery(string(callback.Body)); err == nil {
for key, value := range values {
if len(value) > 0 {
fields[key] = value[0]
}
}
}
}
var object map[string]any
if json.Unmarshal(callback.Body, &object) == nil {
for key, value := range object {
if text, ok := value.(string); ok {
fields[key] = text
}
}
}
return fields
}
func first(values map[string]string, keys ...string) string {
for _, key := range keys {
if values[key] != "" {
return values[key]
}
}
return ""
}
func contains(values []string, value string) bool {
for _, item := range values {
if item == value {
return true
}
}
return false
}
func validatePaymentConfig(provider string, values map[string]any) error {
return biz.ValidateIntegrationConfig(biz.IntegrationKindPayment, provider, values)
}