Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions internal/apps/payment/errs.go
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,8 @@ const (
OrderNotFound = "订单不存在或已完成"
OrderStatusInvalid = "订单状态不允许支付"
OrderExpired = "订单已过期"
OrderPayerMismatch = "当前用户不是该订单的预期付款人"
OrderRequestConflict = "同一业务订单号的订单信息不一致"
MerchantInfoNotFound = "商户信息不存在"
RecipientNotFound = "收款人不存在"
OrderNoFormatError = "订单号格式错误"
Expand Down
113 changes: 67 additions & 46 deletions internal/apps/payment/routers.go
Original file line number Diff line number Diff line change
Expand Up @@ -117,48 +117,46 @@ 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 {
return err
}

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)
}

payURL = fmt.Sprintf("%s?order_no=%s", config.Config.App.FrontendPayURL, url.QueryEscape(encryptString))
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
}

Expand Down Expand Up @@ -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,
},
}))
Expand All @@ -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
Expand All @@ -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

Expand All @@ -548,7 +568,6 @@ func PayMerchantOrder(c *gin.Context) {

// 更新订单状态
order.Status = model.OrderStatusSuccess
order.PayerUserID = orderCtx.CurrentUser.ID
order.TradeTime = time.Now()

if isTestMode {
Expand Down Expand Up @@ -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:
Expand Down
70 changes: 70 additions & 0 deletions internal/apps/payment/utils.go
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@ import (
"sort"
"strconv"
"strings"
"time"

"github.com/gin-gonic/gin"
"github.com/linux-do/credit/internal/apps/oauth"
Expand All @@ -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 响应
Expand Down Expand Up @@ -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()
Expand Down
Loading