279 lines
8.9 KiB
Go
279 lines
8.9 KiB
Go
package payment
|
|
|
|
import (
|
|
"context"
|
|
"crypto/sha256"
|
|
"encoding/hex"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"net/url"
|
|
"strconv"
|
|
"strings"
|
|
|
|
"kra/internal/biz"
|
|
datapayment "kra/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)
|
|
}
|