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:
erio
2026-04-06 01:13:30 +08:00
parent c69929bbb7
commit 96b1a96cf9
15 changed files with 359 additions and 245 deletions
@@ -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).
+17 -11
View File
@@ -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")
+8 -4
View File
@@ -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)
}
})
}
+38 -28
View File
@@ -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,
+12 -5
View File
@@ -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
+19 -6
View File
@@ -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
+57 -26
View File
@@ -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 {
+1 -4
View File
@@ -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
}
+3 -3
View File
@@ -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
View File
@@ -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)
+128 -107
View File
@@ -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) {