package payment import ( "context" "crypto" "crypto/rand" "crypto/rsa" "crypto/sha1" "crypto/tls" "crypto/x509" "encoding/base64" "encoding/json" "io" "net" "net/http" "net/http/httptest" "net/url" "strings" "testing" "kra/internal/biz" "github.com/go-pay/gopay" "github.com/go-pay/gopay/allinpay" ) func allinPayTestKeys(t *testing.T) (*rsa.PrivateKey, string, string) { t.Helper() key, err := rsa.GenerateKey(rand.Reader, 2048) if err != nil { t.Fatal(err) } privateKey := base64.StdEncoding.EncodeToString(x509.MarshalPKCS1PrivateKey(key)) publicDER, err := x509.MarshalPKIXPublicKey(&key.PublicKey) if err != nil { t.Fatal(err) } publicKey := base64.StdEncoding.EncodeToString(publicDER) return key, privateKey, publicKey } func allinPayTestConfig(privateKey, publicKey string, orderType string) map[string]any { config := map[string]any{ "cus_id": "CUS-TEST", "app_id": "APP-TEST", "private_key": privateKey, "public_key": publicKey, "environment": "sandbox", } if orderType != "" { config["query_order_type"] = orderType } return config } func redirectAllinPayHTTPS(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 signedAllinPayResponse(t *testing.T, key *rsa.PrivateKey, fields gopay.BodyMap) []byte { t.Helper() signData := fields.EncodeAliPaySignParams() digest := sha1.Sum([]byte(signData)) signature, err := rsa.SignPKCS1v15(rand.Reader, key, crypto.SHA1, digest[:]) if err != nil { t.Fatal(err) } fields.Set("sign", base64.StdEncoding.EncodeToString(signature)) body, err := json.Marshal(fields) if err != nil { t.Fatal(err) } return body } func TestAllinPayTrxIDCreateRejectsResponseWithoutTransactionID(t *testing.T) { key, privateKey, publicKey := allinPayTestKeys(t) redirectAllinPayHTTPS(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost || r.URL.Path != "/apiweb/unitorder/pay" { t.Errorf("request = %s %s", r.Method, r.URL.Path) } body, err := io.ReadAll(r.Body) if err != nil { t.Errorf("read request: %v", err) return } values, err := url.ParseQuery(string(body)) if err != nil { t.Errorf("parse request form: %v", err) } if values.Get("reqsn") != "MERCHANT-TRXID-MISSING" || values.Get("paytype") == "" { t.Errorf("request fields = %v", values) } response := signedAllinPayResponse(t, key, gopay.BodyMap{ "retcode": "SUCCESS", "retmsg": "ok", "reqsn": "MERCHANT-TRXID-MISSING", "trxstatus": "SUCCESS", "payinfo": "client-only-payinfo", }) w.Header().Set("Content-Type", "application/json") _, _ = w.Write(response) })) _, err := (&allinpayAdapter{}).Create(context.Background(), &biz.PaymentRequest{ Provider: biz.PaymentAllinPay, TradeNo: "MERCHANT-TRXID-MISSING", Subject: "subject", Amount: 100, Currency: "CNY", }, allinPayTestConfig(privateKey, publicKey, allinpay.OrderTypeTrxId)) if err == nil || !strings.Contains(strings.ToLower(err.Error()), "trxid") { t.Fatalf("missing trxid error = %v", err) } } func TestAllinPayNativeCreateRejectsTrxIDLookupMode(t *testing.T) { _, privateKey, publicKey := allinPayTestKeys(t) _, err := (&allinpayAdapter{}).Create(context.Background(), &biz.PaymentRequest{ Provider: biz.PaymentAllinPay, TradeNo: "MERCHANT-NATIVE-TRXID", Subject: "subject", Amount: 100, Currency: "CNY", Extra: map[string]any{"method": "native", "expiretime": "20261231235959"}, }, allinPayTestConfig(privateKey, publicKey, allinpay.OrderTypeTrxId)) if err == nil || !strings.Contains(strings.ToLower(err.Error()), "native") || !strings.Contains(strings.ToLower(err.Error()), "trxid") { t.Fatalf("native trxid error = %v", err) } } func TestAllinPayCreateRejectsUnknownMethodBeforeClient(t *testing.T) { _, err := (&allinpayAdapter{}).Create(context.Background(), &biz.PaymentRequest{ Provider: biz.PaymentAllinPay, TradeNo: "MERCHANT-UNKNOWN-METHOD", Subject: "subject", Amount: 100, Currency: "CNY", Extra: map[string]any{"method": "unsupported"}, }, nil) if err == nil || !strings.Contains(err.Error(), "不支持的下单方式") { t.Fatalf("unknown method error = %v", err) } } func TestNormalizeAllinState(t *testing.T) { for _, tc := range []struct { state string want string }{ {state: "SUCCESS", want: "success"}, {state: "0000", want: "success"}, {state: "2000", want: "pending"}, {state: "2008", want: "pending"}, {state: "3000", want: "failed"}, {state: "3040", want: "failed"}, {state: "3045", want: "failed"}, {state: "3999", want: "failed"}, } { t.Run(tc.state, func(t *testing.T) { if got := normalizeAllinState(tc.state); got != tc.want { t.Fatalf("normalizeAllinState(%q) = %q, want %q", tc.state, got, tc.want) } }) } } func TestAllinPayTrxIDQueryRejectsMismatchedResponseTransactionID(t *testing.T) { key, privateKey, publicKey := allinPayTestKeys(t) redirectAllinPayHTTPS(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if err := r.ParseForm(); err != nil { t.Errorf("parse form: %v", err) } if r.Form.Get("trxid") != "TRX-REQUESTED" { t.Errorf("query form = %v", r.Form) } response := signedAllinPayResponse(t, key, gopay.BodyMap{ "retcode": "SUCCESS", "retmsg": "ok", "reqsn": "MERCHANT-1", "trxid": "TRX-OTHER", "trxstatus": "0000", "trxamt": "100", }) w.Header().Set("Content-Type", "application/json") _, _ = w.Write(response) })) _, err := (&allinpayAdapter{}).Query(context.Background(), "TRX-REQUESTED", allinPayTestConfig(privateKey, publicKey, allinpay.OrderTypeTrxId)) if err == nil || !strings.Contains(strings.ToLower(err.Error()), "trxid") || !strings.Contains(err.Error(), "不匹配") { t.Fatalf("mismatched trxid error = %v", err) } } func TestAllinPayRefundUsesConfiguredLookupIdentity(t *testing.T) { for _, tc := range []struct { name string orderType string tradeNo string queryID string wantOldTrxID string wantOldReqSN string }{ { name: "trxid uses durable query id", orderType: allinpay.OrderTypeTrxId, tradeNo: "MERCHANT-REFUND-TRXID", queryID: "TRX-ORIGINAL-1", wantOldTrxID: "TRX-ORIGINAL-1", }, { name: "reqsn uses merchant trade number", orderType: allinpay.OrderTypeReqSN, tradeNo: "MERCHANT-REFUND-REQSN", queryID: "SHOULD-NOT-BE-USED", wantOldReqSN: "MERCHANT-REFUND-REQSN", }, } { t.Run(tc.name, func(t *testing.T) { key, privateKey, publicKey := allinPayTestKeys(t) redirectAllinPayHTTPS(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost || r.URL.Path != "/apiweb/tranx/refund" { t.Errorf("request = %s %s", r.Method, r.URL.Path) } body, err := io.ReadAll(r.Body) if err != nil { t.Errorf("read request: %v", err) return } values, err := url.ParseQuery(string(body)) if err != nil { t.Errorf("parse request form: %v", err) } if values.Get("oldtrxid") != tc.wantOldTrxID || values.Get("oldreqsn") != tc.wantOldReqSN { t.Errorf("refund identity fields = %v", values) } if values.Get("reqsn") != "REFUND-1" || values.Get("trxamt") != "40" { t.Errorf("refund fields = %v", values) } response := signedAllinPayResponse(t, key, gopay.BodyMap{ "retcode": "SUCCESS", "retmsg": "ok", "reqsn": "REFUND-1", "trxid": "REFUND-TRX-1", "trxstatus": "SUCCESS", "fee": "40", }) w.Header().Set("Content-Type", "application/json") _, _ = w.Write(response) })) result, err := (&allinpayAdapter{}).Refund(context.Background(), &biz.PaymentRefundRequest{ Provider: biz.PaymentAllinPay, TradeNo: tc.tradeNo, QueryID: tc.queryID, RefundNo: "REFUND-1", Amount: 40, TotalAmount: 100, Currency: "CNY", }, allinPayTestConfig(privateKey, publicKey, tc.orderType)) if err != nil { t.Fatal(err) } if result.ProviderTradeNo != "REFUND-TRX-1" || result.Amount != 40 { t.Fatalf("refund result = %+v", result) } }) } } func TestAllinPayRefundResultBindsRefundIdentity(t *testing.T) { req := &biz.PaymentRefundRequest{TradeNo: "LOCAL-1", RefundNo: "REFUND-1", Amount: 40, Currency: "CNY"} base := func() *allinpay.RefundRsp { return &allinpay.RefundRsp{Reqsn: "REFUND-1", Trxid: "ALLIN-REFUND-1", TrxStatus: "SUCCESS", Fee: "40"} } for _, tc := range []struct { name string mutate func(*allinpay.RefundRsp) }{ {name: "missing merchant refund number", mutate: func(rsp *allinpay.RefundRsp) { rsp.Reqsn = "" }}, {name: "mismatched merchant refund number", mutate: func(rsp *allinpay.RefundRsp) { rsp.Reqsn = "OTHER" }}, {name: "missing provider refund id", mutate: func(rsp *allinpay.RefundRsp) { rsp.Trxid = "" }}, {name: "missing amount", mutate: func(rsp *allinpay.RefundRsp) { rsp.Fee = "" }}, {name: "invalid amount", mutate: func(rsp *allinpay.RefundRsp) { rsp.Fee = "invalid" }}, {name: "mismatched amount", mutate: func(rsp *allinpay.RefundRsp) { rsp.Fee = "41" }}, } { t.Run(tc.name, func(t *testing.T) { rsp := base() tc.mutate(rsp) if _, err := allinpayRefundResult(req, rsp, "CNY"); err == nil { t.Fatal("allinpayRefundResult() accepted an unbound refund response") } }) } }