Files
pay/internal/provider/provider_test.go
T

50 lines
1.8 KiB
Go

package provider_test
import (
"context"
"errors"
"testing"
"github.com/wangjia/pay/internal/provider"
)
// stubProvider 最小实现,驱动 Provider 接口 + Registry 成型。
type stubProvider struct{ method string }
func (s stubProvider) Method() string { return s.method }
func (s stubProvider) Capabilities() provider.Capabilities { return provider.Capabilities{RenderTypes: []provider.RenderType{provider.RenderQR}} }
func (s stubProvider) Create(context.Context, provider.CreateRequest) (*provider.Session, error) {
return &provider.Session{ProviderRef: "R-1", RenderType: provider.RenderQR}, nil
}
func (s stubProvider) VerifyCallback(context.Context, provider.CallbackInput) (*provider.PaidEvent, error) {
return &provider.PaidEvent{ProviderRef: "R-1", Status: provider.PaidSucceeded}, nil
}
func (s stubProvider) Query(context.Context, provider.QueryRequest) (*provider.PaidEvent, error) {
return &provider.PaidEvent{ProviderRef: "R-1", Status: provider.PaidPending}, nil
}
func TestRegistryRegisterGet(t *testing.T) {
r := provider.NewRegistry()
r.Register(stubProvider{method: "alipay"})
r.Register(stubProvider{method: "crypto"})
p, err := r.Get("crypto")
if err != nil || p.Method() != "crypto" {
t.Fatalf("Get crypto = %v, %v", p, err)
}
if _, err := r.Get("nope"); !errors.Is(err, provider.ErrUnknownMethod) {
t.Fatalf("未知 method 应返回 ErrUnknownMethod, got %v", err)
}
if got := r.Methods(); len(got) != 2 || got[0] != "alipay" || got[1] != "crypto" {
t.Fatalf("Methods 应按字典序返回 [alipay crypto], got %v", got)
}
}
func TestSessionAndCaps(t *testing.T) {
var _ provider.Provider = stubProvider{} // 编译期断言 stub 满足接口
caps := stubProvider{}.Capabilities()
if len(caps.RenderTypes) != 1 || caps.RenderTypes[0] != provider.RenderQR {
t.Fatalf("caps = %+v", caps)
}
}