From 96b1a96cf9677a22d20fa830f0fabbe97a01c7bf Mon Sep 17 00:00:00 2001 From: erio Date: Mon, 6 Apr 2026 01:13:30 +0800 Subject: [PATCH] refactor(payment): code quality improvements per project conventions MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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 --- .../internal/handler/admin/payment_handler.go | 51 ++-- backend/internal/handler/payment_handler.go | 28 ++- .../handler/payment_webhook_handler.go | 5 +- backend/internal/payment/load_balancer.go | 12 +- .../internal/payment/load_balancer_test.go | 10 +- backend/internal/payment/provider/alipay.go | 66 ++--- backend/internal/payment/provider/easypay.go | 17 +- backend/internal/payment/provider/stripe.go | 25 +- backend/internal/payment/provider/wxpay.go | 83 +++++-- backend/internal/payment/registry.go | 5 +- backend/internal/payment/registry_test.go | 6 +- backend/internal/payment/types.go | 16 ++ .../service/payment_config_service.go | 41 +-- .../service/payment_order_expiry_service.go | 4 +- backend/internal/service/payment_service.go | 235 ++++++++++-------- 15 files changed, 359 insertions(+), 245 deletions(-) diff --git a/backend/internal/handler/admin/payment_handler.go b/backend/internal/handler/admin/payment_handler.go index 4a773e0c52..5079a35c83 100644 --- a/backend/internal/handler/admin/payment_handler.go +++ b/backend/internal/handler/admin/payment_handler.go @@ -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). diff --git a/backend/internal/handler/payment_handler.go b/backend/internal/handler/payment_handler.go index cee80e72b8..6172dae7b0 100644 --- a/backend/internal/handler/payment_handler.go +++ b/backend/internal/handler/payment_handler.go @@ -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")) diff --git a/backend/internal/handler/payment_webhook_handler.go b/backend/internal/handler/payment_webhook_handler.go index 09650081b7..60c2d16493 100644 --- a/backend/internal/handler/payment_webhook_handler.go +++ b/backend/internal/handler/payment_webhook_handler.go @@ -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") diff --git a/backend/internal/payment/load_balancer.go b/backend/internal/payment/load_balancer.go index 1a7371138c..4a53d3d5ee 100644 --- a/backend/internal/payment/load_balancer.go +++ b/backend/internal/payment/load_balancer.go @@ -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 diff --git a/backend/internal/payment/load_balancer_test.go b/backend/internal/payment/load_balancer_test.go index 78316cc4e6..e8bc4c8652 100644 --- a/backend/internal/payment/load_balancer_test.go +++ b/backend/internal/payment/load_balancer_test.go @@ -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) } }) } diff --git a/backend/internal/payment/provider/alipay.go b/backend/internal/payment/provider/alipay.go index 55c119b232..6ce599f275 100644 --- a/backend/internal/payment/provider/alipay.go +++ b/backend/internal/payment/provider/alipay.go @@ -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, diff --git a/backend/internal/payment/provider/easypay.go b/backend/internal/payment/provider/easypay.go index 482712cde4..f9d36c4786 100644 --- a/backend/internal/payment/provider/easypay.go +++ b/backend/internal/payment/provider/easypay.go @@ -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 diff --git a/backend/internal/payment/provider/stripe.go b/backend/internal/payment/provider/stripe.go index 850dfd88fa..8167325f00 100644 --- a/backend/internal/payment/provider/stripe.go +++ b/backend/internal/payment/provider/stripe.go @@ -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 diff --git a/backend/internal/payment/provider/wxpay.go b/backend/internal/payment/provider/wxpay.go index e1946b7c40..96b252c35e 100644 --- a/backend/internal/payment/provider/wxpay.go +++ b/backend/internal/payment/provider/wxpay.go @@ -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 { diff --git a/backend/internal/payment/registry.go b/backend/internal/payment/registry.go index 77fb61907e..259eb4bb33 100644 --- a/backend/internal/payment/registry.go +++ b/backend/internal/payment/registry.go @@ -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 } diff --git a/backend/internal/payment/registry_test.go b/backend/internal/payment/registry_test.go index 6acfeb6abf..9684945c76 100644 --- a/backend/internal/payment/registry_test.go +++ b/backend/internal/payment/registry_test.go @@ -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 } diff --git a/backend/internal/payment/types.go b/backend/internal/payment/types.go index 40d44418a6..f2173c1a6b 100644 --- a/backend/internal/payment/types.go +++ b/backend/internal/payment/types.go @@ -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 { diff --git a/backend/internal/service/payment_config_service.go b/backend/internal/service/payment_config_service.go index d617563b49..3936a4ccfa 100644 --- a/backend/internal/service/payment_config_service.go +++ b/backend/internal/service/payment_config_service.go @@ -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) diff --git a/backend/internal/service/payment_order_expiry_service.go b/backend/internal/service/payment_order_expiry_service.go index 850d8f737d..b0cda3e591 100644 --- a/backend/internal/service/payment_order_expiry_service.go +++ b/backend/internal/service/payment_order_expiry_service.go @@ -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) diff --git a/backend/internal/service/payment_service.go b/backend/internal/service/payment_service.go index f4e2cfff65..c1d8a9a342 100644 --- a/backend/internal/service/payment_service.go +++ b/backend/internal/service/payment_service.go @@ -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) {