diff --git a/internal/apps/payment/errs.go b/internal/apps/payment/errs.go index f082b23..1e223fa 100644 --- a/internal/apps/payment/errs.go +++ b/internal/apps/payment/errs.go @@ -20,6 +20,8 @@ const ( OrderNotFound = "订单不存在或已完成" OrderStatusInvalid = "订单状态不允许支付" OrderExpired = "订单已过期" + OrderPayerMismatch = "当前用户不是该订单的预期付款人" + OrderRequestConflict = "同一业务订单号的订单信息不一致" MerchantInfoNotFound = "商户信息不存在" RecipientNotFound = "收款人不存在" OrderNoFormatError = "订单号格式错误" diff --git a/internal/apps/payment/routers.go b/internal/apps/payment/routers.go index b6d6fa4..32f25b6 100644 --- a/internal/apps/payment/routers.go +++ b/internal/apps/payment/routers.go @@ -117,27 +117,18 @@ func CreateMerchantOrder(c *gin.Context) { } var payURL string + expiresAt := time.Now().Add(time.Duration(expireMinutes) * time.Minute) if err := db.DB(c.Request.Context()).Transaction( func(tx *gorm.DB) error { - // 创建订单 - order := model.Order{ - OrderName: req.OrderName, - ClientID: apiKey.ClientID, - MerchantOrderNo: req.MerchantOrderNo, - PayeeUserID: merchantUser.ID, - Amount: req.Amount, - Status: model.OrderStatusPending, - Type: model.OrderTypePayment, - Remark: req.Remark, - PaymentType: req.PaymentType, - RedirectURI: req.ReturnURL, - NotifyURL: req.NotifyURL, - ExpiresAt: time.Now().Add(time.Duration(expireMinutes) * time.Minute), - } - if err := tx.Create(&order).Error; err != nil { + order, err := createOrReuseMerchantOrder(tx, req, apiKey, merchantUser.ID, expiresAt) + if err != nil { return err } + remainingTTL := time.Until(order.ExpiresAt) + if remainingTTL <= 0 { + return errors.New(OrderExpired) + } encryptString, err := util.Encrypt(merchantUser.SignKey, strconv.FormatUint(order.ID, 10)) if err != nil { @@ -145,12 +136,12 @@ func CreateMerchantOrder(c *gin.Context) { } merchantIDStr := strconv.FormatUint(merchantUser.ID, 10) - if errSet := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OrderMerchantIDCacheKeyFormat, encryptString)), merchantIDStr, time.Duration(expireMinutes)*time.Minute).Err(); errSet != nil { + if errSet := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OrderMerchantIDCacheKeyFormat, encryptString)), merchantIDStr, remainingTTL).Err(); errSet != nil { return fmt.Errorf("failed to set redis key: %w", errSet) } expireKey := db.PrefixedKey(fmt.Sprintf(OrderExpireKeyFormat, order.ID)) - if errSet := db.Redis.Set(c.Request.Context(), expireKey, order.ID, time.Duration(expireMinutes)*time.Minute).Err(); errSet != nil { + if errSet := db.Redis.Set(c.Request.Context(), expireKey, order.ID, remainingTTL).Err(); errSet != nil { return fmt.Errorf("failed to set order expire key: %w", errSet) } @@ -158,7 +149,14 @@ func CreateMerchantOrder(c *gin.Context) { return nil }, ); err != nil { - c.JSON(http.StatusInternalServerError, util.Err(err.Error())) + switch err.Error() { + case OrderRequestConflict: + c.JSON(http.StatusConflict, util.Err(err.Error())) + case OrderStatusInvalid, OrderExpired: + c.JSON(http.StatusBadRequest, util.Err(err.Error())) + default: + c.JSON(http.StatusInternalServerError, util.Err(err.Error())) + } return } @@ -460,35 +458,56 @@ func GetPaymentPageDetails(c *gin.Context) { } var order model.Order - if err := db.DB(c.Request.Context()). - Select("orders.*, payee_user.username as payee_username"). - Joins("LEFT JOIN users as payee_user ON orders.payee_user_id = payee_user.id"). - Where("orders.id = ? AND orders.status = ?", orderCtx.OrderID, model.OrderStatusPending). - First(&order).Error; err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - c.JSON(http.StatusNotFound, util.Err(OrderNotFound)) - return + if err := db.DB(c.Request.Context()).Transaction(func(tx *gorm.DB) error { + now := time.Now() + result := tx.Model(&model.Order{}). + Where("id = ? AND status = ? AND payer_user_id = ? AND expires_at > ?", + orderCtx.OrderID, model.OrderStatusPending, 0, now). + Update("payer_user_id", orderCtx.CurrentUser.ID) + if result.Error != nil { + return result.Error } - c.JSON(http.StatusInternalServerError, util.Err(err.Error())) - return - } - order.PayerUsername = orderCtx.CurrentUser.Username - var merchant model.MerchantAPIKey - if err := db.DB(c.Request.Context()). - Where("client_id = ?", order.ClientID). - First(&merchant).Error; err != nil { - c.JSON(http.StatusNotFound, util.Err(MerchantInfoNotFound)) + if err := tx. + Select("orders.*, payee_user.username as payee_username"). + Joins("LEFT JOIN users as payee_user ON orders.payee_user_id = payee_user.id"). + Where("orders.id = ? AND orders.status = ?", orderCtx.OrderID, model.OrderStatusPending). + First(&order).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return errors.New(OrderNotFound) + } + return err + } + if !order.ExpiresAt.After(time.Now()) { + return errors.New(OrderExpired) + } + if err := validateExpectedPayer(order.PayerUserID, orderCtx.CurrentUser.ID); err != nil { + return err + } + + order.PayerUsername = orderCtx.CurrentUser.Username + return nil + }); err != nil { + switch err.Error() { + case OrderNotFound: + c.JSON(http.StatusNotFound, util.Err(err.Error())) + case OrderExpired: + c.JSON(http.StatusBadRequest, util.Err(err.Error())) + case OrderPayerMismatch: + c.JSON(http.StatusForbidden, util.Err(err.Error())) + default: + c.JSON(http.StatusInternalServerError, util.Err(err.Error())) + } return } - redirectURI := cmp.Or(order.RedirectURI, merchant.RedirectURI) + redirectURI := cmp.Or(order.RedirectURI, orderCtx.MerchantAPIKey.RedirectURI) c.JSON(http.StatusOK, util.OK(GetOrderResponse{ Order: &order, FeeRate: orderCtx.MerchantPayConfig.FeeRate, Merchant: MerchantInfo{ - AppName: merchant.AppName, + AppName: orderCtx.MerchantAPIKey.AppName, RedirectURI: redirectURI, }, })) @@ -512,11 +531,6 @@ func PayMerchantOrder(c *gin.Context) { return } - if !orderCtx.CurrentUser.VerifyPayKey(req.PayKey) { - c.JSON(http.StatusBadRequest, util.Err(common.PayKeyIncorrect)) - return - } - if err := db.DB(c.Request.Context()).Transaction( func(tx *gorm.DB) error { var order model.Order @@ -530,9 +544,15 @@ func PayMerchantOrder(c *gin.Context) { } // 检查订单是否过期 - if order.ExpiresAt.Before(time.Now()) { + if !order.ExpiresAt.After(time.Now()) { return errors.New(OrderExpired) } + if err := validateExpectedPayer(order.PayerUserID, orderCtx.CurrentUser.ID); err != nil { + return err + } + if !orderCtx.CurrentUser.VerifyPayKey(req.PayKey) { + return errors.New(common.PayKeyIncorrect) + } isTestMode := orderCtx.MerchantAPIKey.TestMode @@ -548,7 +568,6 @@ func PayMerchantOrder(c *gin.Context) { // 更新订单状态 order.Status = model.OrderStatusSuccess - order.PayerUserID = orderCtx.CurrentUser.ID order.TradeTime = time.Now() if isTestMode { @@ -618,8 +637,10 @@ func PayMerchantOrder(c *gin.Context) { ); err != nil { errMsg := err.Error() switch errMsg { - case common.InsufficientBalance, OrderExpired, common.DailyLimitExceeded: + case common.InsufficientBalance, OrderExpired, common.DailyLimitExceeded, common.PayKeyIncorrect: c.JSON(http.StatusBadRequest, util.Err(errMsg)) + case OrderPayerMismatch: + c.JSON(http.StatusForbidden, util.Err(errMsg)) case OrderNotFound: c.JSON(http.StatusNotFound, util.Err(errMsg)) default: diff --git a/internal/apps/payment/utils.go b/internal/apps/payment/utils.go index 8e4a46b..74f0edd 100644 --- a/internal/apps/payment/utils.go +++ b/internal/apps/payment/utils.go @@ -25,6 +25,7 @@ import ( "sort" "strconv" "strings" + "time" "github.com/gin-gonic/gin" "github.com/linux-do/credit/internal/apps/oauth" @@ -35,6 +36,7 @@ import ( "github.com/linux-do/credit/internal/util" "github.com/redis/go-redis/v9" "gorm.io/gorm" + "gorm.io/gorm/clause" ) // HandleParseOrderNoError 处理 ParseOrderNo 返回的错误,返回对应的 HTTP 响应 @@ -69,6 +71,74 @@ type OrderContext struct { MerchantAPIKey *model.MerchantAPIKey } +// validateExpectedPayer 校验当前用户是否为订单绑定的预期付款人。 +func validateExpectedPayer(expectedPayerUserID, currentUserID uint64) error { + if expectedPayerUserID == 0 || expectedPayerUserID != currentUserID { + return errors.New(OrderPayerMismatch) + } + return nil +} + +// createOrReuseMerchantOrder 创建商户订单;幂等键冲突时复用原待支付订单。 +func createOrReuseMerchantOrder(tx *gorm.DB, req *CreateOrderRequest, apiKey *model.MerchantAPIKey, merchantUserID uint64, expiresAt time.Time) (*model.Order, error) { + order := model.Order{ + OrderName: req.OrderName, + ClientID: apiKey.ClientID, + MerchantOrderNo: req.MerchantOrderNo, + PayeeUserID: merchantUserID, + Amount: req.Amount, + Status: model.OrderStatusPending, + Type: model.OrderTypePayment, + Remark: req.Remark, + PaymentType: req.PaymentType, + RedirectURI: req.ReturnURL, + NotifyURL: req.NotifyURL, + ExpiresAt: expiresAt, + } + + result := tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "client_id"}, {Name: "merchant_order_no"}}, + DoNothing: true, + }).Create(&order) + if result.Error != nil { + return nil, result.Error + } + if result.RowsAffected > 0 { + return &order, nil + } + + // BeforeCreate 已为冲突请求生成新 ID,查询原订单前必须清空主键条件。 + order = model.Order{} + if err := tx.Where("client_id = ? AND merchant_order_no = ?", apiKey.ClientID, req.MerchantOrderNo). + First(&order).Error; err != nil { + return nil, err + } + if order.Status != model.OrderStatusPending { + return nil, errors.New(OrderStatusInvalid) + } + if !order.ExpiresAt.After(time.Now()) { + return nil, errors.New(OrderExpired) + } + + // 同一幂等键仅允许复用原始参数完全一致的订单,避免商户误用业务单号。 + merchantOrderNoMatches := (order.MerchantOrderNo == nil && req.MerchantOrderNo == nil) || + (order.MerchantOrderNo != nil && req.MerchantOrderNo != nil && *order.MerchantOrderNo == *req.MerchantOrderNo) + if order.ClientID != apiKey.ClientID || + !merchantOrderNoMatches || + order.OrderName != req.OrderName || + order.PayeeUserID != merchantUserID || + !order.Amount.Equal(req.Amount) || + order.Type != model.OrderTypePayment || + order.Remark != req.Remark || + order.PaymentType != req.PaymentType || + order.RedirectURI != req.ReturnURL || + order.NotifyURL != req.NotifyURL { + return nil, errors.New(OrderRequestConflict) + } + + return &order, nil +} + // ParseOrderNo 解析订单号,获取订单上下文信息 func ParseOrderNo(c *gin.Context, orderNo string) (*OrderContext, error) { merchantIDStr, errGet := db.Redis.Get(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OrderMerchantIDCacheKeyFormat, orderNo))).Result()