package payment import ( "context" "crypto/tls" "encoding/json" "fmt" "io" "net" "net/http" "net/http/httptest" "testing" "kra/internal/biz" "kra/internal/utils/paymentutil" "github.com/go-pay/gopay" gopayQQ "github.com/go-pay/gopay/qq" ) func redirectQQHTTPS(t *testing.T, handler http.Handler) { t.Helper() server := httptest.NewTLSServer(handler) oldDefault := http.DefaultTransport base, ok := oldDefault.(*http.Transport) if !ok { server.Close() t.Fatal("http.DefaultTransport is not *http.Transport") } transport := base.Clone() transport.Proxy = nil transport.DisableKeepAlives = true transport.ForceAttemptHTTP2 = false transport.TLSClientConfig = &tls.Config{MinVersion: tls.VersionTLS12, InsecureSkipVerify: true} //nolint:gosec // local test server target := server.Listener.Addr().String() transport.DialContext = func(ctx context.Context, network, _ string) (net.Conn, error) { return (&net.Dialer{}).DialContext(ctx, network, target) } http.DefaultTransport = transport t.Cleanup(func() { http.DefaultTransport = oldDefault transport.CloseIdleConnections() server.Close() }) } func TestQQCreateKeepsClientPaymentDataOutOfDurableIdentities(t *testing.T) { redirectQQHTTPS(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost || r.URL.Path != "/cgi-bin/pay/qpay_unified_order.cgi" { t.Errorf("request = %s %s", r.Method, r.URL.Path) } body, err := io.ReadAll(r.Body) if err != nil { t.Errorf("read request XML: %v", err) http.Error(w, "invalid XML", http.StatusBadRequest) return } values, err := paymentutil.XMLValues(body) if err != nil { t.Errorf("parse request XML: %v", err) http.Error(w, "invalid XML", http.StatusBadRequest) return } if values["mch_id"] != "QQ-MERCHANT" || values["out_trade_no"] != "MERCHANT-QQ-1" { t.Errorf("request identities = %#v", values) } if values["trade_type"] != "NATIVE" || values["total_fee"] != "100" || values["sign"] == "" { t.Errorf("request fields = %#v", values) } w.Header().Set("Content-Type", "application/xml") response := gopay.BodyMap{ "return_code": "SUCCESS", "result_code": "SUCCESS", "mch_id": "QQ-MERCHANT", "trade_type": "NATIVE", "prepay_id": "QQ-PREPAY-1", "code_url": "https://qpay.qq.com/qr/QQ-PREPAY-1", } response.Set("sign", gopayQQ.GetReleaseSign("qq-api-key", gopayQQ.SignType_MD5, response)) encoded := make(map[string]string, len(response)) for key, value := range response { encoded[key] = fmt.Sprint(value) } _, _ = w.Write(paymentutil.XMLEncode(encoded)) })) result, err := (&qqAdapter{}).Create(context.Background(), &biz.PaymentRequest{ TradeNo: "MERCHANT-QQ-1", Subject: "subject", Amount: 100, Currency: "CNY", ClientIP: "127.0.0.1", NotifyURL: "https://merchant.example/qq/notify", }, map[string]any{"mch_id": "QQ-MERCHANT", "api_key": "qq-api-key"}) if err != nil { t.Fatal(err) } if result.Provider != biz.PaymentQQ || result.Status != "created" || result.TradeNo != "MERCHANT-QQ-1" { t.Fatalf("create result = %+v", result) } if result.QueryID != "" || result.ProviderTradeNo != "" { t.Fatalf("client payment data leaked into durable identities: %+v", result) } var payload struct { PrepayID string `json:"prepay_id"` CodeURL string `json:"code_url"` } if err = json.Unmarshal(result.Payload, &payload); err != nil { t.Fatal(err) } if payload.PrepayID != "QQ-PREPAY-1" || payload.CodeURL != "https://qpay.qq.com/qr/QQ-PREPAY-1" { t.Fatalf("client payment payload = %+v", payload) } } func TestQQCreateMethodMapsSupportedModes(t *testing.T) { for _, tc := range []struct { value string want string }{ {value: "", want: gopayQQ.TradeType_Native}, {value: "MICROPAY", want: gopayQQ.TradeType_MicroPay}, {value: "barcode", want: gopayQQ.TradeType_MicroPay}, {value: "JSAPI", want: gopayQQ.TradeType_JsApi}, {value: "APP", want: gopayQQ.TradeType_App}, {value: "MINIAPP", want: gopayQQ.TradeType_Mini}, } { if got, err := qqCreateMethod(map[string]any{"trade_type": tc.value}, nil); err != nil || got != tc.want { t.Fatalf("mode %q = %q, err = %v, want %q", tc.value, got, err, tc.want) } } } func TestQQCreateMethodRejectsUnknownMode(t *testing.T) { if got, err := qqCreateMethod(map[string]any{"method": "unsupported"}, nil); err == nil || got != "" { t.Fatalf("unknown mode = %q, err = %v", got, err) } } func TestNormalizeQQStateTreatsRefundAsFailed(t *testing.T) { if got := normalizeQQState("REFUND"); got != "failed" { t.Fatalf("normalizeQQState(REFUND) = %q, want failed", got) } }