package nezha_test import ( "context" "crypto" "crypto/rand" "crypto/rsa" "crypto/sha256" "crypto/x509" "encoding/base64" "encoding/json" "encoding/pem" "net/http" "net/http/httptest" "net/url" "sort" "strconv" "strings" "testing" "time" "github.com/wangjia/pay/internal/provider" nz "github.com/wangjia/pay/internal/provider/nezha" ) const testNotifyURL = "https://pay.example.com/api/v2/callback/nezha" // --------------------------------------------------------------------------- // 独立(不复用 nezha 包内部实现)签名/验签工具,照 sign_note 字面重新实现一遍。 // 与 alipay_test.go 的 signRSA2 同一套惯例:测试自己造签名/验签逻辑去构造 fixture、 // 校验被测代码产出,避免"拿被测代码自己的实现验自己"的假绿。exclusion/ordering 的 // 精确规则另有 TestBuildSignSourceOrderingAndExclusion 通过 export_test.go 直接钉死。 // --------------------------------------------------------------------------- func canonicalSource(params map[string]string) string { keys := make([]string, 0, len(params)) for k, v := range params { if k == "sign" || k == "sign_type" || v == "" { continue } keys = append(keys, k) } sort.Strings(keys) parts := make([]string, 0, len(keys)) for _, k := range keys { parts = append(parts, k+"="+params[k]) } return strings.Join(parts, "&") } func signWith(t *testing.T, priv *rsa.PrivateKey, params map[string]string) string { t.Helper() h := sha256.Sum256([]byte(canonicalSource(params))) sig, err := rsa.SignPKCS1v15(rand.Reader, priv, crypto.SHA256, h[:]) if err != nil { t.Fatalf("sign: %v", err) } return base64.StdEncoding.EncodeToString(sig) } func verifyWith(pub *rsa.PublicKey, params map[string]string, sigB64 string) error { sig, err := base64.StdEncoding.DecodeString(sigB64) if err != nil { return err } h := sha256.Sum256([]byte(canonicalSource(params))) return rsa.VerifyPKCS1v15(pub, crypto.SHA256, h[:], sig) } func genKey(t *testing.T) *rsa.PrivateKey { t.Helper() k, err := rsa.GenerateKey(rand.Reader, 2048) if err != nil { t.Fatalf("gen rsa key: %v", err) } return k } func privPEM(k *rsa.PrivateKey) string { return string(pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(k)})) } func pubPEM(k *rsa.PublicKey) string { der, _ := x509.MarshalPKIXPublicKey(k) return string(pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: der})) } // formToParams 把 url.Values(单值)拍平成 map[string]string,供签名串重算用。 func formToParams(v url.Values) map[string]string { out := make(map[string]string, len(v)) for k := range v { out[k] = v.Get(k) } return out } // jsonToParams 把 JSON 响应节点(map[string]any,数字统一是 float64)拍平成 // map[string]string,与 nezha 包内 stringifyParams 的转换规则保持一致 // (仅标量参与签名;整数不带小数点,以便与商户/平台侧对同一份数值的字符串化结果一致)。 func jsonToParams(m map[string]any) map[string]string { out := make(map[string]string, len(m)) for k, v := range m { switch t := v.(type) { case string: out[k] = t case float64: out[k] = strconv.FormatFloat(t, 'f', -1, 64) } } return out } func newProvider(t *testing.T, baseURL, notifyURL string, merchantPriv, platformPriv *rsa.PrivateKey) *nz.Provider { t.Helper() p, err := nz.New("test-pid-1", privPEM(merchantPriv), pubPEM(&platformPriv.PublicKey), notifyURL, nz.WithBaseURL(baseURL)) if err != nil { t.Fatalf("new: %v", err) } return p } // --------------------------------------------------------------------------- // 签名串构造规则:精确钉死排序/排除(sign/sign_type/空值)。 // --------------------------------------------------------------------------- func TestBuildSignSourceOrderingAndExclusion(t *testing.T) { got := nz.BuildSignSourceForTest(map[string]string{ "b": "2", "a": "1", "sign": "should-be-excluded", "sign_type": "RSA", "empty": "", "c": "3", }) want := "a=1&b=2&c=3" if got != want { t.Fatalf("got %q want %q", got, want) } } // --------------------------------------------------------------------------- // New: 凭证校验。 // --------------------------------------------------------------------------- func TestNewRejectsBadCredentials(t *testing.T) { merchant := genKey(t) if _, err := nz.New("", privPEM(merchant), pubPEM(&merchant.PublicKey), testNotifyURL); err == nil { t.Fatal("空 pid 应报错") } if _, err := nz.New("pid", privPEM(merchant), "not-a-key", testNotifyURL); err == nil { t.Fatal("非法平台公钥应报错") } if _, err := nz.New("pid", "not-a-key", pubPEM(&merchant.PublicKey), testNotifyURL); err == nil { t.Fatal("非法商户私钥应报错") } if _, err := nz.New("pid", privPEM(merchant), pubPEM(&merchant.PublicKey), ""); err == nil { t.Fatal("空 notify_url 应报错") } } // TestNewAcceptsBareBase64Keys 兼容无 PEM 头的裸 base64 DER(同 alipay 账户凭证 // "无 PEM 头亦可" 的既有惯例)。 func TestNewAcceptsBareBase64Keys(t *testing.T) { merchant := genKey(t) platform := genKey(t) privDER := base64.StdEncoding.EncodeToString(x509.MarshalPKCS1PrivateKey(merchant)) pubDER, err := x509.MarshalPKIXPublicKey(&platform.PublicKey) if err != nil { t.Fatalf("marshal platform pub: %v", err) } pubB64 := base64.StdEncoding.EncodeToString(pubDER) if _, err := nz.New("pid", privDER, pubB64, testNotifyURL); err != nil { t.Fatalf("裸 base64 密钥应可解析: %v", err) } } // --------------------------------------------------------------------------- // Create // --------------------------------------------------------------------------- func TestCreateRedirect(t *testing.T) { merchant := genKey(t) platform := genKey(t) var gotForm url.Values ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if err := r.ParseForm(); err != nil { t.Fatalf("parse form: %v", err) } gotForm = r.Form if got := r.Header.Get("Content-Type"); got != "application/x-www-form-urlencoded" { t.Fatalf("content-type = %q", got) } resp := map[string]any{ "code": float64(0), "msg": "success", "trade_no": "NZ2026071100001", "out_trade_no": gotForm.Get("out_trade_no"), "pay_type": gotForm.Get("type"), "payurl": "https://nzzf.org/pay/checkout/abc", "qrcode": "", "timestamp": strconv.FormatInt(time.Now().Unix(), 10), "sign_type": "RSA", } resp["sign"] = signWith(t, platform, jsonToParams(resp)) w.Header().Set("Content-Type", "application/json") _ = json.NewEncoder(w).Encode(resp) })) defer ts.Close() p := newProvider(t, ts.URL, testNotifyURL, merchant, platform) sess, err := p.Create(context.Background(), provider.CreateRequest{ OutTradeNo: "PAY-NZ-1", Subject: "Pro 年付", AmountMinor: 19900, Currency: "CNY", ReturnURL: "https://x/return", }) if err != nil { t.Fatalf("create: %v", err) } if sess.RenderType != provider.RenderRedirect || sess.ProviderRef != "NZ2026071100001" { t.Fatalf("session = %+v", sess) } if u, _ := sess.Payload["url"].(string); u != "https://nzzf.org/pay/checkout/abc" { t.Fatalf("payload.url = %v", sess.Payload["url"]) } // 下单请求本身的签名须逐字节对:用"平台侧"验签工具、商户公钥重验一遍。 if err := verifyWith(&merchant.PublicKey, formToParams(gotForm), gotForm.Get("sign")); err != nil { t.Fatalf("请求签名验证失败(逐字节不对): %v", err) } if gotForm.Get("sign_type") != "RSA" { t.Fatalf("sign_type = %q, want RSA", gotForm.Get("sign_type")) } if gotForm.Get("pid") != "test-pid-1" { t.Fatalf("pid = %q", gotForm.Get("pid")) } if gotForm.Get("type") != "alipay" { t.Fatalf("未传 Metadata[type] 时默认应为 alipay, got %q", gotForm.Get("type")) } if gotForm.Get("notify_url") != testNotifyURL { t.Fatalf("notify_url = %q, want %q", gotForm.Get("notify_url"), testNotifyURL) } if gotForm.Get("return_url") != "https://x/return" { t.Fatalf("return_url = %q", gotForm.Get("return_url")) } if gotForm.Get("money") != "199" { t.Fatalf("money = %q, want 199(money.Format 去尾零)", gotForm.Get("money")) } } func TestCreateQRWithExplicitType(t *testing.T) { merchant := genKey(t) platform := genKey(t) var gotForm url.Values ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { _ = r.ParseForm() gotForm = r.Form resp := map[string]any{ "code": float64(0), "msg": "success", "trade_no": "NZ-QR-1", "out_trade_no": gotForm.Get("out_trade_no"), "pay_type": gotForm.Get("type"), "payurl": "", "qrcode": "weixin://wxpay/bizpayurl?pr=abc123", "timestamp": strconv.FormatInt(time.Now().Unix(), 10), "sign_type": "RSA", } resp["sign"] = signWith(t, platform, jsonToParams(resp)) w.Header().Set("Content-Type", "application/json") _ = json.NewEncoder(w).Encode(resp) })) defer ts.Close() p := newProvider(t, ts.URL, testNotifyURL, merchant, platform) sess, err := p.Create(context.Background(), provider.CreateRequest{ OutTradeNo: "PAY-NZ-QR", Subject: "Pro 月付", AmountMinor: 2990, Currency: "CNY", Metadata: map[string]string{"render": "qr", "type": "wxpay"}, }) if err != nil { t.Fatalf("create qr: %v", err) } if sess.RenderType != provider.RenderQR || sess.ProviderRef != "NZ-QR-1" { t.Fatalf("session = %+v", sess) } if sess.Payload["qr_content"] != "weixin://wxpay/bizpayurl?pr=abc123" { t.Fatalf("qr_content = %v", sess.Payload["qr_content"]) } if sess.Payload["currency"] != "CNY" { t.Fatalf("currency = %v", sess.Payload["currency"]) } if sess.Payload["display_amount"] != "29.9" { t.Fatalf("display_amount = %v, want 29.9", sess.Payload["display_amount"]) } if gotForm.Get("type") != "wxpay" { t.Fatalf("Metadata[type]=wxpay 应透传, got %q", gotForm.Get("type")) } } // TestCreateQRAutoSelectedWhenNoPayURL: 未显式要求 render=qr,但渠道响应只带 qrcode // 不带 payurl 时,也应自动落 RenderQR(不能因缺 url 报错)。 func TestCreateQRAutoSelectedWhenNoPayURL(t *testing.T) { merchant := genKey(t) platform := genKey(t) ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { _ = r.ParseForm() resp := map[string]any{ "code": float64(0), "msg": "success", "trade_no": "NZ-AUTO-QR", "payurl": "", "qrcode": "https://qr.example/abc", "timestamp": strconv.FormatInt(time.Now().Unix(), 10), "sign_type": "RSA", } resp["sign"] = signWith(t, platform, jsonToParams(resp)) w.Header().Set("Content-Type", "application/json") _ = json.NewEncoder(w).Encode(resp) })) defer ts.Close() p := newProvider(t, ts.URL, testNotifyURL, merchant, platform) sess, err := p.Create(context.Background(), provider.CreateRequest{ OutTradeNo: "PAY-AUTO-QR", Subject: "x", AmountMinor: 100, Currency: "CNY", }) if err != nil { t.Fatalf("create: %v", err) } if sess.RenderType != provider.RenderQR { t.Fatalf("render_type = %v, want qr(响应只带 qrcode)", sess.RenderType) } } func TestCreateChannelRejected(t *testing.T) { merchant := genKey(t) platform := genKey(t) ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { resp := map[string]any{ "code": float64(-1), "msg": "商户不存在", "timestamp": strconv.FormatInt(time.Now().Unix(), 10), "sign_type": "RSA", } resp["sign"] = signWith(t, platform, jsonToParams(resp)) w.Header().Set("Content-Type", "application/json") _ = json.NewEncoder(w).Encode(resp) })) defer ts.Close() p := newProvider(t, ts.URL, testNotifyURL, merchant, platform) sess, err := p.Create(context.Background(), provider.CreateRequest{ OutTradeNo: "PAY-BAD", Subject: "x", AmountMinor: 100, Currency: "CNY", }) if err == nil { t.Fatalf("want 渠道拒绝返回 error, got session=%+v", sess) } if sess != nil { t.Fatalf("渠道拒绝时不应返回 session, got %+v", sess) } } // TestCreateResponseBadSignRejected: 攻击者/中间人用不相干私钥冒充平台响应签名, // Create 必须报错、绝不能把伪造的 payurl/trade_no 当真返回。 func TestCreateResponseBadSignRejected(t *testing.T) { merchant := genKey(t) platform := genKey(t) wrongPlatform := genKey(t) // 冒充平台侧的另一把毫不相干的私钥 ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { resp := map[string]any{ "code": float64(0), "msg": "success", "trade_no": "NZ-EVIL", "payurl": "https://evil.example/pay", "timestamp": strconv.FormatInt(time.Now().Unix(), 10), "sign_type": "RSA", } resp["sign"] = signWith(t, wrongPlatform, jsonToParams(resp)) w.Header().Set("Content-Type", "application/json") _ = json.NewEncoder(w).Encode(resp) })) defer ts.Close() p := newProvider(t, ts.URL, testNotifyURL, merchant, platform) sess, err := p.Create(context.Background(), provider.CreateRequest{ OutTradeNo: "PAY-EVIL", Subject: "x", AmountMinor: 100, Currency: "CNY", }) if err == nil { t.Fatalf("want 验签失败 error, got session=%+v", sess) } if sess != nil { t.Fatalf("验签失败时不应返回 session(哪怕字段看起来正常), got %+v", sess) } } func TestCreateRejectsNonCNY(t *testing.T) { merchant := genKey(t) platform := genKey(t) p := newProvider(t, "https://unused.invalid", testNotifyURL, merchant, platform) if _, err := p.Create(context.Background(), provider.CreateRequest{ OutTradeNo: "PAY-USD", Subject: "x", AmountMinor: 100, Currency: "USD", }); err == nil { t.Fatal("want 非 CNY 报错") } } // --------------------------------------------------------------------------- // VerifyCallback(GET query,平台公钥验签) // --------------------------------------------------------------------------- func buildCallbackQuery(t *testing.T, priv *rsa.PrivateKey, overrides map[string]string) map[string]string { q := map[string]string{ "pid": "test-pid-1", "trade_no": "NZ-CB-1", "out_trade_no": "PAY-CB-1", "type": "alipay", "name": "Pro 年付", "money": "199.00", "trade_status": "TRADE_SUCCESS", "timestamp": strconv.FormatInt(time.Now().Unix(), 10), "sign_type": "RSA", } for k, v := range overrides { q[k] = v } q["sign"] = signWith(t, priv, q) return q } func TestVerifyCallbackSuccess(t *testing.T) { merchant := genKey(t) platform := genKey(t) p := newProvider(t, "https://unused.invalid", testNotifyURL, merchant, platform) q := buildCallbackQuery(t, platform, nil) ev, err := p.VerifyCallback(context.Background(), provider.CallbackInput{Query: q}) if err != nil { t.Fatalf("verify: %v", err) } if ev.ProviderRef != "NZ-CB-1" || ev.Status != provider.PaidSucceeded { t.Fatalf("event = %+v", ev) } if ev.PaidAmountMinor != 19900 || ev.PaidCurrency != "CNY" { t.Fatalf("amount = %d %s", ev.PaidAmountMinor, ev.PaidCurrency) } if ev.PaidAt != nil { t.Fatalf("回调未带独立付款时间字段,PaidAt 应为 nil(交 settle 用收到时间兜底), got %v", ev.PaidAt) } } // TestVerifyCallbackTamperedAmountRejected: 签名后篡改金额,验签必须失败——逐字节对 // 是本任务的硬要求,篡改任何一个参与签名的字段都不能蒙混过关。 func TestVerifyCallbackTamperedAmountRejected(t *testing.T) { merchant := genKey(t) platform := genKey(t) p := newProvider(t, "https://unused.invalid", testNotifyURL, merchant, platform) q := buildCallbackQuery(t, platform, nil) q["money"] = "999999.00" // 签名之后篡改 ev, err := p.VerifyCallback(context.Background(), provider.CallbackInput{Query: q}) if err == nil { t.Fatalf("want 验签失败, got event=%+v", ev) } if ev != nil { t.Fatalf("验签失败不应返回 event, got %+v", ev) } } func TestVerifyCallbackWrongKeyRejected(t *testing.T) { merchant := genKey(t) platform := genKey(t) wrongPlatform := genKey(t) p := newProvider(t, "https://unused.invalid", testNotifyURL, merchant, platform) q := buildCallbackQuery(t, wrongPlatform, nil) // 攻击者用不相干私钥签名,冒充平台 ev, err := p.VerifyCallback(context.Background(), provider.CallbackInput{Query: q}) if err == nil { t.Fatalf("want 验签失败(签名私钥与装配平台公钥不匹配), got event=%+v", ev) } if ev != nil { t.Fatalf("验签失败不应返回 event, got %+v", ev) } } func TestVerifyCallbackNonSuccessStatusIsPending(t *testing.T) { merchant := genKey(t) platform := genKey(t) p := newProvider(t, "https://unused.invalid", testNotifyURL, merchant, platform) q := buildCallbackQuery(t, platform, map[string]string{"trade_status": "TRADE_CLOSED"}) ev, err := p.VerifyCallback(context.Background(), provider.CallbackInput{Query: q}) if err != nil { t.Fatalf("verify: %v", err) } if ev.Status != provider.PaidPending { t.Fatalf("status = %v, want pending", ev.Status) } } func TestVerifyCallbackMissingSignRejected(t *testing.T) { merchant := genKey(t) platform := genKey(t) p := newProvider(t, "https://unused.invalid", testNotifyURL, merchant, platform) q := map[string]string{"trade_no": "NZ-CB-2", "trade_status": "TRADE_SUCCESS"} if _, err := p.VerifyCallback(context.Background(), provider.CallbackInput{Query: q}); err == nil { t.Fatal("want 缺 sign 报错") } } // --------------------------------------------------------------------------- // Query // --------------------------------------------------------------------------- func TestQuerySucceeded(t *testing.T) { merchant := genKey(t) platform := genKey(t) var gotForm url.Values ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { _ = r.ParseForm() gotForm = r.Form resp := map[string]any{ "code": float64(0), "trade_no": "NZ-Q-1", "out_trade_no": "PAY-Q-1", "status": float64(1), "money": "199.00", "type": "alipay", "addtime": "2026-07-11 10:00:00", "endtime": "2026-07-11 10:00:05", } w.Header().Set("Content-Type", "application/json") _ = json.NewEncoder(w).Encode(resp) })) defer ts.Close() p := newProvider(t, ts.URL, testNotifyURL, merchant, platform) ev, err := p.Query(context.Background(), provider.QueryRequest{ProviderRef: "NZ-Q-1", OutTradeNo: "PAY-Q-1"}) if err != nil { t.Fatalf("query: %v", err) } if ev.Status != provider.PaidSucceeded || ev.PaidAmountMinor != 19900 || ev.PaidCurrency != "CNY" { t.Fatalf("event = %+v", ev) } if gotForm.Get("trade_no") != "NZ-Q-1" { t.Fatalf("查单请求应优先带 trade_no(ProviderRef), got form=%v", gotForm) } if err := verifyWith(&merchant.PublicKey, formToParams(gotForm), gotForm.Get("sign")); err != nil { t.Fatalf("查单请求签名验证失败: %v", err) } } func TestQueryPending(t *testing.T) { merchant := genKey(t) platform := genKey(t) ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { resp := map[string]any{"code": float64(0), "trade_no": "NZ-Q-2", "status": float64(0)} w.Header().Set("Content-Type", "application/json") _ = json.NewEncoder(w).Encode(resp) })) defer ts.Close() p := newProvider(t, ts.URL, testNotifyURL, merchant, platform) ev, err := p.Query(context.Background(), provider.QueryRequest{ProviderRef: "NZ-Q-2"}) if err != nil { t.Fatalf("query: %v", err) } if ev.Status != provider.PaidPending { t.Fatalf("status = %v, want pending", ev.Status) } } func TestQueryPrefersOutTradeNoWhenNoProviderRef(t *testing.T) { merchant := genKey(t) platform := genKey(t) var gotForm url.Values ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { _ = r.ParseForm() gotForm = r.Form resp := map[string]any{"code": float64(0), "status": float64(0)} w.Header().Set("Content-Type", "application/json") _ = json.NewEncoder(w).Encode(resp) })) defer ts.Close() p := newProvider(t, ts.URL, testNotifyURL, merchant, platform) if _, err := p.Query(context.Background(), provider.QueryRequest{OutTradeNo: "PAY-NOREF"}); err != nil { t.Fatalf("query: %v", err) } if gotForm.Get("out_trade_no") != "PAY-NOREF" || gotForm.Get("trade_no") != "" { t.Fatalf("无 ProviderRef 时应用 out_trade_no 而非 trade_no, form=%v", gotForm) } } // --------------------------------------------------------------------------- // Method / Capabilities // --------------------------------------------------------------------------- func TestMethodAndCapabilities(t *testing.T) { merchant := genKey(t) platform := genKey(t) p := newProvider(t, "https://unused.invalid", testNotifyURL, merchant, platform) if p.Method() != "nezha" { t.Fatalf("Method() = %q, want nezha", p.Method()) } caps := p.Capabilities() if caps.SupportsRefund { t.Fatal("SupportsRefund 应为 false(官方文档未见退款接口)") } if len(caps.SettleCurrencies) != 1 || caps.SettleCurrencies[0] != "CNY" { t.Fatalf("SettleCurrencies = %v", caps.SettleCurrencies) } if len(caps.Regions) != 1 || caps.Regions[0] != "cn" { t.Fatalf("Regions = %v", caps.Regions) } wantRender := map[provider.RenderType]bool{provider.RenderRedirect: true, provider.RenderQR: true} if len(caps.RenderTypes) != len(wantRender) { t.Fatalf("RenderTypes = %v", caps.RenderTypes) } for _, rt := range caps.RenderTypes { if !wantRender[rt] { t.Fatalf("意外的 render_type: %v", rt) } } }