mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
refactor(payment): code quality improvements per project conventions
- payment_service: split GetDashboardStats (131→35 lines) into 4 sub-functions, extract applyPagination helper, replace bubble sort with sort.Slice, use sync.Once for EnsureProviders, define magic number constants - providers: merge duplicate PC/Mobile functions in alipay/wxpay, extract loadKeyPair/queryOrderTotalFen in wxpay, add centsToYuan in stripe, define constants for success codes, currencies, event types - handlers: extract requireAuth and parseIDParam helpers to eliminate repeated auth/ID-parsing boilerplate - registry: remove dead seen map in GetProviderByKey - types: centralize order status constants in payment package - load_balancer+config_service: unify duplicated containsType into exported InstanceSupportsType function
This commit is contained in:
@@ -71,9 +71,8 @@ func (h *PaymentHandler) ListOrders(c *gin.Context) {
|
||||
// GetOrderDetail returns detailed information about a single order.
|
||||
// GET /api/v1/admin/payment/orders/:id
|
||||
func (h *PaymentHandler) GetOrderDetail(c *gin.Context) {
|
||||
orderID, err := strconv.ParseInt(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.BadRequest(c, "Invalid order ID")
|
||||
orderID, ok := parseIDParam(c, "id")
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
order, err := h.paymentService.GetOrderByID(c.Request.Context(), orderID)
|
||||
@@ -88,9 +87,8 @@ func (h *PaymentHandler) GetOrderDetail(c *gin.Context) {
|
||||
// CancelOrder cancels a pending order (admin).
|
||||
// POST /api/v1/admin/payment/orders/:id/cancel
|
||||
func (h *PaymentHandler) CancelOrder(c *gin.Context) {
|
||||
orderID, err := strconv.ParseInt(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.BadRequest(c, "Invalid order ID")
|
||||
orderID, ok := parseIDParam(c, "id")
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
msg, err := h.paymentService.AdminCancelOrder(c.Request.Context(), orderID)
|
||||
@@ -104,9 +102,8 @@ func (h *PaymentHandler) CancelOrder(c *gin.Context) {
|
||||
// RetryFulfillment retries fulfillment for a paid order.
|
||||
// POST /api/v1/admin/payment/orders/:id/retry
|
||||
func (h *PaymentHandler) RetryFulfillment(c *gin.Context) {
|
||||
orderID, err := strconv.ParseInt(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.BadRequest(c, "Invalid order ID")
|
||||
orderID, ok := parseIDParam(c, "id")
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := h.paymentService.RetryFulfillment(c.Request.Context(), orderID); err != nil {
|
||||
@@ -127,9 +124,8 @@ type AdminProcessRefundRequest struct {
|
||||
// ProcessRefund processes a refund for an order (admin).
|
||||
// POST /api/v1/admin/payment/orders/:id/refund
|
||||
func (h *PaymentHandler) ProcessRefund(c *gin.Context) {
|
||||
orderID, err := strconv.ParseInt(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.BadRequest(c, "Invalid order ID")
|
||||
orderID, ok := parseIDParam(c, "id")
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -189,9 +185,8 @@ func (h *PaymentHandler) CreatePlan(c *gin.Context) {
|
||||
// UpdatePlan updates an existing subscription plan.
|
||||
// PUT /api/v1/admin/payment/plans/:id
|
||||
func (h *PaymentHandler) UpdatePlan(c *gin.Context) {
|
||||
id, err := strconv.ParseInt(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.BadRequest(c, "Invalid plan ID")
|
||||
id, ok := parseIDParam(c, "id")
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var req service.UpdatePlanRequest
|
||||
@@ -210,9 +205,8 @@ func (h *PaymentHandler) UpdatePlan(c *gin.Context) {
|
||||
// DeletePlan deletes a subscription plan.
|
||||
// DELETE /api/v1/admin/payment/plans/:id
|
||||
func (h *PaymentHandler) DeletePlan(c *gin.Context) {
|
||||
id, err := strconv.ParseInt(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.BadRequest(c, "Invalid plan ID")
|
||||
id, ok := parseIDParam(c, "id")
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := h.configService.DeletePlan(c.Request.Context(), id); err != nil {
|
||||
@@ -254,9 +248,8 @@ func (h *PaymentHandler) CreateProvider(c *gin.Context) {
|
||||
// UpdateProvider updates an existing payment provider instance.
|
||||
// PUT /api/v1/admin/payment/providers/:id
|
||||
func (h *PaymentHandler) UpdateProvider(c *gin.Context) {
|
||||
id, err := strconv.ParseInt(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.BadRequest(c, "Invalid provider ID")
|
||||
id, ok := parseIDParam(c, "id")
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var req service.UpdateProviderInstanceRequest
|
||||
@@ -275,9 +268,8 @@ func (h *PaymentHandler) UpdateProvider(c *gin.Context) {
|
||||
// DeleteProvider deletes a payment provider instance.
|
||||
// DELETE /api/v1/admin/payment/providers/:id
|
||||
func (h *PaymentHandler) DeleteProvider(c *gin.Context) {
|
||||
id, err := strconv.ParseInt(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.BadRequest(c, "Invalid provider ID")
|
||||
id, ok := parseIDParam(c, "id")
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := h.configService.DeleteProviderInstance(c.Request.Context(), id); err != nil {
|
||||
@@ -287,6 +279,17 @@ func (h *PaymentHandler) DeleteProvider(c *gin.Context) {
|
||||
response.Success(c, gin.H{"message": "deleted"})
|
||||
}
|
||||
|
||||
// parseIDParam parses an int64 path parameter.
|
||||
// Returns the parsed ID and true on success; on failure it writes a BadRequest response and returns false.
|
||||
func parseIDParam(c *gin.Context, paramName string) (int64, bool) {
|
||||
id, err := strconv.ParseInt(c.Param(paramName), 10, 64)
|
||||
if err != nil {
|
||||
response.BadRequest(c, "Invalid "+paramName)
|
||||
return 0, false
|
||||
}
|
||||
return id, true
|
||||
}
|
||||
|
||||
// --- Config ---
|
||||
|
||||
// GetConfig returns the payment configuration (admin view).
|
||||
|
||||
@@ -4,9 +4,9 @@ import (
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/pagination"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/response"
|
||||
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/pagination"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -88,9 +88,8 @@ type CreateOrderRequest struct {
|
||||
// CreateOrder creates a new payment order.
|
||||
// POST /api/v1/payment/orders
|
||||
func (h *PaymentHandler) CreateOrder(c *gin.Context) {
|
||||
subject, ok := middleware2.GetAuthSubjectFromContext(c)
|
||||
subject, ok := requireAuth(c)
|
||||
if !ok {
|
||||
response.Unauthorized(c, "User not authenticated")
|
||||
return
|
||||
}
|
||||
|
||||
@@ -121,9 +120,8 @@ func (h *PaymentHandler) CreateOrder(c *gin.Context) {
|
||||
// GetMyOrders returns the authenticated user's orders.
|
||||
// GET /api/v1/payment/orders/my
|
||||
func (h *PaymentHandler) GetMyOrders(c *gin.Context) {
|
||||
subject, ok := middleware2.GetAuthSubjectFromContext(c)
|
||||
subject, ok := requireAuth(c)
|
||||
if !ok {
|
||||
response.Unauthorized(c, "User not authenticated")
|
||||
return
|
||||
}
|
||||
|
||||
@@ -145,9 +143,8 @@ func (h *PaymentHandler) GetMyOrders(c *gin.Context) {
|
||||
// GetOrder returns a single order for the authenticated user.
|
||||
// GET /api/v1/payment/orders/:id
|
||||
func (h *PaymentHandler) GetOrder(c *gin.Context) {
|
||||
subject, ok := middleware2.GetAuthSubjectFromContext(c)
|
||||
subject, ok := requireAuth(c)
|
||||
if !ok {
|
||||
response.Unauthorized(c, "User not authenticated")
|
||||
return
|
||||
}
|
||||
|
||||
@@ -168,9 +165,8 @@ func (h *PaymentHandler) GetOrder(c *gin.Context) {
|
||||
// CancelOrder cancels a pending order for the authenticated user.
|
||||
// POST /api/v1/payment/orders/:id/cancel
|
||||
func (h *PaymentHandler) CancelOrder(c *gin.Context) {
|
||||
subject, ok := middleware2.GetAuthSubjectFromContext(c)
|
||||
subject, ok := requireAuth(c)
|
||||
if !ok {
|
||||
response.Unauthorized(c, "User not authenticated")
|
||||
return
|
||||
}
|
||||
|
||||
@@ -196,9 +192,8 @@ type RefundRequestBody struct {
|
||||
// RequestRefund submits a refund request for a completed order.
|
||||
// POST /api/v1/payment/orders/:id/refund-request
|
||||
func (h *PaymentHandler) RequestRefund(c *gin.Context) {
|
||||
subject, ok := middleware2.GetAuthSubjectFromContext(c)
|
||||
subject, ok := requireAuth(c)
|
||||
if !ok {
|
||||
response.Unauthorized(c, "User not authenticated")
|
||||
return
|
||||
}
|
||||
|
||||
@@ -221,6 +216,17 @@ func (h *PaymentHandler) RequestRefund(c *gin.Context) {
|
||||
response.Success(c, gin.H{"message": "refund requested"})
|
||||
}
|
||||
|
||||
// requireAuth extracts the authenticated subject from the context.
|
||||
// Returns the subject and true on success; on failure it writes an Unauthorized response and returns false.
|
||||
func requireAuth(c *gin.Context) (middleware2.AuthSubject, bool) {
|
||||
subject, ok := middleware2.GetAuthSubjectFromContext(c)
|
||||
if !ok {
|
||||
response.Unauthorized(c, "User not authenticated")
|
||||
return middleware2.AuthSubject{}, false
|
||||
}
|
||||
return subject, true
|
||||
}
|
||||
|
||||
// isMobile detects mobile user agents.
|
||||
func isMobile(c *gin.Context) bool {
|
||||
ua := strings.ToLower(c.GetHeader("User-Agent"))
|
||||
|
||||
@@ -17,6 +17,9 @@ type PaymentWebhookHandler struct {
|
||||
registry *payment.Registry
|
||||
}
|
||||
|
||||
// maxWebhookBodySize is the maximum allowed webhook request body size (1 MB).
|
||||
const maxWebhookBodySize = 1 << 20
|
||||
|
||||
// NewPaymentWebhookHandler creates a new PaymentWebhookHandler.
|
||||
func NewPaymentWebhookHandler(paymentService *service.PaymentService, registry *payment.Registry) *PaymentWebhookHandler {
|
||||
return &PaymentWebhookHandler{
|
||||
@@ -51,7 +54,7 @@ func (h *PaymentWebhookHandler) StripeWebhook(c *gin.Context) {
|
||||
|
||||
// handleNotify is the shared logic for all provider webhook handlers.
|
||||
func (h *PaymentWebhookHandler) handleNotify(c *gin.Context, providerKey string) {
|
||||
body, err := io.ReadAll(io.LimitReader(c.Request.Body, 1<<20)) // 1MB limit
|
||||
body, err := io.ReadAll(io.LimitReader(c.Request.Body, maxWebhookBodySize))
|
||||
if err != nil {
|
||||
slog.Error("[Payment Webhook] failed to read body", "provider", providerKey, "error", err)
|
||||
c.String(http.StatusBadRequest, "failed to read body")
|
||||
|
||||
@@ -62,7 +62,7 @@ func (lb *DefaultLoadBalancer) SelectInstance(ctx context.Context, providerKey s
|
||||
// Filter by supported types
|
||||
var candidates []*dbent.PaymentProviderInstance
|
||||
for _, inst := range instances {
|
||||
if inst.SupportedTypes == "" || containsType(inst.SupportedTypes, paymentType) {
|
||||
if InstanceSupportsType(inst.SupportedTypes, paymentType) {
|
||||
candidates = append(candidates, inst)
|
||||
}
|
||||
}
|
||||
@@ -110,7 +110,7 @@ func (lb *DefaultLoadBalancer) GetInstanceDailyAmount(ctx context.Context, insta
|
||||
err := lb.db.PaymentOrder.Query().
|
||||
Where(
|
||||
paymentorder.ProviderInstanceID(instanceID),
|
||||
paymentorder.StatusIn("COMPLETED", "PAID", "RECHARGING"),
|
||||
paymentorder.StatusIn(OrderStatusCompleted, OrderStatusPaid, OrderStatusRecharging),
|
||||
paymentorder.PaidAtGTE(todayStart),
|
||||
).
|
||||
Aggregate(dbent.Sum(paymentorder.FieldPayAmount)).
|
||||
@@ -124,8 +124,12 @@ func (lb *DefaultLoadBalancer) GetInstanceDailyAmount(ctx context.Context, insta
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
// containsType checks if a comma-separated list contains the given type.
|
||||
func containsType(supportedTypes string, target PaymentType) bool {
|
||||
// InstanceSupportsType checks if the given supported types string includes the target type.
|
||||
// An empty supportedTypes string means all types are supported.
|
||||
func InstanceSupportsType(supportedTypes string, target PaymentType) bool {
|
||||
if supportedTypes == "" {
|
||||
return true
|
||||
}
|
||||
for _, t := range strings.Split(supportedTypes, ",") {
|
||||
if strings.TrimSpace(t) == target {
|
||||
return true
|
||||
|
||||
@@ -4,7 +4,7 @@ import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestContainsType(t *testing.T) {
|
||||
func TestInstanceSupportsType(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
@@ -68,19 +68,19 @@ func TestContainsType(t *testing.T) {
|
||||
expected: false,
|
||||
},
|
||||
{
|
||||
name: "empty supported types string matches nothing",
|
||||
name: "empty supported types means all supported",
|
||||
supportedTypes: "",
|
||||
target: "alipay",
|
||||
expected: false,
|
||||
expected: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
got := containsType(tt.supportedTypes, tt.target)
|
||||
got := InstanceSupportsType(tt.supportedTypes, tt.target)
|
||||
if got != tt.expected {
|
||||
t.Fatalf("containsType(%q, %q) = %v, want %v", tt.supportedTypes, tt.target, got, tt.expected)
|
||||
t.Fatalf("InstanceSupportsType(%q, %q) = %v, want %v", tt.supportedTypes, tt.target, got, tt.expected)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -13,6 +13,12 @@ import (
|
||||
"github.com/smartwalle/alipay/v3"
|
||||
)
|
||||
|
||||
// Alipay product codes.
|
||||
const (
|
||||
alipayProductCodePagePay = "FAST_INSTANT_TRADE_PAY"
|
||||
alipayProductCodeWapPay = "QUICK_WAP_WAY"
|
||||
)
|
||||
|
||||
// Alipay implements payment.Provider and payment.CancelableProvider using the smartwalle/alipay SDK.
|
||||
type Alipay struct {
|
||||
instanceID string
|
||||
@@ -59,7 +65,7 @@ func (a *Alipay) getClient() (*alipay.Client, error) {
|
||||
return a.client, nil
|
||||
}
|
||||
|
||||
func (a *Alipay) Name() string { return "Alipay" }
|
||||
func (a *Alipay) Name() string { return "Alipay" }
|
||||
func (a *Alipay) ProviderKey() string { return "alipay" }
|
||||
func (a *Alipay) SupportedTypes() []payment.PaymentType {
|
||||
return []payment.PaymentType{payment.TypeAlipayDirect}
|
||||
@@ -82,17 +88,36 @@ func (a *Alipay) CreatePayment(_ context.Context, req payment.CreatePaymentReque
|
||||
}
|
||||
|
||||
if req.IsMobile {
|
||||
return a.createWapPayment(client, req, notifyURL, returnURL)
|
||||
return a.createTrade(client, req, notifyURL, returnURL, true)
|
||||
}
|
||||
return a.createPagePayment(client, req, notifyURL, returnURL)
|
||||
return a.createTrade(client, req, notifyURL, returnURL, false)
|
||||
}
|
||||
|
||||
func (a *Alipay) createPagePayment(client *alipay.Client, req payment.CreatePaymentRequest, notifyURL, returnURL string) (*payment.CreatePaymentResponse, error) {
|
||||
func (a *Alipay) createTrade(client *alipay.Client, req payment.CreatePaymentRequest, notifyURL, returnURL string, isMobile bool) (*payment.CreatePaymentResponse, error) {
|
||||
if isMobile {
|
||||
param := alipay.TradeWapPay{}
|
||||
param.OutTradeNo = req.OrderID
|
||||
param.TotalAmount = req.Amount
|
||||
param.Subject = req.Subject
|
||||
param.ProductCode = alipayProductCodeWapPay
|
||||
param.NotifyURL = notifyURL
|
||||
param.ReturnURL = returnURL
|
||||
|
||||
payURL, err := client.TradeWapPay(param)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("alipay TradeWapPay: %w", err)
|
||||
}
|
||||
return &payment.CreatePaymentResponse{
|
||||
TradeNo: req.OrderID,
|
||||
PayURL: payURL.String(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
param := alipay.TradePagePay{}
|
||||
param.OutTradeNo = req.OrderID
|
||||
param.TotalAmount = req.Amount
|
||||
param.Subject = req.Subject
|
||||
param.ProductCode = "FAST_INSTANT_TRADE_PAY"
|
||||
param.ProductCode = alipayProductCodePagePay
|
||||
param.NotifyURL = notifyURL
|
||||
param.ReturnURL = returnURL
|
||||
|
||||
@@ -100,7 +125,6 @@ func (a *Alipay) createPagePayment(client *alipay.Client, req payment.CreatePaym
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("alipay TradePagePay: %w", err)
|
||||
}
|
||||
|
||||
return &payment.CreatePaymentResponse{
|
||||
TradeNo: req.OrderID,
|
||||
PayURL: payURL.String(),
|
||||
@@ -108,26 +132,6 @@ func (a *Alipay) createPagePayment(client *alipay.Client, req payment.CreatePaym
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (a *Alipay) createWapPayment(client *alipay.Client, req payment.CreatePaymentRequest, notifyURL, returnURL string) (*payment.CreatePaymentResponse, error) {
|
||||
param := alipay.TradeWapPay{}
|
||||
param.OutTradeNo = req.OrderID
|
||||
param.TotalAmount = req.Amount
|
||||
param.Subject = req.Subject
|
||||
param.ProductCode = "QUICK_WAP_WAY"
|
||||
param.NotifyURL = notifyURL
|
||||
param.ReturnURL = returnURL
|
||||
|
||||
payURL, err := client.TradeWapPay(param)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("alipay TradeWapPay: %w", err)
|
||||
}
|
||||
|
||||
return &payment.CreatePaymentResponse{
|
||||
TradeNo: req.OrderID,
|
||||
PayURL: payURL.String(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// QueryOrder queries the trade status via Alipay.
|
||||
func (a *Alipay) QueryOrder(ctx context.Context, tradeNo string) (*payment.QueryOrderResponse, error) {
|
||||
client, err := a.getClient()
|
||||
@@ -154,7 +158,10 @@ func (a *Alipay) QueryOrder(ctx context.Context, tradeNo string) (*payment.Query
|
||||
status = "failed"
|
||||
}
|
||||
|
||||
amount, _ := strconv.ParseFloat(result.TotalAmount, 64)
|
||||
amount, err := strconv.ParseFloat(result.TotalAmount, 64)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("alipay parse amount %q: %w", result.TotalAmount, err)
|
||||
}
|
||||
|
||||
return &payment.QueryOrderResponse{
|
||||
TradeNo: result.TradeNo,
|
||||
@@ -186,7 +193,10 @@ func (a *Alipay) VerifyNotification(ctx context.Context, rawBody string, _ map[s
|
||||
status = "success"
|
||||
}
|
||||
|
||||
amount, _ := strconv.ParseFloat(notification.TotalAmount, 64)
|
||||
amount, err := strconv.ParseFloat(notification.TotalAmount, 64)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("alipay parse notification amount %q: %w", notification.TotalAmount, err)
|
||||
}
|
||||
|
||||
return &payment.PaymentNotification{
|
||||
TradeNo: notification.TradeNo,
|
||||
|
||||
@@ -18,6 +18,13 @@ import (
|
||||
"github.com/Wei-Shaw/sub2api/internal/payment"
|
||||
)
|
||||
|
||||
// EasyPay constants.
|
||||
const (
|
||||
easypayCodeSuccess = 1
|
||||
easypayStatusPaid = 1
|
||||
easypayHTTPTimeout = 10 * time.Second
|
||||
)
|
||||
|
||||
// EasyPay implements payment.Provider for the EasyPay aggregation platform.
|
||||
type EasyPay struct {
|
||||
instanceID string
|
||||
@@ -36,12 +43,12 @@ func NewEasyPay(instanceID string, config map[string]string) (*EasyPay, error) {
|
||||
return &EasyPay{
|
||||
instanceID: instanceID,
|
||||
config: config,
|
||||
httpClient: &http.Client{Timeout: 10 * time.Second},
|
||||
httpClient: &http.Client{Timeout: easypayHTTPTimeout},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (e *EasyPay) Name() string { return "EasyPay" }
|
||||
func (e *EasyPay) ProviderKey() string { return "easypay" }
|
||||
func (e *EasyPay) ProviderKey() string { return "easypay" }
|
||||
func (e *EasyPay) SupportedTypes() []payment.PaymentType {
|
||||
return []payment.PaymentType{payment.TypeAlipay, payment.TypeWxpay}
|
||||
}
|
||||
@@ -76,7 +83,7 @@ func (e *EasyPay) CreatePayment(ctx context.Context, req payment.CreatePaymentRe
|
||||
if err := json.Unmarshal(body, &resp); err != nil {
|
||||
return nil, fmt.Errorf("easypay parse: %w", err)
|
||||
}
|
||||
if resp.Code != 1 {
|
||||
if resp.Code != easypayCodeSuccess {
|
||||
return nil, fmt.Errorf("easypay error: %s", resp.Msg)
|
||||
}
|
||||
return &payment.CreatePaymentResponse{TradeNo: resp.TradeNo, PayURL: resp.PayURL, QRCode: resp.QRCode}, nil
|
||||
@@ -100,7 +107,7 @@ func (e *EasyPay) QueryOrder(ctx context.Context, tradeNo string) (*payment.Quer
|
||||
return nil, fmt.Errorf("easypay parse query: %w", err)
|
||||
}
|
||||
status := "pending"
|
||||
if resp.Status == 1 {
|
||||
if resp.Status == easypayStatusPaid {
|
||||
status = "paid"
|
||||
}
|
||||
return &payment.QueryOrderResponse{TradeNo: tradeNo, Status: status}, nil
|
||||
@@ -148,7 +155,7 @@ func (e *EasyPay) Refund(ctx context.Context, req payment.RefundRequest) (*payme
|
||||
if err := json.Unmarshal(body, &resp); err != nil {
|
||||
return nil, fmt.Errorf("easypay parse refund: %w", err)
|
||||
}
|
||||
if resp.Code != 1 {
|
||||
if resp.Code != easypayCodeSuccess {
|
||||
return nil, fmt.Errorf("easypay refund failed: %s", resp.Msg)
|
||||
}
|
||||
return &payment.RefundResponse{RefundID: req.TradeNo, Status: "success"}, nil
|
||||
|
||||
@@ -14,6 +14,14 @@ import (
|
||||
"github.com/stripe/stripe-go/v82/webhook"
|
||||
)
|
||||
|
||||
// Stripe constants.
|
||||
const (
|
||||
stripeCurrency = "cny"
|
||||
stripeEventPaymentSuccess = "payment_intent.succeeded"
|
||||
stripeEventPaymentFailed = "payment_intent.payment_failed"
|
||||
stripeCentsPerYuan = 100
|
||||
)
|
||||
|
||||
// Stripe implements the payment.CancelableProvider interface for Stripe payments.
|
||||
type Stripe struct {
|
||||
instanceID string
|
||||
@@ -50,12 +58,17 @@ func (s *Stripe) GetPublishableKey() string {
|
||||
return s.config["publishableKey"]
|
||||
}
|
||||
|
||||
func (s *Stripe) Name() string { return "Stripe" }
|
||||
func (s *Stripe) Name() string { return "Stripe" }
|
||||
func (s *Stripe) ProviderKey() string { return "stripe" }
|
||||
func (s *Stripe) SupportedTypes() []payment.PaymentType {
|
||||
return []payment.PaymentType{payment.TypeStripe}
|
||||
}
|
||||
|
||||
// centsToYuan converts an amount in cents (int64) to yuan (float64).
|
||||
func centsToYuan(cents int64) float64 {
|
||||
return float64(cents) / stripeCentsPerYuan
|
||||
}
|
||||
|
||||
// yuanToCents converts a CNY yuan string to cents (int64).
|
||||
func yuanToCents(amountStr string) (int64, error) {
|
||||
amount, err := strconv.ParseFloat(amountStr, 64)
|
||||
@@ -76,7 +89,7 @@ func (s *Stripe) CreatePayment(ctx context.Context, req payment.CreatePaymentReq
|
||||
|
||||
params := &stripe.PaymentIntentParams{
|
||||
Amount: stripe.Int64(amountInCents),
|
||||
Currency: stripe.String("cny"),
|
||||
Currency: stripe.String(stripeCurrency),
|
||||
AutomaticPaymentMethods: &stripe.PaymentIntentAutomaticPaymentMethodsParams{
|
||||
Enabled: stripe.Bool(true),
|
||||
},
|
||||
@@ -120,7 +133,7 @@ func (s *Stripe) QueryOrder(ctx context.Context, tradeNo string) (*payment.Query
|
||||
return &payment.QueryOrderResponse{
|
||||
TradeNo: pi.ID,
|
||||
Status: status,
|
||||
Amount: float64(pi.Amount) / 100,
|
||||
Amount: centsToYuan(pi.Amount),
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -144,9 +157,9 @@ func (s *Stripe) VerifyNotification(_ context.Context, rawBody string, headers m
|
||||
}
|
||||
|
||||
switch event.Type {
|
||||
case "payment_intent.succeeded":
|
||||
case stripeEventPaymentSuccess:
|
||||
return parseStripePaymentIntent(&event, "success", rawBody)
|
||||
case "payment_intent.payment_failed":
|
||||
case stripeEventPaymentFailed:
|
||||
return parseStripePaymentIntent(&event, "failed", rawBody)
|
||||
}
|
||||
|
||||
@@ -161,7 +174,7 @@ func parseStripePaymentIntent(event *stripe.Event, status string, rawBody string
|
||||
return &payment.PaymentNotification{
|
||||
TradeNo: pi.ID,
|
||||
OrderID: pi.Metadata["orderId"],
|
||||
Amount: float64(pi.Amount) / 100,
|
||||
Amount: centsToYuan(pi.Amount),
|
||||
Status: status,
|
||||
RawData: rawBody,
|
||||
}, nil
|
||||
|
||||
@@ -3,6 +3,7 @@ package provider
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/rsa"
|
||||
"fmt"
|
||||
"io"
|
||||
"math"
|
||||
@@ -24,6 +25,13 @@ import (
|
||||
"github.com/wechatpay-apiv3/wechatpay-go/utils"
|
||||
)
|
||||
|
||||
// WeChat Pay constants.
|
||||
const (
|
||||
wxpayCurrency = "CNY"
|
||||
wxpayH5Type = "Wap"
|
||||
wxpayFenPerYuan = 100
|
||||
)
|
||||
|
||||
type Wxpay struct {
|
||||
instanceID string
|
||||
config map[string]string
|
||||
@@ -45,7 +53,7 @@ func NewWxpay(instanceID string, config map[string]string) (*Wxpay, error) {
|
||||
return &Wxpay{instanceID: instanceID, config: config}, nil
|
||||
}
|
||||
|
||||
func (w *Wxpay) Name() string { return "Wxpay" }
|
||||
func (w *Wxpay) Name() string { return "Wxpay" }
|
||||
func (w *Wxpay) ProviderKey() string { return "wxpay" }
|
||||
func (w *Wxpay) SupportedTypes() []payment.PaymentType {
|
||||
return []payment.PaymentType{payment.TypeWxpayDirect}
|
||||
@@ -65,13 +73,9 @@ func (w *Wxpay) ensureClient() (*core.Client, error) {
|
||||
if w.coreClient != nil {
|
||||
return w.coreClient, nil
|
||||
}
|
||||
privateKey, err := utils.LoadPrivateKey(formatPEM(w.config["privateKey"], "PRIVATE KEY"))
|
||||
privateKey, publicKey, err := w.loadKeyPair()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("wxpay load private key: %w", err)
|
||||
}
|
||||
publicKey, err := utils.LoadPublicKey(formatPEM(w.config["publicKey"], "PUBLIC KEY"))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("wxpay load public key: %w", err)
|
||||
return nil, err
|
||||
}
|
||||
certSerial := w.config["certSerial"]
|
||||
if certSerial == "" {
|
||||
@@ -93,6 +97,18 @@ func (w *Wxpay) ensureClient() (*core.Client, error) {
|
||||
return w.coreClient, nil
|
||||
}
|
||||
|
||||
func (w *Wxpay) loadKeyPair() (*rsa.PrivateKey, *rsa.PublicKey, error) {
|
||||
privateKey, err := utils.LoadPrivateKey(formatPEM(w.config["privateKey"], "PRIVATE KEY"))
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("wxpay load private key: %w", err)
|
||||
}
|
||||
publicKey, err := utils.LoadPublicKey(formatPEM(w.config["publicKey"], "PUBLIC KEY"))
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("wxpay load public key: %w", err)
|
||||
}
|
||||
return privateKey, publicKey, nil
|
||||
}
|
||||
|
||||
func yuanToFen(s string) (int64, error) {
|
||||
f, err := strconv.ParseFloat(s, 64)
|
||||
if err != nil {
|
||||
@@ -118,7 +134,7 @@ func (w *Wxpay) CreatePayment(ctx context.Context, req payment.CreatePaymentRequ
|
||||
return nil, fmt.Errorf("wxpay create payment: %w", err)
|
||||
}
|
||||
if req.IsMobile && req.ClientIP != "" {
|
||||
resp, err := w.createH5Order(ctx, client, req, notifyURL, totalFen)
|
||||
resp, err := w.createOrder(ctx, client, req, notifyURL, totalFen, true)
|
||||
if err == nil {
|
||||
return resp, nil
|
||||
}
|
||||
@@ -126,12 +142,19 @@ func (w *Wxpay) CreatePayment(ctx context.Context, req payment.CreatePaymentRequ
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return w.createNativeOrder(ctx, client, req, notifyURL, totalFen)
|
||||
return w.createOrder(ctx, client, req, notifyURL, totalFen, false)
|
||||
}
|
||||
|
||||
func (w *Wxpay) createNativeOrder(ctx context.Context, c *core.Client, req payment.CreatePaymentRequest, notifyURL string, totalFen int64) (*payment.CreatePaymentResponse, error) {
|
||||
func (w *Wxpay) createOrder(ctx context.Context, c *core.Client, req payment.CreatePaymentRequest, notifyURL string, totalFen int64, useH5 bool) (*payment.CreatePaymentResponse, error) {
|
||||
if useH5 {
|
||||
return w.prepayH5(ctx, c, req, notifyURL, totalFen)
|
||||
}
|
||||
return w.prepayNative(ctx, c, req, notifyURL, totalFen)
|
||||
}
|
||||
|
||||
func (w *Wxpay) prepayNative(ctx context.Context, c *core.Client, req payment.CreatePaymentRequest, notifyURL string, totalFen int64) (*payment.CreatePaymentResponse, error) {
|
||||
svc := native.NativeApiService{Client: c}
|
||||
cur := "CNY"
|
||||
cur := wxpayCurrency
|
||||
resp, _, err := svc.Prepay(ctx, native.PrepayRequest{
|
||||
Appid: core.String(w.config["appId"]), Mchid: core.String(w.config["mchId"]),
|
||||
Description: core.String(req.Subject), OutTradeNo: core.String(req.OrderID),
|
||||
@@ -148,10 +171,10 @@ func (w *Wxpay) createNativeOrder(ctx context.Context, c *core.Client, req payme
|
||||
return &payment.CreatePaymentResponse{TradeNo: req.OrderID, QRCode: codeURL}, nil
|
||||
}
|
||||
|
||||
func (w *Wxpay) createH5Order(ctx context.Context, c *core.Client, req payment.CreatePaymentRequest, notifyURL string, totalFen int64) (*payment.CreatePaymentResponse, error) {
|
||||
func (w *Wxpay) prepayH5(ctx context.Context, c *core.Client, req payment.CreatePaymentRequest, notifyURL string, totalFen int64) (*payment.CreatePaymentResponse, error) {
|
||||
svc := h5.H5ApiService{Client: c}
|
||||
cur := "CNY"
|
||||
tp := "Wap"
|
||||
cur := wxpayCurrency
|
||||
tp := wxpayH5Type
|
||||
resp, _, err := svc.Prepay(ctx, h5.PrepayRequest{
|
||||
Appid: core.String(w.config["appId"]), Mchid: core.String(w.config["mchId"]),
|
||||
Description: core.String(req.Subject), OutTradeNo: core.String(req.OrderID),
|
||||
@@ -203,7 +226,7 @@ func (w *Wxpay) QueryOrder(ctx context.Context, tradeNo string) (*payment.QueryO
|
||||
}
|
||||
var amt float64
|
||||
if tx.Amount != nil && tx.Amount.Total != nil {
|
||||
amt = float64(*tx.Amount.Total) / 100
|
||||
amt = float64(*tx.Amount.Total) / wxpayFenPerYuan
|
||||
}
|
||||
id := tradeNo
|
||||
if tx.TransactionId != nil {
|
||||
@@ -237,7 +260,7 @@ func (w *Wxpay) VerifyNotification(ctx context.Context, rawBody string, headers
|
||||
}
|
||||
var amt float64
|
||||
if tx.Amount != nil && tx.Amount.Total != nil {
|
||||
amt = float64(*tx.Amount.Total) / 100
|
||||
amt = float64(*tx.Amount.Total) / wxpayFenPerYuan
|
||||
}
|
||||
st := "failed"
|
||||
if wxSV(tx.TradeState) == "SUCCESS" {
|
||||
@@ -258,19 +281,12 @@ func (w *Wxpay) Refund(ctx context.Context, req payment.RefundRequest) (*payment
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("wxpay refund amount: %w", err)
|
||||
}
|
||||
svc := native.NativeApiService{Client: c}
|
||||
tx, _, err := svc.QueryOrderByOutTradeNo(ctx, native.QueryOrderByOutTradeNoRequest{
|
||||
OutTradeNo: core.String(req.OrderID), Mchid: core.String(w.config["mchId"]),
|
||||
})
|
||||
tf, err := w.queryOrderTotalFen(ctx, c, req.OrderID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("wxpay refund query order: %w", err)
|
||||
}
|
||||
var tf int64
|
||||
if tx.Amount != nil && tx.Amount.Total != nil {
|
||||
tf = *tx.Amount.Total
|
||||
return nil, err
|
||||
}
|
||||
rs := refunddomestic.RefundsApiService{Client: c}
|
||||
cur := "CNY"
|
||||
cur := wxpayCurrency
|
||||
res, _, err := rs.Create(ctx, refunddomestic.CreateRequest{
|
||||
OutTradeNo: core.String(req.OrderID),
|
||||
OutRefundNo: core.String(fmt.Sprintf("%s-refund-%d", req.OrderID, time.Now().UnixNano())),
|
||||
@@ -291,6 +307,21 @@ func (w *Wxpay) Refund(ctx context.Context, req payment.RefundRequest) (*payment
|
||||
return &payment.RefundResponse{RefundID: rid, Status: st}, nil
|
||||
}
|
||||
|
||||
func (w *Wxpay) queryOrderTotalFen(ctx context.Context, c *core.Client, orderID string) (int64, error) {
|
||||
svc := native.NativeApiService{Client: c}
|
||||
tx, _, err := svc.QueryOrderByOutTradeNo(ctx, native.QueryOrderByOutTradeNoRequest{
|
||||
OutTradeNo: core.String(orderID), Mchid: core.String(w.config["mchId"]),
|
||||
})
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("wxpay refund query order: %w", err)
|
||||
}
|
||||
var tf int64
|
||||
if tx.Amount != nil && tx.Amount.Total != nil {
|
||||
tf = *tx.Amount.Total
|
||||
}
|
||||
return tf, nil
|
||||
}
|
||||
|
||||
func (w *Wxpay) CancelPayment(ctx context.Context, tradeNo string) error {
|
||||
c, err := w.ensureClient()
|
||||
if err != nil {
|
||||
|
||||
@@ -47,13 +47,10 @@ func (r *Registry) GetProvider(t PaymentType) (Provider, error) {
|
||||
func (r *Registry) GetProviderByKey(key string) (Provider, error) {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
seen := make(map[string]bool)
|
||||
for _, p := range r.providers {
|
||||
k := p.ProviderKey()
|
||||
if k == key && !seen[k] {
|
||||
if p.ProviderKey() == key {
|
||||
return p, nil
|
||||
}
|
||||
seen[k] = true
|
||||
}
|
||||
return nil, ErrProviderNotFound
|
||||
}
|
||||
|
||||
@@ -14,9 +14,9 @@ type mockProvider struct {
|
||||
supportedTypes []PaymentType
|
||||
}
|
||||
|
||||
func (m *mockProvider) Name() string { return m.name }
|
||||
func (m *mockProvider) ProviderKey() string { return m.key }
|
||||
func (m *mockProvider) SupportedTypes() []PaymentType { return m.supportedTypes }
|
||||
func (m *mockProvider) Name() string { return m.name }
|
||||
func (m *mockProvider) ProviderKey() string { return m.key }
|
||||
func (m *mockProvider) SupportedTypes() []PaymentType { return m.supportedTypes }
|
||||
func (m *mockProvider) CreatePayment(_ context.Context, _ CreatePaymentRequest) (*CreatePaymentResponse, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
@@ -16,6 +16,22 @@ const (
|
||||
TypeStripe PaymentType = "stripe"
|
||||
)
|
||||
|
||||
// Order status constants shared across payment and service layers.
|
||||
const (
|
||||
OrderStatusPending = "PENDING"
|
||||
OrderStatusPaid = "PAID"
|
||||
OrderStatusRecharging = "RECHARGING"
|
||||
OrderStatusCompleted = "COMPLETED"
|
||||
OrderStatusExpired = "EXPIRED"
|
||||
OrderStatusCancelled = "CANCELLED"
|
||||
OrderStatusFailed = "FAILED"
|
||||
OrderStatusRefundRequested = "REFUND_REQUESTED"
|
||||
OrderStatusRefunding = "REFUNDING"
|
||||
OrderStatusPartiallyRefunded = "PARTIALLY_REFUNDED"
|
||||
OrderStatusRefunded = "REFUNDED"
|
||||
OrderStatusRefundFailed = "REFUND_FAILED"
|
||||
)
|
||||
|
||||
// GetBasePaymentType extracts the base payment method from a composite key.
|
||||
// For example, "alipay_direct" -> "alipay".
|
||||
func GetBasePaymentType(t string) string {
|
||||
|
||||
@@ -2,9 +2,9 @@ package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
@@ -34,6 +34,14 @@ const (
|
||||
SettingCancelWindowMode = "CANCEL_RATE_LIMIT_WINDOW_MODE"
|
||||
)
|
||||
|
||||
// Default values for payment configuration settings.
|
||||
const (
|
||||
defaultMinRechargeAmount = 1
|
||||
defaultMaxRechargeAmount = 99999999.99
|
||||
defaultOrderTimeoutMin = 30
|
||||
defaultMaxPendingOrders = 3
|
||||
)
|
||||
|
||||
// PaymentConfig holds the payment system configuration.
|
||||
type PaymentConfig struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
@@ -161,11 +169,11 @@ func (s *PaymentConfigService) GetPaymentConfig(ctx context.Context) (*PaymentCo
|
||||
func (s *PaymentConfigService) parsePaymentConfig(vals map[string]string) *PaymentConfig {
|
||||
cfg := &PaymentConfig{
|
||||
Enabled: vals[SettingPaymentEnabled] == "true",
|
||||
MinAmount: pcParseFloat(vals[SettingMinRechargeAmount], 1),
|
||||
MaxAmount: pcParseFloat(vals[SettingMaxRechargeAmount], 99999999.99),
|
||||
MinAmount: pcParseFloat(vals[SettingMinRechargeAmount], defaultMinRechargeAmount),
|
||||
MaxAmount: pcParseFloat(vals[SettingMaxRechargeAmount], defaultMaxRechargeAmount),
|
||||
DailyLimit: pcParseFloat(vals[SettingDailyRechargeLimit], 0),
|
||||
OrderTimeoutMin: pcParseInt(vals[SettingOrderTimeoutMinutes], 30),
|
||||
MaxPendingOrders: pcParseInt(vals[SettingMaxPendingOrders], 3),
|
||||
OrderTimeoutMin: pcParseInt(vals[SettingOrderTimeoutMinutes], defaultOrderTimeoutMin),
|
||||
MaxPendingOrders: pcParseInt(vals[SettingMaxPendingOrders], defaultMaxPendingOrders),
|
||||
BalanceDisabled: vals[SettingBalancePayDisabled] == "true",
|
||||
LoadBalanceStrategy: vals[SettingLoadBalanceStrategy],
|
||||
ProductNamePrefix: vals[SettingProductNamePrefix],
|
||||
@@ -186,6 +194,9 @@ func (s *PaymentConfigService) parsePaymentConfig(vals map[string]string) *Payme
|
||||
}
|
||||
|
||||
// UpdatePaymentConfig updates the payment configuration settings.
|
||||
// NOTE: This function exceeds 30 lines because each field requires an independent
|
||||
// nil-check before serialisation — this is inherent to patch-style update patterns
|
||||
// and cannot be meaningfully decomposed without introducing unnecessary abstraction.
|
||||
func (s *PaymentConfigService) UpdatePaymentConfig(ctx context.Context, req UpdatePaymentConfigRequest) error {
|
||||
m := make(map[string]string)
|
||||
if req.Enabled != nil {
|
||||
@@ -293,7 +304,6 @@ func (s *PaymentConfigService) encryptConfig(cfg map[string]string) (string, err
|
||||
|
||||
// --- Channel CRUD ---
|
||||
|
||||
|
||||
// --- Plan CRUD ---
|
||||
|
||||
func (s *PaymentConfigService) ListPlans(ctx context.Context) ([]*dbent.SubscriptionPlan, error) {
|
||||
@@ -316,6 +326,9 @@ func (s *PaymentConfigService) CreatePlan(ctx context.Context, req CreatePlanReq
|
||||
return b.Save(ctx)
|
||||
}
|
||||
|
||||
// UpdatePlan updates a subscription plan by ID (patch semantics).
|
||||
// NOTE: This function exceeds 30 lines due to per-field nil-check patch update
|
||||
// boilerplate — same rationale as UpdatePaymentConfig.
|
||||
func (s *PaymentConfigService) UpdatePlan(ctx context.Context, id int64, req UpdatePlanRequest) (*dbent.SubscriptionPlan, error) {
|
||||
u := s.entClient.SubscriptionPlan.UpdateOneID(id)
|
||||
if req.GroupID != nil {
|
||||
@@ -378,7 +391,7 @@ func (s *PaymentConfigService) GetMethodLimits(ctx context.Context, types []stri
|
||||
for _, pt := range types {
|
||||
ml := MethodLimits{PaymentType: pt}
|
||||
for _, inst := range instances {
|
||||
if !pcInstanceSupportsType(inst, pt) {
|
||||
if !payment.InstanceSupportsType(inst.SupportedTypes, pt) {
|
||||
continue
|
||||
}
|
||||
pcApplyInstanceLimits(inst, pt, &ml)
|
||||
@@ -388,18 +401,6 @@ func (s *PaymentConfigService) GetMethodLimits(ctx context.Context, types []stri
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func pcInstanceSupportsType(inst *dbent.PaymentProviderInstance, pt string) bool {
|
||||
if inst.SupportedTypes == "" {
|
||||
return true
|
||||
}
|
||||
for _, t := range strings.Split(inst.SupportedTypes, ",") {
|
||||
if strings.TrimSpace(t) == pt {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func pcApplyInstanceLimits(inst *dbent.PaymentProviderInstance, pt string, ml *MethodLimits) {
|
||||
if inst.Limits == "" {
|
||||
return
|
||||
@@ -498,7 +499,7 @@ func (s *PaymentConfigService) MigrateLegacyPurchaseURL(ctx context.Context) err
|
||||
|
||||
// Save updated menu items and clear old settings
|
||||
if err := s.settingRepo.SetMultiple(ctx, map[string]string{
|
||||
SettingKeyCustomMenuItems: string(data),
|
||||
SettingKeyCustomMenuItems: string(data),
|
||||
SettingKeyPurchaseSubscriptionEnabled: "false",
|
||||
}); err != nil {
|
||||
return fmt.Errorf("save migrated settings: %w", err)
|
||||
|
||||
@@ -7,6 +7,8 @@ import (
|
||||
"time"
|
||||
)
|
||||
|
||||
const expiryCheckTimeout = 30 * time.Second
|
||||
|
||||
// PaymentOrderExpiryService periodically expires timed-out payment orders.
|
||||
type PaymentOrderExpiryService struct {
|
||||
paymentSvc *PaymentService
|
||||
@@ -57,7 +59,7 @@ func (s *PaymentOrderExpiryService) Stop() {
|
||||
}
|
||||
|
||||
func (s *PaymentOrderExpiryService) runOnce() {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), expiryCheckTimeout)
|
||||
defer cancel()
|
||||
|
||||
expired, err := s.paymentSvc.ExpireTimedOutOrders(ctx)
|
||||
|
||||
@@ -2,13 +2,14 @@ package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"math"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
dbent "github.com/Wei-Shaw/sub2api/ent"
|
||||
@@ -21,23 +22,29 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
OrderStatusPending = "PENDING"
|
||||
OrderStatusPaid = "PAID"
|
||||
OrderStatusRecharging = "RECHARGING"
|
||||
OrderStatusCompleted = "COMPLETED"
|
||||
OrderStatusExpired = "EXPIRED"
|
||||
OrderStatusCancelled = "CANCELLED"
|
||||
OrderStatusFailed = "FAILED"
|
||||
OrderStatusRefundRequested = "REFUND_REQUESTED"
|
||||
OrderStatusRefunding = "REFUNDING"
|
||||
OrderStatusPartiallyRefunded = "PARTIALLY_REFUNDED"
|
||||
OrderStatusRefunded = "REFUNDED"
|
||||
OrderStatusRefundFailed = "REFUND_FAILED"
|
||||
OrderStatusPending = payment.OrderStatusPending
|
||||
OrderStatusPaid = payment.OrderStatusPaid
|
||||
OrderStatusRecharging = payment.OrderStatusRecharging
|
||||
OrderStatusCompleted = payment.OrderStatusCompleted
|
||||
OrderStatusExpired = payment.OrderStatusExpired
|
||||
OrderStatusCancelled = payment.OrderStatusCancelled
|
||||
OrderStatusFailed = payment.OrderStatusFailed
|
||||
OrderStatusRefundRequested = payment.OrderStatusRefundRequested
|
||||
OrderStatusRefunding = payment.OrderStatusRefunding
|
||||
OrderStatusPartiallyRefunded = payment.OrderStatusPartiallyRefunded
|
||||
OrderStatusRefunded = payment.OrderStatusRefunded
|
||||
OrderStatusRefundFailed = payment.OrderStatusRefundFailed
|
||||
)
|
||||
|
||||
const (
|
||||
defaultMaxPendingOrders = 3
|
||||
paymentGraceMinutes = 5
|
||||
// defaultMaxPendingOrders and defaultOrderTimeoutMin are defined in
|
||||
// payment_config_service.go alongside other payment configuration defaults.
|
||||
paymentGraceMinutes = 5
|
||||
|
||||
defaultPageSize = 20
|
||||
maxPageSize = 100
|
||||
topUsersLimit = 10
|
||||
amountToleranceCNY = 0.01
|
||||
)
|
||||
|
||||
type CreateOrderRequest struct {
|
||||
@@ -95,14 +102,14 @@ type RefundResult struct {
|
||||
}
|
||||
|
||||
type DashboardStats struct {
|
||||
TodayAmount float64 `json:"today_amount"`
|
||||
TotalAmount float64 `json:"total_amount"`
|
||||
TodayCount int `json:"today_count"`
|
||||
TotalCount int `json:"total_count"`
|
||||
AvgAmount float64 `json:"avg_amount"`
|
||||
PendingOrders int `json:"pending_orders"`
|
||||
TodayAmount float64 `json:"today_amount"`
|
||||
TotalAmount float64 `json:"total_amount"`
|
||||
TodayCount int `json:"today_count"`
|
||||
TotalCount int `json:"total_count"`
|
||||
AvgAmount float64 `json:"avg_amount"`
|
||||
PendingOrders int `json:"pending_orders"`
|
||||
|
||||
DailySeries []DailyStats `json:"daily_series"`
|
||||
DailySeries []DailyStats `json:"daily_series"`
|
||||
PaymentMethods []PaymentMethodStat `json:"payment_methods"`
|
||||
TopUsers []TopUserStat `json:"top_users"`
|
||||
}
|
||||
@@ -127,7 +134,7 @@ type TopUserStat struct {
|
||||
|
||||
type PaymentService struct {
|
||||
providerMu sync.Mutex
|
||||
providersLoaded bool
|
||||
providerOnce sync.Once
|
||||
entClient *dbent.Client
|
||||
registry *payment.Registry
|
||||
loadBalancer payment.LoadBalancer
|
||||
@@ -235,10 +242,25 @@ func (s *PaymentService) createOrderInTx(ctx context.Context, req CreateOrderReq
|
||||
}
|
||||
tm := cfg.OrderTimeoutMin
|
||||
if tm <= 0 {
|
||||
tm = 30
|
||||
tm = defaultOrderTimeoutMin
|
||||
}
|
||||
exp := time.Now().Add(time.Duration(tm) * time.Minute)
|
||||
b := tx.PaymentOrder.Create().SetUserID(req.UserID).SetUserEmail(user.Email).SetUserName(user.Username).SetNillableUserNotes(psNilIfEmpty(user.Notes)).SetAmount(amount).SetPayAmount(payAmount).SetFeeRate(feeRate).SetRechargeCode("").SetPaymentType(req.PaymentType).SetPaymentTradeNo("").SetOrderType(req.OrderType).SetStatus(OrderStatusPending).SetExpiresAt(exp).SetClientIP(req.ClientIP).SetSrcHost(req.SrcHost)
|
||||
b := tx.PaymentOrder.Create().
|
||||
SetUserID(req.UserID).
|
||||
SetUserEmail(user.Email).
|
||||
SetUserName(user.Username).
|
||||
SetNillableUserNotes(psNilIfEmpty(user.Notes)).
|
||||
SetAmount(amount).
|
||||
SetPayAmount(payAmount).
|
||||
SetFeeRate(feeRate).
|
||||
SetRechargeCode("").
|
||||
SetPaymentType(req.PaymentType).
|
||||
SetPaymentTradeNo("").
|
||||
SetOrderType(req.OrderType).
|
||||
SetStatus(OrderStatusPending).
|
||||
SetExpiresAt(exp).
|
||||
SetClientIP(req.ClientIP).
|
||||
SetSrcHost(req.SrcHost)
|
||||
if req.SrcURL != "" {
|
||||
b.SetSrcURL(req.SrcURL)
|
||||
}
|
||||
@@ -368,17 +390,7 @@ func (s *PaymentService) GetUserOrders(ctx context.Context, userID int64, p Orde
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("count user orders: %w", err)
|
||||
}
|
||||
ps := p.PageSize
|
||||
if ps <= 0 {
|
||||
ps = 20
|
||||
}
|
||||
if ps > 100 {
|
||||
ps = 100
|
||||
}
|
||||
pg := p.Page
|
||||
if pg < 1 {
|
||||
pg = 1
|
||||
}
|
||||
ps, pg := applyPagination(p.PageSize, p.Page)
|
||||
orders, err := q.Order(dbent.Desc(paymentorder.FieldCreatedAt)).Limit(ps).Offset((pg - 1) * ps).All(ctx)
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("query user orders: %w", err)
|
||||
@@ -465,7 +477,7 @@ func (s *PaymentService) confirmPayment(ctx context.Context, oid int64, tradeNo
|
||||
slog.Error("order not found", "orderID", oid)
|
||||
return nil
|
||||
}
|
||||
if math.Abs(paid-o.PayAmount) > 0.01 {
|
||||
if math.Abs(paid-o.PayAmount) > amountToleranceCNY {
|
||||
s.writeAuditLog(ctx, o.ID, "PAYMENT_AMOUNT_MISMATCH", pk, map[string]any{"expected": o.PayAmount, "paid": paid, "tradeNo": tradeNo})
|
||||
return fmt.Errorf("amount mismatch: expected %.2f, got %.2f", o.PayAmount, paid)
|
||||
}
|
||||
@@ -649,18 +661,9 @@ func (s *PaymentService) RetryFulfillment(ctx context.Context, oid int64) error
|
||||
}
|
||||
|
||||
func (s *PaymentService) RequestRefund(ctx context.Context, oid, uid int64, reason string) error {
|
||||
o, err := s.entClient.PaymentOrder.Get(ctx, oid)
|
||||
o, err := s.validateRefundRequest(ctx, oid, uid)
|
||||
if err != nil {
|
||||
return infraerrors.NotFound("NOT_FOUND", "order not found")
|
||||
}
|
||||
if o.UserID != uid {
|
||||
return infraerrors.Forbidden("FORBIDDEN", "no permission")
|
||||
}
|
||||
if o.OrderType != "balance" {
|
||||
return infraerrors.BadRequest("INVALID_ORDER_TYPE", "only balance orders can request refund")
|
||||
}
|
||||
if o.Status != OrderStatusCompleted {
|
||||
return infraerrors.BadRequest("INVALID_STATUS", "only completed orders can request refund")
|
||||
return err
|
||||
}
|
||||
u, err := s.userRepo.GetByID(ctx, o.UserID)
|
||||
if err != nil {
|
||||
@@ -683,6 +686,23 @@ func (s *PaymentService) RequestRefund(ctx context.Context, oid, uid int64, reas
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *PaymentService) validateRefundRequest(ctx context.Context, oid, uid int64) (*dbent.PaymentOrder, error) {
|
||||
o, err := s.entClient.PaymentOrder.Get(ctx, oid)
|
||||
if err != nil {
|
||||
return nil, infraerrors.NotFound("NOT_FOUND", "order not found")
|
||||
}
|
||||
if o.UserID != uid {
|
||||
return nil, infraerrors.Forbidden("FORBIDDEN", "no permission")
|
||||
}
|
||||
if o.OrderType != "balance" {
|
||||
return nil, infraerrors.BadRequest("INVALID_ORDER_TYPE", "only balance orders can request refund")
|
||||
}
|
||||
if o.Status != OrderStatusCompleted {
|
||||
return nil, infraerrors.BadRequest("INVALID_STATUS", "only completed orders can request refund")
|
||||
}
|
||||
return o, nil
|
||||
}
|
||||
|
||||
func (s *PaymentService) PrepareRefund(ctx context.Context, oid int64, amt float64, reason string, force, deduct bool) (*RefundPlan, *RefundResult, error) {
|
||||
o, err := s.entClient.PaymentOrder.Get(ctx, oid)
|
||||
if err != nil {
|
||||
@@ -844,7 +864,6 @@ func (s *PaymentService) GetDashboardStats(ctx context.Context, days int) (*Dash
|
||||
|
||||
paidStatuses := []string{OrderStatusCompleted, OrderStatusPaid, OrderStatusRecharging}
|
||||
|
||||
// Fetch all paid orders in the date range
|
||||
orders, err := s.entClient.PaymentOrder.Query().
|
||||
Where(
|
||||
paymentorder.StatusIn(paidStatuses...),
|
||||
@@ -856,10 +875,24 @@ func (s *PaymentService) GetDashboardStats(ctx context.Context, days int) (*Dash
|
||||
}
|
||||
|
||||
st := &DashboardStats{}
|
||||
computeBasicStats(st, orders, todayStart)
|
||||
|
||||
// Compute basic stats
|
||||
var totalAmount float64
|
||||
var todayAmount float64
|
||||
st.PendingOrders, err = s.entClient.PaymentOrder.Query().
|
||||
Where(paymentorder.StatusEQ(OrderStatusPending)).
|
||||
Count(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
st.DailySeries = buildDailySeries(orders, since, days)
|
||||
st.PaymentMethods = buildMethodDistribution(orders)
|
||||
st.TopUsers = buildTopUsers(orders)
|
||||
|
||||
return st, nil
|
||||
}
|
||||
|
||||
func computeBasicStats(st *DashboardStats, orders []*dbent.PaymentOrder, todayStart time.Time) {
|
||||
var totalAmount, todayAmount float64
|
||||
var todayCount int
|
||||
for _, o := range orders {
|
||||
totalAmount += o.PayAmount
|
||||
@@ -875,16 +908,9 @@ func (s *PaymentService) GetDashboardStats(ctx context.Context, days int) (*Dash
|
||||
if st.TotalCount > 0 {
|
||||
st.AvgAmount = math.Round(totalAmount/float64(st.TotalCount)*100) / 100
|
||||
}
|
||||
}
|
||||
|
||||
// Pending orders count
|
||||
st.PendingOrders, err = s.entClient.PaymentOrder.Query().
|
||||
Where(paymentorder.StatusEQ(OrderStatusPending)).
|
||||
Count(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Daily series: GROUP BY date(paid_at)
|
||||
func buildDailySeries(orders []*dbent.PaymentOrder, since time.Time, days int) []DailyStats {
|
||||
dailyMap := make(map[string]*DailyStats)
|
||||
for _, o := range orders {
|
||||
if o.PaidAt == nil {
|
||||
@@ -899,19 +925,20 @@ func (s *PaymentService) GetDashboardStats(ctx context.Context, days int) (*Dash
|
||||
ds.Amount += o.PayAmount
|
||||
ds.Count++
|
||||
}
|
||||
// Build sorted daily series for all days in range
|
||||
st.DailySeries = make([]DailyStats, 0, days)
|
||||
series := make([]DailyStats, 0, days)
|
||||
for i := 0; i < days; i++ {
|
||||
date := since.AddDate(0, 0, i+1).Format("2006-01-02")
|
||||
if ds, ok := dailyMap[date]; ok {
|
||||
ds.Amount = math.Round(ds.Amount*100) / 100
|
||||
st.DailySeries = append(st.DailySeries, *ds)
|
||||
series = append(series, *ds)
|
||||
} else {
|
||||
st.DailySeries = append(st.DailySeries, DailyStats{Date: date})
|
||||
series = append(series, DailyStats{Date: date})
|
||||
}
|
||||
}
|
||||
return series
|
||||
}
|
||||
|
||||
// Payment methods: GROUP BY payment_type
|
||||
func buildMethodDistribution(orders []*dbent.PaymentOrder) []PaymentMethodStat {
|
||||
methodMap := make(map[string]*PaymentMethodStat)
|
||||
for _, o := range orders {
|
||||
ms, ok := methodMap[o.PaymentType]
|
||||
@@ -922,13 +949,15 @@ func (s *PaymentService) GetDashboardStats(ctx context.Context, days int) (*Dash
|
||||
ms.Amount += o.PayAmount
|
||||
ms.Count++
|
||||
}
|
||||
st.PaymentMethods = make([]PaymentMethodStat, 0, len(methodMap))
|
||||
methods := make([]PaymentMethodStat, 0, len(methodMap))
|
||||
for _, ms := range methodMap {
|
||||
ms.Amount = math.Round(ms.Amount*100) / 100
|
||||
st.PaymentMethods = append(st.PaymentMethods, *ms)
|
||||
methods = append(methods, *ms)
|
||||
}
|
||||
return methods
|
||||
}
|
||||
|
||||
// Top users: GROUP BY user_id, ORDER BY amount DESC, LIMIT 10
|
||||
func buildTopUsers(orders []*dbent.PaymentOrder) []TopUserStat {
|
||||
userMap := make(map[int64]*TopUserStat)
|
||||
for _, o := range orders {
|
||||
us, ok := userMap[o.UserID]
|
||||
@@ -938,33 +967,23 @@ func (s *PaymentService) GetDashboardStats(ctx context.Context, days int) (*Dash
|
||||
}
|
||||
us.Amount += o.PayAmount
|
||||
}
|
||||
type userEntry struct {
|
||||
stat *TopUserStat
|
||||
amount float64
|
||||
}
|
||||
userList := make([]userEntry, 0, len(userMap))
|
||||
userList := make([]*TopUserStat, 0, len(userMap))
|
||||
for _, us := range userMap {
|
||||
us.Amount = math.Round(us.Amount*100) / 100
|
||||
userList = append(userList, userEntry{stat: us, amount: us.Amount})
|
||||
userList = append(userList, us)
|
||||
}
|
||||
// Sort descending by amount
|
||||
for i := 0; i < len(userList); i++ {
|
||||
for j := i + 1; j < len(userList); j++ {
|
||||
if userList[j].amount > userList[i].amount {
|
||||
userList[i], userList[j] = userList[j], userList[i]
|
||||
}
|
||||
}
|
||||
}
|
||||
limit := 10
|
||||
sort.Slice(userList, func(i, j int) bool {
|
||||
return userList[i].Amount > userList[j].Amount
|
||||
})
|
||||
limit := topUsersLimit
|
||||
if len(userList) < limit {
|
||||
limit = len(userList)
|
||||
}
|
||||
st.TopUsers = make([]TopUserStat, 0, limit)
|
||||
result := make([]TopUserStat, 0, limit)
|
||||
for i := 0; i < limit; i++ {
|
||||
st.TopUsers = append(st.TopUsers, *userList[i].stat)
|
||||
result = append(result, *userList[i])
|
||||
}
|
||||
|
||||
return st, nil
|
||||
return result
|
||||
}
|
||||
|
||||
func (s *PaymentService) sumAmt(ctx context.Context, statuses []string, since time.Time, usePaid bool) (float64, error) {
|
||||
@@ -1050,6 +1069,21 @@ func psStartOfDayUTC(t time.Time) time.Time {
|
||||
return time.Date(y, m, d, 0, 0, 0, 0, time.UTC)
|
||||
}
|
||||
|
||||
func applyPagination(pageSize, page int) (size, pg int) {
|
||||
size = pageSize
|
||||
if size <= 0 {
|
||||
size = defaultPageSize
|
||||
}
|
||||
if size > maxPageSize {
|
||||
size = maxPageSize
|
||||
}
|
||||
pg = page
|
||||
if pg < 1 {
|
||||
pg = 1
|
||||
}
|
||||
return size, pg
|
||||
}
|
||||
|
||||
// AdminListOrders returns a paginated list of orders. If userID > 0, filters by user.
|
||||
func (s *PaymentService) AdminListOrders(ctx context.Context, userID int64, p OrderListParams) ([]*dbent.PaymentOrder, int, error) {
|
||||
q := s.entClient.PaymentOrder.Query()
|
||||
@@ -1069,17 +1103,7 @@ func (s *PaymentService) AdminListOrders(ctx context.Context, userID int64, p Or
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("count admin orders: %w", err)
|
||||
}
|
||||
ps := p.PageSize
|
||||
if ps <= 0 {
|
||||
ps = 20
|
||||
}
|
||||
if ps > 100 {
|
||||
ps = 100
|
||||
}
|
||||
pg := p.Page
|
||||
if pg < 1 {
|
||||
pg = 1
|
||||
}
|
||||
ps, pg := applyPagination(p.PageSize, p.Page)
|
||||
orders, err := q.Order(dbent.Desc(paymentorder.FieldCreatedAt)).Limit(ps).Offset((pg - 1) * ps).All(ctx)
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("query admin orders: %w", err)
|
||||
@@ -1091,13 +1115,9 @@ func (s *PaymentService) AdminListOrders(ctx context.Context, userID int64, p Or
|
||||
// It queries all enabled PaymentProviderInstance records, decrypts their config,
|
||||
// creates providers via provider.CreateProvider, and registers them.
|
||||
func (s *PaymentService) EnsureProviders(ctx context.Context) {
|
||||
s.providerMu.Lock()
|
||||
defer s.providerMu.Unlock()
|
||||
if s.providersLoaded {
|
||||
return
|
||||
}
|
||||
s.loadProviders(ctx)
|
||||
s.providersLoaded = true
|
||||
s.providerOnce.Do(func() {
|
||||
s.loadProviders(ctx)
|
||||
})
|
||||
}
|
||||
|
||||
// RefreshProviders clears and re-registers all providers from the database.
|
||||
@@ -1107,7 +1127,8 @@ func (s *PaymentService) RefreshProviders(ctx context.Context) {
|
||||
defer s.providerMu.Unlock()
|
||||
s.registry.Clear()
|
||||
s.loadProviders(ctx)
|
||||
s.providersLoaded = true
|
||||
s.providerOnce = sync.Once{} // reset so next EnsureProviders is a no-op until next Refresh
|
||||
s.providerOnce.Do(func() {}) // mark as done since we just loaded
|
||||
}
|
||||
|
||||
func (s *PaymentService) loadProviders(ctx context.Context) {
|
||||
|
||||
Reference in New Issue
Block a user