292 lines
9.8 KiB
Go
292 lines
9.8 KiB
Go
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")
|
|
}
|
|
})
|
|
}
|
|
}
|