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) }